Refactor FastMKS to allow sparse datasets.

This commit is contained in:
ryan
2015-04-05 18:48:09 -04:00
parent a62125b7f7
commit a8637c19dd
5 changed files with 50 additions and 48 deletions
+11 -11
View File
@@ -63,7 +63,7 @@ class FastMKS
* @param single Whether or not to run single-tree search.
* @param naive Whether or not to run brute-force (naive) search.
*/
FastMKS(const arma::mat& referenceSet,
FastMKS(const typename TreeType::Mat& referenceSet,
const bool single = false,
const bool naive = false);
@@ -77,8 +77,8 @@ class FastMKS
* @param single Whether or not to run single-tree search.
* @param naive Whether or not to run brute-force (naive) search.
*/
FastMKS(const arma::mat& referenceSet,
const arma::mat& querySet,
FastMKS(const typename TreeType::Mat& referenceSet,
const typename TreeType::Mat& querySet,
const bool single = false,
const bool naive = false);
@@ -93,7 +93,7 @@ class FastMKS
* @param single Whether or not to run single-tree search.
* @param naive Whether or not to run brute-force (naive) search.
*/
FastMKS(const arma::mat& referenceSet,
FastMKS(const typename TreeType::Mat& referenceSet,
KernelType& kernel,
const bool single = false,
const bool naive = false);
@@ -110,8 +110,8 @@ class FastMKS
* @param single Whether or not to run single-tree search.
* @param naive Whether or not to run brute-force (naive) search.
*/
FastMKS(const arma::mat& referenceSet,
const arma::mat& querySet,
FastMKS(const typename TreeType::Mat& referenceSet,
const typename TreeType::Mat& querySet,
KernelType& kernel,
const bool single = false,
const bool naive = false);
@@ -128,7 +128,7 @@ class FastMKS
* @param single Whether or not to run single-tree search.
* @param naive Whether or not to run brute-force (naive) search.
*/
FastMKS(const arma::mat& referenceSet,
FastMKS(const typename TreeType::Mat& referenceSet,
TreeType* referenceTree,
const bool single = false,
const bool naive = false);
@@ -146,9 +146,9 @@ class FastMKS
* @param single Whether or not to use single-tree search.
* @param naive Whether or not to use naive (brute-force) search.
*/
FastMKS(const arma::mat& referenceSet,
FastMKS(const typename TreeType::Mat& referenceSet,
TreeType* referenceTree,
const arma::mat& querySet,
const typename TreeType::Mat& querySet,
TreeType* queryTree,
const bool single = false,
const bool naive = false);
@@ -186,9 +186,9 @@ class FastMKS
private:
//! The reference dataset.
const arma::mat& referenceSet;
const typename TreeType::Mat& referenceSet;
//! The query dataset.
const arma::mat& querySet;
const typename TreeType::Mat& querySet;
//! The tree built on the reference dataset.
TreeType* referenceTree;
+20 -18
View File
@@ -20,7 +20,7 @@ namespace fastmks {
// Single dataset, no instantiated kernel.
template<typename KernelType, typename TreeType>
FastMKS<KernelType, TreeType>::FastMKS(const arma::mat& referenceSet,
FastMKS<KernelType, TreeType>::FastMKS(const typename TreeType::Mat& referenceSet,
const bool single,
const bool naive) :
referenceSet(referenceSet),
@@ -44,8 +44,8 @@ FastMKS<KernelType, TreeType>::FastMKS(const arma::mat& referenceSet,
// Two datasets, no instantiated kernel.
template<typename KernelType, typename TreeType>
FastMKS<KernelType, TreeType>::FastMKS(const arma::mat& referenceSet,
const arma::mat& querySet,
FastMKS<KernelType, TreeType>::FastMKS(const typename TreeType::Mat& referenceSet,
const typename TreeType::Mat& querySet,
const bool single,
const bool naive) :
referenceSet(referenceSet),
@@ -70,7 +70,7 @@ FastMKS<KernelType, TreeType>::FastMKS(const arma::mat& referenceSet,
// One dataset, instantiated kernel.
template<typename KernelType, typename TreeType>
FastMKS<KernelType, TreeType>::FastMKS(const arma::mat& referenceSet,
FastMKS<KernelType, TreeType>::FastMKS(const typename TreeType::Mat& referenceSet,
KernelType& kernel,
const bool single,
const bool naive) :
@@ -97,8 +97,8 @@ FastMKS<KernelType, TreeType>::FastMKS(const arma::mat& referenceSet,
// Two datasets, instantiated kernel.
template<typename KernelType, typename TreeType>
FastMKS<KernelType, TreeType>::FastMKS(const arma::mat& referenceSet,
const arma::mat& querySet,
FastMKS<KernelType, TreeType>::FastMKS(const typename TreeType::Mat& referenceSet,
const typename TreeType::Mat& querySet,
KernelType& kernel,
const bool single,
const bool naive) :
@@ -125,10 +125,11 @@ FastMKS<KernelType, TreeType>::FastMKS(const arma::mat& referenceSet,
// One dataset, pre-built tree.
template<typename KernelType, typename TreeType>
FastMKS<KernelType, TreeType>::FastMKS(const arma::mat& referenceSet,
TreeType* referenceTree,
const bool single,
const bool naive) :
FastMKS<KernelType, TreeType>::FastMKS(
const typename TreeType::Mat& referenceSet,
TreeType* referenceTree,
const bool single,
const bool naive) :
referenceSet(referenceSet),
querySet(referenceSet),
referenceTree(referenceTree),
@@ -145,12 +146,13 @@ FastMKS<KernelType, TreeType>::FastMKS(const arma::mat& referenceSet,
// Two datasets, pre-built trees.
template<typename KernelType, typename TreeType>
FastMKS<KernelType, TreeType>::FastMKS(const arma::mat& referenceSet,
TreeType* referenceTree,
const arma::mat& querySet,
TreeType* queryTree,
const bool single,
const bool naive) :
FastMKS<KernelType, TreeType>::FastMKS(
const typename TreeType::Mat& referenceSet,
TreeType* referenceTree,
const typename TreeType::Mat& querySet,
TreeType* queryTree,
const bool single,
const bool naive) :
referenceSet(referenceSet),
querySet(querySet),
referenceTree(referenceTree),
@@ -205,8 +207,8 @@ void FastMKS<KernelType, TreeType>::Search(const size_t k,
if ((&querySet == &referenceSet) && (q == r))
continue;
const double eval = metric.Kernel().Evaluate(querySet.unsafe_col(q),
referenceSet.unsafe_col(r));
const double eval = metric.Kernel().Evaluate(querySet.col(q),
referenceSet.col(r));
size_t insertPosition;
for (insertPosition = 0; insertPosition < indices.n_rows;
+4 -4
View File
@@ -22,8 +22,8 @@ template<typename KernelType, typename TreeType>
class FastMKSRules
{
public:
FastMKSRules(const arma::mat& referenceSet,
const arma::mat& querySet,
FastMKSRules(const typename TreeType::Mat& referenceSet,
const typename TreeType::Mat& querySet,
arma::Mat<size_t>& indices,
arma::mat& products,
KernelType& kernel);
@@ -98,9 +98,9 @@ class FastMKSRules
private:
//! The reference dataset.
const arma::mat& referenceSet;
const typename TreeType::Mat& referenceSet;
//! The query dataset.
const arma::mat& querySet;
const typename TreeType::Mat& querySet;
//! The indices of the maximum kernel results.
arma::Mat<size_t>& indices;
@@ -14,11 +14,12 @@ namespace mlpack {
namespace fastmks {
template<typename KernelType, typename TreeType>
FastMKSRules<KernelType, TreeType>::FastMKSRules(const arma::mat& referenceSet,
const arma::mat& querySet,
arma::Mat<size_t>& indices,
arma::mat& products,
KernelType& kernel) :
FastMKSRules<KernelType, TreeType>::FastMKSRules(
const typename TreeType::Mat& referenceSet,
const typename TreeType::Mat& querySet,
arma::Mat<size_t>& indices,
arma::mat& products,
KernelType& kernel) :
referenceSet(referenceSet),
querySet(querySet),
indices(indices),
@@ -33,13 +34,13 @@ FastMKSRules<KernelType, TreeType>::FastMKSRules(const arma::mat& referenceSet,
// Precompute each self-kernel.
queryKernels.set_size(querySet.n_cols);
for (size_t i = 0; i < querySet.n_cols; ++i)
queryKernels[i] = sqrt(kernel.Evaluate(querySet.unsafe_col(i),
querySet.unsafe_col(i)));
queryKernels[i] = sqrt(kernel.Evaluate(querySet.col(i),
querySet.col(i)));
referenceKernels.set_size(referenceSet.n_cols);
for (size_t i = 0; i < referenceSet.n_cols; ++i)
referenceKernels[i] = sqrt(kernel.Evaluate(referenceSet.unsafe_col(i),
referenceSet.unsafe_col(i)));
referenceKernels[i] = sqrt(kernel.Evaluate(referenceSet.col(i),
referenceSet.col(i)));
// Set to invalid memory, so that the first node combination does not try to
// dereference null pointers.
@@ -69,8 +70,8 @@ double FastMKSRules<KernelType, TreeType>::BaseCase(
}
++baseCases;
double kernelEval = kernel.Evaluate(querySet.unsafe_col(queryIndex),
referenceSet.unsafe_col(referenceIndex));
double kernelEval = kernel.Evaluate(querySet.col(queryIndex),
referenceSet.col(referenceIndex));
// Update the last kernel value, if we need to.
if (tree::TreeTraits<TreeType>::FirstPointIsCentroid)
@@ -156,11 +157,10 @@ double FastMKSRules<KernelType, TreeType>::Score(const size_t queryIndex,
}
else
{
const arma::vec queryPoint = querySet.unsafe_col(queryIndex);
arma::vec refCentroid;
referenceNode.Centroid(refCentroid);
kernelEval = kernel.Evaluate(queryPoint, refCentroid);
kernelEval = kernel.Evaluate(querySet.col(queryIndex), refCentroid);
}
referenceNode.Stat().LastKernel() = kernelEval;
+2 -2
View File
@@ -58,8 +58,8 @@ class FastMKSStat
else
{
selfKernel = sqrt(node.Metric().Kernel().Evaluate(
node.Dataset().unsafe_col(node.Point(0)),
node.Dataset().unsafe_col(node.Point(0))));
node.Dataset().col(node.Point(0)),
node.Dataset().col(node.Point(0))));
}
}
else