From 174fc0514d92a5baebd255f699c04dbbd0eedd1e Mon Sep 17 00:00:00 2001 From: Ryan Curtin Date: Fri, 18 May 2012 03:30:20 +0000 Subject: [PATCH] unsafe_col() is faster. --- src/mlpack/methods/maxip/max_ip_impl.hpp | 4 ++-- src/mlpack/methods/maxip/max_ip_rules_impl.hpp | 9 +++++---- 2 files changed, 7 insertions(+), 6 deletions(-) diff --git a/src/mlpack/methods/maxip/max_ip_impl.hpp b/src/mlpack/methods/maxip/max_ip_impl.hpp index bd4a287555..3191f2844d 100644 --- a/src/mlpack/methods/maxip/max_ip_impl.hpp +++ b/src/mlpack/methods/maxip/max_ip_impl.hpp @@ -91,8 +91,8 @@ void MaxIP::Search(const size_t k, { for (size_t r = 0; r < referenceSet.n_cols; ++r) { - const double eval = KernelType::Evaluate(querySet.col(q), - referenceSet.col(r)); + const double eval = KernelType::Evaluate(querySet.unsafe_col(q), + referenceSet.unsafe_col(r)); size_t insertPosition; for (insertPosition = 0; insertPosition < indices.n_rows; diff --git a/src/mlpack/methods/maxip/max_ip_rules_impl.hpp b/src/mlpack/methods/maxip/max_ip_rules_impl.hpp index 369cd3dfec..4832282e6b 100644 --- a/src/mlpack/methods/maxip/max_ip_rules_impl.hpp +++ b/src/mlpack/methods/maxip/max_ip_rules_impl.hpp @@ -40,8 +40,9 @@ bool MaxIPRules::CanPrune(const size_t queryIndex, // and since we are using cover trees, p_0 is the point referred to by the // node, and R_p will be the expansion constant to the power of the scale plus // one. - const double eval = MetricType::Kernel::Evaluate(querySet.col(queryIndex), - referenceSet.col(referenceNode.Point())); + const double eval = MetricType::Kernel::Evaluate( + querySet.unsafe_col(queryIndex), + referenceSet.unsafe_col(referenceNode.Point())); // See if base case can be added. if (eval > products(products.n_rows - 1, queryIndex)) @@ -57,8 +58,8 @@ bool MaxIPRules::CanPrune(const size_t queryIndex, double maxProduct = eval + std::pow(referenceNode.ExpansionConstant(), referenceNode.Scale() + 1) * - sqrt(MetricType::Kernel::Evaluate(querySet.col(queryIndex), - querySet.col(queryIndex))); + sqrt(MetricType::Kernel::Evaluate(querySet.unsafe_col(queryIndex), + querySet.unsafe_col(queryIndex))); if (maxProduct > products(products.n_rows - 1, queryIndex)) return false;