Distributed kde test passes.
This commit is contained in:
@@ -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);
|
||||
}
|
||||
|
||||
+4
-3
@@ -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();
|
||||
|
||||
Reference in New Issue
Block a user