use unsafe_col to speed up

This commit is contained in:
HurricaneTong
2015-04-29 14:31:09 -04:00
committed by Ryan Curtin
parent 823b8d755f
commit f381fe8ae7
+5 -28
View File
@@ -93,8 +93,7 @@ class MeanShift
private:
/**
* If the kernel doesn't include a squared distance,
* general way will be applied to calculate the weight of a data point.
* A general approach to calculate the weight for a point.
*
* @param centroid The centroid to calculate the weight
* @param point Calculate its weight
@@ -102,23 +101,7 @@ class MeanShift
* @return If true, the @point is near enough to the @centroid and @weight is valid,
* If false, the @point is far from the @centroid and @weight is invalid.
*/
template <typename Kernel = KernelType>
typename std::enable_if<!kernel::KernelTraits<Kernel>::UsesSquaredDistance, bool>::type
CalcWeight(const arma::colvec& centroid, const arma::colvec& point, double& weight);
/**
* If the kernel includes a squared distance,
* the weight of a data point can be calculated faster.
*
* @param centroid The centroid to calculate the weight
* @param point Calculate its weight
* @param weight Store the weight
* @return If true, the @point is near enough to the @centroid and @weight is valid,
* If false, the @point is far from the @centroid and @weight is invalid.
*/
template <typename Kernel = KernelType>
typename std::enable_if<kernel::KernelTraits<Kernel>::UsesSquaredDistance, bool>::type
CalcWeight(const arma::colvec& centroid, const arma::colvec& point, double& weight);
bool CalcWeight(const arma::colvec& centroid, const arma::colvec& point, double& weight);
/**
* If distance of two centroids is less than radius, one will be removed.
@@ -127,24 +110,18 @@ class MeanShift
*/
double radius;
// By storing radius * radius, we can speed up a little.
double squaredRadius;
//! Maximum number of iterations before giving up.
size_t maxIterations;
//! Instantiated kernel.
KernelType kernel;
metric::EuclideanDistance metric;
};
}; // namespace meanshift
}; // namespace mlpack
} // namespace meanshift
} // namespace mlpack
// Include implementation.
#include "mean_shift_impl.hpp"
#endif // __MLPACK_METHODS_MEAN_SHIFT_MEAN_SHIFT_HPP
#endif // __MLPACK_METHODS_MEAN_SHIFT_MEAN_SHIFT_HPP