Distributed kde test passes.

This commit is contained in:
Dongryeol Lee
2010-12-11 04:09:44 +00:00
parent ecef8980fe
commit 14aa6d1de3
5 changed files with 31 additions and 25 deletions
@@ -30,8 +30,6 @@ void core::gnp::DistributedDualtreeDfs<DistributedProblemType>::AllReduce_(
self_argument.Init(problem_->global());
self_problem.Init(self_argument);
self_engine.Init(self_problem);
self_problem.global().set_effective_num_reference_points(
problem_->global().effective_num_reference_points());
self_engine.Compute(metric, query_results);
world_->barrier();
@@ -103,8 +101,6 @@ void core::gnp::DistributedDualtreeDfs<DistributedProblemType>::AllReduce_(
query_table_->local_table(),
problem_->global());
sub_problem.Init(sub_argument);
sub_problem.global().set_effective_num_reference_points(
problem_->global().effective_num_reference_points());
sub_engine.Init(sub_problem);
sub_engine.Compute(metric, query_results, false);
}
@@ -86,7 +86,8 @@ void DistributedKde<DistributedTableType>::Init(
// Declare the global constants.
global_.Init(
reference_table_, query_table_, arguments_in.bandwidth_, is_monochromatic_,
reference_table_, query_table_, reference_table_->n_entries(),
arguments_in.bandwidth_, is_monochromatic_,
arguments_in.relative_error_, arguments_in.probability_,
arguments_in.kernel_, false);
global_.set_effective_num_reference_points(
@@ -119,11 +120,11 @@ bool DistributedKde<DistributedTableType>::ConstructBoostVariableMap_(
"the leave-one-out density at each reference point."
)(
"random_generate_n_attributes",
boost::program_options::value<int>()->default_value(3),
boost::program_options::value<int>(),
"Generate the datasets on the fly of the specified dimension."
)(
"random_generate_n_entries",
boost::program_options::value<int>()->default_value(20),
boost::program_options::value<int>(),
"Generate the datasets on the fly of the specified number of points."
)(
"densities_out",
@@ -23,6 +23,8 @@ class KdeArguments {
TableType *query_table_;
double effective_num_reference_points_;
double bandwidth_;
double relative_error_;
@@ -35,6 +37,8 @@ class KdeArguments {
bool tables_are_aliased_;
bool normalize_densities_;
public:
template<typename GlobalType>
@@ -43,11 +47,14 @@ class KdeArguments {
GlobalType &global_in) {
reference_table_ = reference_table_in;
query_table_ = query_table_in;
effective_num_reference_points_ =
global_in.effective_num_reference_points();
bandwidth_ = global_in.bandwidth();
relative_error_ = global_in.relative_error();
probability_ = global_in.probability();
kernel_ = global_in.kernel().name();
tables_are_aliased_ = true;
normalize_densities_ = global_in.normalize_densities();
}
template<typename GlobalType>
@@ -56,23 +63,28 @@ class KdeArguments {
if(reference_table_ != query_table_) {
query_table_ = global_in.query_table()->local_table();
}
effective_num_reference_points_ =
global_in.effective_num_reference_points();
bandwidth_ = global_in.bandwidth();
relative_error_ = global_in.relative_error();
probability_ = global_in.probability();
kernel_ = global_in.kernel().name();
tables_are_aliased_ = true;
normalize_densities_ = global_in.normalize_densities();
}
KdeArguments() {
leaf_size_ = 0;
reference_table_ = NULL;
query_table_ = NULL;
effective_num_reference_points_ = 0.0;
bandwidth_ = 0.0;
relative_error_ = 0.0;
probability_ = 0.0;
kernel_ = "";
metric_ = NULL;
tables_are_aliased_ = false;
normalize_densities_ = true;
}
~KdeArguments() {
@@ -60,9 +60,11 @@ void mlpack::kde::Kde<TableType>::Init(
// Declare the global constants.
global_.Init(
reference_table_, query_table_, arguments_in.bandwidth_, is_monochromatic_,
reference_table_, query_table_,
arguments_in.effective_num_reference_points_,
arguments_in.bandwidth_, is_monochromatic_,
arguments_in.relative_error_, arguments_in.probability_,
arguments_in.kernel_);
arguments_in.kernel_, arguments_in.normalize_densities_);
}
template<typename TableType>
@@ -221,9 +223,13 @@ void mlpack::kde::Kde<TableType>::ParseArguments(
arguments_out->query_table_->IndexData(
*(arguments_out->metric_), arguments_out->leaf_size_);
std::cout << "Finished building the query tree.\n";
arguments_out->effective_num_reference_points_ =
arguments_out->reference_table_->n_entries();
}
else {
arguments_out->query_table_ = arguments_out->reference_table_;
arguments_out->effective_num_reference_points_ =
arguments_out->reference_table_->n_entries() - 1;
}
// Parse the bandwidth.
@@ -144,16 +144,6 @@ class KdeGlobal {
return effective_num_reference_points_;
}
void set_effective_num_reference_points(
double effective_num_reference_points_in) {
effective_num_reference_points_ = effective_num_reference_points_in;
mult_const_ = 1.0 /
(kernel_->CalcNormConstant(
reference_table_->n_attributes()) *
((double) effective_num_reference_points_));
}
template<typename DistributedTableType>
void set_effective_num_reference_points(
boost::mpi::communicator &comm,
@@ -164,10 +154,13 @@ class KdeGlobal {
for(int i = 0; i < comm.size(); i++) {
total_sum += reference_table_in->local_n_entries(i);
}
double effective_num_reference_points_in =
effective_num_reference_points_ =
(reference_table_in == query_table_in) ?
(total_sum - 1.0) : total_sum;
set_effective_num_reference_points(effective_num_reference_points_in);
mult_const_ = 1.0 /
(kernel_->CalcNormConstant(
reference_table_in->n_attributes()) *
((double) effective_num_reference_points_));
}
~KdeGlobal() {
@@ -228,15 +221,13 @@ class KdeGlobal {
void Init(
TableType *reference_table_in,
TableType *query_table_in,
int effective_num_reference_points_in,
double bandwidth_in, const bool is_monochromatic,
double relative_error_in, double probability_in,
const std::string &kernel_type_in,
bool normalize_densities_in = true) {
effective_num_reference_points_ =
(is_monochromatic) ?
(reference_table_in->n_entries() - 1) :
reference_table_in->n_entries();
effective_num_reference_points_ = effective_num_reference_points_in;
if(kernel_type_in == "gaussian") {
kernel_ = new core::metric_kernels::GaussianKernel();