diff --git a/src/mlpack/methods/maxip/ip_metric_impl.hpp b/src/mlpack/methods/maxip/ip_metric_impl.hpp index eaa5f5a220..5c7e8e38d1 100644 --- a/src/mlpack/methods/maxip/ip_metric_impl.hpp +++ b/src/mlpack/methods/maxip/ip_metric_impl.hpp @@ -18,19 +18,21 @@ namespace maxip { template template -double IPMetric::Evaluate(const Vec1Type& a, const Vec2Type& b) +inline double IPMetric::Evaluate(const Vec1Type& a, const Vec2Type& b) { // This is the metric induced by the kernel function. // Maybe we can do better by caching some of this? + ++distanceEvaluations; return KernelType::Evaluate(a, a) + KernelType::Evaluate(b, b) - 2 * KernelType::Evaluate(a, b); } template<> template -double IPMetric::Evaluate(const Vec1Type& a, - const Vec2Type& b) +inline double IPMetric::Evaluate(const Vec1Type& a, + const Vec2Type& b) { + ++distanceEvaluations; return metric::LMetric<2>::Evaluate(a, b); } diff --git a/src/mlpack/methods/maxip/max_ip_impl.hpp b/src/mlpack/methods/maxip/max_ip_impl.hpp index 19d6b9715e..217fa513f3 100644 --- a/src/mlpack/methods/maxip/max_ip_impl.hpp +++ b/src/mlpack/methods/maxip/max_ip_impl.hpp @@ -91,6 +91,8 @@ void MaxIP::Search(const size_t k, Timer::Start("computing_products"); + size_t kernelEvaluations = 0; + // Naive implementation. if (naive) { @@ -101,6 +103,7 @@ void MaxIP::Search(const size_t k, { const double eval = KernelType::Evaluate(querySet.unsafe_col(q), referenceSet.unsafe_col(r)); + ++kernelEvaluations; size_t insertPosition; for (insertPosition = 0; insertPosition < indices.n_rows; @@ -114,6 +117,8 @@ void MaxIP::Search(const size_t k, } Timer::Stop("computing_products"); + + Log::Info << "Kernel evaluations: " << kernelEvaluations << "." << std::endl; return; } @@ -169,6 +174,7 @@ void MaxIP::Search(const size_t k, // Evaluate the kernel. Then see if it is a result to keep. eval = KernelType::Evaluate(querySet.unsafe_col(queryIndex), referenceSet.unsafe_col(referenceNode->Point())); + ++kernelEvaluations; // Is the result good enough to be saved? if (eval > products(products.n_rows - 1, queryIndex)) @@ -211,6 +217,8 @@ void MaxIP::Search(const size_t k, } Log::Info << "Pruned " << numPrunes << " nodes." << std::endl; + Log::Info << "Kernel evaluations: " << kernelEvaluations << "." << std::endl; + Log::Info << "Distance evaluations: " << distanceEvaluations << "." << std::endl; Timer::Stop("computing_products"); return; diff --git a/src/mlpack/methods/maxip/max_ip_main.cpp b/src/mlpack/methods/maxip/max_ip_main.cpp index 5d4a7c8c8c..ee6c474567 100644 --- a/src/mlpack/methods/maxip/max_ip_main.cpp +++ b/src/mlpack/methods/maxip/max_ip_main.cpp @@ -4,9 +4,15 @@ * * Main executable for maximum inner product search. */ +// I can't believe I'm doing this. +#include +size_t distanceEvaluations; + #include #include -#include 1 +#include +#include +#include #include "max_ip.hpp" @@ -35,8 +41,8 @@ PARAM_STRING("products_file", "File to save inner products into.", "p", ""); PARAM_STRING("indices_file", "File to save indices of inner products into.", "i", ""); -PARAM_STRING("kernel", "Kernel type to use: 'linear', 'polynomial'.", "K", - "linear"); +PARAM_STRING("kernel", "Kernel type to use: 'linear', 'polynomial', 'cosine'.", + "K", "linear"); PARAM_FLAG("naive", "If true, O(n^2) naive mode is used for computation.", "N"); PARAM_FLAG("single", "If true, single-tree search is used (as opposed to " @@ -44,6 +50,8 @@ PARAM_FLAG("single", "If true, single-tree search is used (as opposed to " int main(int argc, char** argv) { + distanceEvaluations = 0; + CLI::ParseCommandLine(argc, argv); // Get reference dataset filename. @@ -73,7 +81,10 @@ int main(int argc, char** argv) } // Check on kernel type. - if ((kernelType != "linear") && (kernelType != "polynomial")) + if ((kernelType != "linear") && (kernelType != "polynomial") && + (kernelType != "cosine") && (kernelType != "polynomial5") && + (kernelType != "polynomial10") && (kernelType != "polynomial0.5") && + (kernelType != "polynomial0.1") && (kernelType != "gaussian")) { Log::Fatal << "Invalid kernel type: '" << kernelType << "'; must be "; Log::Fatal << "'linear' or 'polynomial'." << endl; @@ -118,6 +129,41 @@ int main(int argc, char** argv) (single && !naive), naive); maxip.Search(k, indices, products); } + else if (kernelType == "cosine") + { + MaxIP maxip(referenceData, (single && !naive), naive); + maxip.Search(k, indices, products); + } + else if (kernelType == "polynomial5") + { + MaxIP > maxip(referenceData, + (single && !naive), naive); + maxip.Search(k, indices, products); + } + else if (kernelType == "polynomial10") + { + MaxIP > maxip(referenceData, + (single && !naive), naive); + maxip.Search(k, indices, products); + } + else if (kernelType == "polynomial0.5") + { + MaxIP > maxip(referenceData, + (single && !naive), naive); + maxip.Search(k, indices, products); + } + else if (kernelType == "polynomial0.1") + { + MaxIP > maxip(referenceData, + (single && !naive), naive); + maxip.Search(k, indices, products); + } + else if (kernelType == "gaussian") + { + MaxIP maxip(referenceData, (single && !naive), + naive); + maxip.Search(k, indices, products); + } } else { @@ -133,6 +179,42 @@ int main(int argc, char** argv) (single && !naive), naive); maxip.Search(k, indices, products); } + else if (kernelType == "cosine") + { + MaxIP maxip(referenceData, queryData, (single && !naive), + naive); + maxip.Search(k, indices, products); + } + else if (kernelType == "polynomial5") + { + MaxIP > maxip(referenceData, queryData, + (single && !naive), naive); + maxip.Search(k, indices, products); + } + else if (kernelType == "polynomial10") + { + MaxIP > maxip(referenceData, queryData, + (single && !naive), naive); + maxip.Search(k, indices, products); + } + else if (kernelType == "polynomial0.5") + { + MaxIP > maxip(referenceData, queryData, + (single && !naive), naive); + maxip.Search(k, indices, products); + } + else if (kernelType == "polynomial0.1") + { + MaxIP > maxip(referenceData, queryData, + (single && !naive), naive); + maxip.Search(k, indices, products); + } + else if (kernelType == "gaussian") + { + MaxIP maxip(referenceData, queryData, + (single && !naive), naive); + maxip.Search(k, indices, products); + } } // Save output, if we were asked to.