More compilation error fix.
This commit is contained in:
@@ -47,6 +47,10 @@ class MatrixFactorizedFMM {
|
||||
/** @brief The root of the reference tree.
|
||||
*/
|
||||
ReferenceTree *reference_tree_root_;
|
||||
|
||||
/** @brief The list of leaf nodes in the reference tree.
|
||||
*/
|
||||
ArrayList<ReferenceTree *> reference_leaf_nodes_;
|
||||
|
||||
/** @brief The permutation mapping indices of reference_set_ to its
|
||||
* original order.
|
||||
@@ -63,6 +67,34 @@ class MatrixFactorizedFMM {
|
||||
const QueryTree *query_node,
|
||||
const ReferenceTree *reference_node,
|
||||
Vector &query_kernel_sums) const;
|
||||
|
||||
/** @brief The canonical case for evaluating the reference
|
||||
* contributions to the given set of query points using the
|
||||
* dual-tree algorithm.
|
||||
*/
|
||||
void CanonicalCase_(const Matrix &query_set,
|
||||
const ArrayList<index_t> &query_index_permutation,
|
||||
const QueryTree *query_node,
|
||||
const ReferenceTree *reference_node,
|
||||
Vector &query_kernel_sums) const;
|
||||
|
||||
/** @brief Traverse the FASTLib tree to get the list of leaf nodes.
|
||||
*/
|
||||
template<typename Tree>
|
||||
void GetLeafNodes_(Tree *node, ArrayList<Tree *> &leaf_nodes);
|
||||
|
||||
/** @brief The method for preprocessing the query tree.
|
||||
*/
|
||||
void PreProcessQueryTree_
|
||||
(const Matrix &query_set, QueryTree *query_node,
|
||||
const Matrix &reference_set,
|
||||
const ArrayList<ReferenceTree *> &reference_leaf_nodes);
|
||||
|
||||
/** @brief The method for preprocessing the reference tree.
|
||||
*/
|
||||
void PreProcessReferenceTree_
|
||||
(ReferenceTree *reference_node, const Matrix &query_set,
|
||||
const ArrayList<QueryTree *> &query_leaf_nodes);
|
||||
|
||||
public:
|
||||
|
||||
@@ -74,6 +106,11 @@ class MatrixFactorizedFMM {
|
||||
*/
|
||||
void Init(const Matrix &references, struct datanode *module_in);
|
||||
|
||||
/** @brief Compute the weighted kernel sums at each point in the
|
||||
* given query set.
|
||||
*/
|
||||
void Compute(const Matrix &queries, Vector *query_kernel_sums);
|
||||
|
||||
};
|
||||
|
||||
#include "matrix_factorized_fmm_impl.h"
|
||||
|
||||
@@ -35,3 +35,211 @@ void MatrixFactorizedFMM<TKernelAux>::BaseCase_
|
||||
|
||||
} // end of looping over each query point.
|
||||
}
|
||||
|
||||
template<typename TKernelAux>
|
||||
void MatrixFactorizedFMM<TKernelAux>::CanonicalCase_
|
||||
(const Matrix &query_set, const ArrayList<index_t> &query_index_permutation,
|
||||
const QueryTree *query_node, const ReferenceTree *reference_node,
|
||||
Vector &query_kernel_sums) const {
|
||||
|
||||
// If the current query/reference node is prunable, then
|
||||
// approximate.
|
||||
|
||||
|
||||
// If the query node is a leaf node,
|
||||
if(query_node->is_leaf()) {
|
||||
|
||||
// ... and the reference node is a leaf node, then we do base
|
||||
// computation.
|
||||
if(reference_node->is_leaf()) {
|
||||
BaseCase_(query_set, query_index_permutation, query_node, reference_node,
|
||||
query_kernel_sums);
|
||||
}
|
||||
|
||||
// ...and the reference node is not a leaf node, then recurse on
|
||||
// the reference side.
|
||||
else {
|
||||
CanonicalCase_(query_set, query_index_permutation, query_node,
|
||||
reference_node->left(), query_kernel_sums);
|
||||
CanonicalCase_(query_set, query_index_permutation, query_node,
|
||||
reference_node->right(), query_kernel_sums);
|
||||
}
|
||||
} // end case for the query node as the leaf node.
|
||||
|
||||
// If the query node is not a leaf node,
|
||||
else {
|
||||
|
||||
// ... and the reference node is a leaf node, then recurse on the
|
||||
// query side.
|
||||
if(reference_node->is_leaf()) {
|
||||
CanonicalCase_(query_set, query_index_permutation, query_node->left(),
|
||||
reference_node, query_kernel_sums);
|
||||
CanonicalCase_(query_set, query_index_permutation, query_node->right(),
|
||||
reference_node, query_kernel_sums);
|
||||
}
|
||||
|
||||
// .. and the reference node is not a leaf node, then do the
|
||||
// four-way recursion.
|
||||
else {
|
||||
CanonicalCase_(query_set, query_index_permutation, query_node->left(),
|
||||
reference_node->left(), query_kernel_sums);
|
||||
CanonicalCase_(query_set, query_index_permutation, query_node->left(),
|
||||
reference_node->right(), query_kernel_sums);
|
||||
CanonicalCase_(query_set, query_index_permutation, query_node->right(),
|
||||
reference_node->left(), query_kernel_sums);
|
||||
CanonicalCase_(query_set, query_index_permutation, query_node->right(),
|
||||
reference_node->right(), query_kernel_sums);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template<typename TKernelAux>
|
||||
template<typename Tree>
|
||||
void MatrixFactorizedFMM<TKernelAux>::GetLeafNodes_
|
||||
(Tree *node, ArrayList<Tree *> &leaf_nodes) {
|
||||
|
||||
if(node->is_leaf()) {
|
||||
leaf_nodes.PushBackCopy(node);
|
||||
}
|
||||
else {
|
||||
GetLeafNodes_(node->left(), leaf_nodes);
|
||||
GetLeafNodes_(node->right(), leaf_nodes);
|
||||
}
|
||||
}
|
||||
|
||||
template<typename TKernelAux>
|
||||
void MatrixFactorizedFMM<TKernelAux>::PreProcessQueryTree_
|
||||
(const Matrix &query_set, QueryTree *query_node, const Matrix &reference_set,
|
||||
const ArrayList<ReferenceTree *> &reference_leaf_nodes) {
|
||||
|
||||
// Initialize the local expansion object.
|
||||
MatrixFactorizedLocalExpansion<TKernelAux> &local_expansion =
|
||||
(query_node->stat()).local_expansion_;
|
||||
local_expansion.Init(ka_);
|
||||
|
||||
// For query leaf nodes, train the incoming representation using the
|
||||
// set of reference leaf nodes using stratified sampling.
|
||||
if(query_node->is_leaf()) {
|
||||
local_expansion.TrainBasisFunctions(query_set, query_node->begin(),
|
||||
query_node->begin() +
|
||||
query_node->count(), &reference_set,
|
||||
&reference_leaf_nodes);
|
||||
}
|
||||
|
||||
// For an internal query node, merge the incoming representations of
|
||||
// its children.
|
||||
else {
|
||||
PreProcessQueryTree_(query_set, query_node->left(), reference_set,
|
||||
reference_leaf_nodes);
|
||||
PreProcessQueryTree_(query_set, query_node->right(), reference_set,
|
||||
reference_leaf_nodes);
|
||||
|
||||
local_expansion.CombineBasisFunctions
|
||||
((query_node->left()->stat()).local_expansion_,
|
||||
(query_node->right()->stat()).local_expansion_);
|
||||
}
|
||||
}
|
||||
|
||||
template<typename TKernelAux>
|
||||
void MatrixFactorizedFMM<TKernelAux>::PreProcessReferenceTree_
|
||||
(ReferenceTree *reference_node, const Matrix &query_set,
|
||||
const ArrayList<QueryTree *> &query_leaf_nodes) {
|
||||
|
||||
// Initialize the far-field expansion object.
|
||||
MatrixFactorizedFarFieldExpansion<TKernelAux> &farfield_expansion =
|
||||
(reference_node->stat()).farfield_expansion_;
|
||||
farfield_expansion.Init(ka_);
|
||||
|
||||
// For reference leaf nodes, train the outgoing representation using
|
||||
// the set of query leaf nodes using stratified sampling.
|
||||
if(reference_node->is_leaf()) {
|
||||
farfield_expansion.AccumulateCoeffs(reference_set_, reference_weights_,
|
||||
reference_node->begin(),
|
||||
reference_node->begin() +
|
||||
reference_node->count(),
|
||||
-1, &query_set, &query_leaf_nodes);
|
||||
}
|
||||
|
||||
// For an internal reference node, merge the representations of its
|
||||
// children.
|
||||
else {
|
||||
PreProcessReferenceTree_(reference_node->left(), query_set,
|
||||
query_leaf_nodes);
|
||||
PreProcessReferenceTree_(reference_node->right(), query_set,
|
||||
query_leaf_nodes);
|
||||
|
||||
farfield_expansion.CombineBasisFunctions
|
||||
((reference_node->left()->stat()).farfield_expansion_,
|
||||
(reference_node->right()->stat()).farfield_expansion_);
|
||||
}
|
||||
}
|
||||
|
||||
template<typename TKernelAux>
|
||||
void MatrixFactorizedFMM<TKernelAux>::Init(const Matrix &references,
|
||||
struct datanode *module_in) {
|
||||
|
||||
// Point to the incoming module.
|
||||
module_ = module_in;
|
||||
|
||||
// Read in the number of points owned by a leaf
|
||||
int leaflen = fx_param_int(module_in, "leaflen", 20);
|
||||
|
||||
// Copy reference dataset and reference weights. Currently supports
|
||||
// only the uniform weight.
|
||||
reference_set_.Copy(references);
|
||||
reference_weights_.Init(reference_set_.n_cols());
|
||||
reference_weights_.SetAll(1);
|
||||
|
||||
// Construct the reference tree.
|
||||
fx_timer_start(fx_root, "reference_tree_construction");
|
||||
reference_tree_root_ = tree::MakeKdTreeMidpoint<ReferenceTree>
|
||||
(reference_set_, leaflen, &old_from_new_references_, NULL);
|
||||
fx_timer_stop(fx_root, "reference_tree_construction");
|
||||
|
||||
// Retrieve the list of reference leaf nodes.
|
||||
reference_leaf_nodes_.Init();
|
||||
GetLeafNodes_(reference_tree_root_, reference_leaf_nodes_);
|
||||
|
||||
// Retrieve the bandwidth and initialize the kernel.
|
||||
double bandwidth = fx_param_double_req(module_, "bandwidth");
|
||||
ka_.Init(bandwidth, 0, references.n_rows());
|
||||
}
|
||||
|
||||
template<typename TKernelAux>
|
||||
void MatrixFactorizedFMM<TKernelAux>::Compute
|
||||
(const Matrix &queries, Vector *query_kernel_sums) {
|
||||
|
||||
// Construct the query tree.
|
||||
int leaflen = fx_param_int(module_, "leaflen", 20);
|
||||
|
||||
// Copy the query dataset.
|
||||
Matrix query_set;
|
||||
query_set.Copy(queries);
|
||||
|
||||
fx_timer_start(fx_root, "query_tree_construction");
|
||||
ArrayList<index_t> old_from_new_queries;
|
||||
QueryTree *query_tree_root =
|
||||
tree::MakeKdTreeMidpoint<QueryTree>
|
||||
(query_set, leaflen, &old_from_new_queries, NULL);
|
||||
fx_timer_stop(fx_root, "query_tree_construction");
|
||||
|
||||
// Retrieve the leaf node lists in the query tree.
|
||||
ArrayList<QueryTree *> query_leaf_nodes;
|
||||
query_leaf_nodes.Init();
|
||||
GetLeafNodes_(query_tree_root, query_leaf_nodes);
|
||||
|
||||
// Train the basis functions in the reference tree and the query
|
||||
// tree.
|
||||
PreProcessReferenceTree_(reference_tree_root_, query_set, query_leaf_nodes);
|
||||
PreProcessQueryTree_(query_set, query_tree_root, reference_set_,
|
||||
reference_leaf_nodes_);
|
||||
|
||||
// Compute the kernel summations.
|
||||
query_kernel_sums->Init(query_set.n_cols());
|
||||
query_kernel_sums->SetZero();
|
||||
CanonicalCase_(query_set, old_from_new_queries, query_tree_root,
|
||||
reference_tree_root_, *query_kernel_sums);
|
||||
|
||||
// Delete the query tree after the computation...
|
||||
delete query_tree_root;
|
||||
}
|
||||
|
||||
@@ -1,7 +1,72 @@
|
||||
#include "matrix_factorized_fmm.h"
|
||||
#include "fastlib/fastlib.h"
|
||||
#include "mlpack/kde/dataset_scaler.h"
|
||||
#include "mlpack/series_expansion/matrix_factorized_kernel_aux.h"
|
||||
|
||||
int main(int argc, char *argv) {
|
||||
int main(int argc, char *argv[]) {
|
||||
|
||||
// Initialize FastExec (parameter handling stuff)
|
||||
fx_init(argc, argv);
|
||||
|
||||
////////// READING PARAMETERS AND LOADING DATA /////////////////////
|
||||
|
||||
// FASTexec organizes parameters and results into submodules. Think
|
||||
// of this as creating a new folder named "kde_module" under the
|
||||
// root directory (NULL) for the Kde object to work inside. Here,
|
||||
// we initialize it with all parameters defined "--kde/...=...".
|
||||
struct datanode* kde_module =
|
||||
fx_submodule(NULL, "kde", "kde_module");
|
||||
|
||||
// The reference data file is a required parameter.
|
||||
const char* references_file_name = fx_param_str_req(fx_root, "data");
|
||||
|
||||
// The query data file defaults to the references.
|
||||
const char* queries_file_name =
|
||||
fx_param_str(fx_root, "query", references_file_name);
|
||||
|
||||
// flag for determining whether to compute naively
|
||||
bool do_naive = fx_param_exists(kde_module, "do_naive");
|
||||
|
||||
// The query and reference datasets
|
||||
Matrix references;
|
||||
Matrix queries;
|
||||
|
||||
// The flag for telling whether references are equal to queries
|
||||
bool queries_equal_references =
|
||||
!strcmp(queries_file_name, references_file_name);
|
||||
|
||||
// data::Load inits a matrix with the contents of a .csv or .arff.
|
||||
data::Load(references_file_name, &references);
|
||||
if(queries_equal_references) {
|
||||
queries.Alias(references);
|
||||
}
|
||||
else {
|
||||
data::Load(queries_file_name, &queries);
|
||||
}
|
||||
|
||||
// Confirm whether the user asked for scaling of the dataset
|
||||
if(!strcmp(fx_param_str(kde_module, "scaling", "none"), "range")) {
|
||||
DatasetScaler::ScaleDataByMinMax(queries, references,
|
||||
queries_equal_references);
|
||||
}
|
||||
|
||||
if(!strcmp(fx_param_str(kde_module, "kernel", "gaussian"), "gaussian")) {
|
||||
|
||||
Vector fast_kde_results;
|
||||
|
||||
printf("Kernel independent expansion for Gaussian kernel KDE\n");
|
||||
MatrixFactorizedFMM<GaussianKernelMatrixFactorizedAux> fast_kde;
|
||||
fast_kde.Init(references, kde_module);
|
||||
fast_kde.Compute(queries, &fast_kde_results);
|
||||
}
|
||||
else if(!strcmp(fx_param_str(kde_module, "kernel", "epan"), "epan")) {
|
||||
MatrixFactorizedFMM<EpanKernelMatrixFactorizedAux> fast_kde;
|
||||
Vector fast_kde_results;
|
||||
|
||||
fast_kde.Init(references, kde_module);
|
||||
fast_kde.Compute(queries, &fast_kde_results);
|
||||
}
|
||||
|
||||
fx_done();
|
||||
return 0;
|
||||
}
|
||||
|
||||
@@ -47,7 +47,7 @@ class MatrixFactorizedFMMQueryNodeStat {
|
||||
|
||||
/** @brief The local expansion for the query points in this node.
|
||||
*/
|
||||
typename TKernelAux::TFarFieldExpansion local_expansion_;
|
||||
typename TKernelAux::TLocalExpansion local_expansion_;
|
||||
|
||||
void Init(const TKernelAux &ka) {
|
||||
local_expansion_.Init(ka);
|
||||
|
||||
Reference in New Issue
Block a user