More compilation error fix.

This commit is contained in:
Dongryeol Lee
2008-05-20 00:51:36 +00:00
parent 1d2c20a944
commit 6015ec3a94
4 changed files with 312 additions and 2 deletions
@@ -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);