Another checkpoint, getting there.
This commit is contained in:
+35
-35
@@ -75,9 +75,7 @@ DualtreeDfs<ProblemType>::iterator::IteratorArgType::IteratorArgType(
|
||||
qnode_ = qnode_in;
|
||||
rnode_ = rnode_in;
|
||||
squared_distance_range_ =
|
||||
(query_table_in->get_node_bound(qnode_in)).RangeDistanceSq(
|
||||
metric_in,
|
||||
reference_table_in->get_node_bound(rnode_in));
|
||||
(qnode_in->bound()).RangeDistanceSq(metric_in, rnode_in->bound());
|
||||
}
|
||||
|
||||
template<typename ProblemType>
|
||||
@@ -157,24 +155,24 @@ void DualtreeDfs<ProblemType>::iterator::operator++() {
|
||||
core::math::Range squared_distance_range_first,
|
||||
squared_distance_range_second;
|
||||
TreeType *rnode_second;
|
||||
engine_->Heuristic_(metric_, qnode, query_table_,
|
||||
reference_table_->get_node_left_child(rnode),
|
||||
reference_table_->get_node_right_child(rnode),
|
||||
reference_table_,
|
||||
&rnode_first, squared_distance_range_first,
|
||||
&rnode_second, squared_distance_range_second);
|
||||
engine_->Heuristic_(
|
||||
metric_, qnode, query_table_, rnode->left(), rnode->right(),
|
||||
reference_table_, &rnode_first, squared_distance_range_first,
|
||||
&rnode_second, squared_distance_range_second);
|
||||
|
||||
// Push the first prioritized reference node on the back
|
||||
// of the trace and the later one on the front of the
|
||||
// trace.
|
||||
trace_.push_back(IteratorArgType(
|
||||
metric_, query_table_, qnode,
|
||||
reference_table_, rnode_first,
|
||||
squared_distance_range_first));
|
||||
trace_.push_front(IteratorArgType(
|
||||
metric_, query_table_, qnode,
|
||||
reference_table_, rnode_second,
|
||||
squared_distance_range_second));
|
||||
trace_.push_back(
|
||||
IteratorArgType(
|
||||
metric_, query_table_, qnode,
|
||||
reference_table_, rnode_first,
|
||||
squared_distance_range_first));
|
||||
trace_.push_front(
|
||||
IteratorArgType(
|
||||
metric_, query_table_, qnode,
|
||||
reference_table_, rnode_second,
|
||||
squared_distance_range_second));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -182,11 +180,11 @@ void DualtreeDfs<ProblemType>::iterator::operator++() {
|
||||
else {
|
||||
|
||||
// Here we split the query.
|
||||
TreeType *qnode_left = query_table_->get_node_left_child(qnode);
|
||||
TreeType *qnode_right = query_table_->get_node_right_child(qnode);
|
||||
TreeType *qnode_left = qnode->left();
|
||||
TreeType *qnode_right = qnode->right();
|
||||
|
||||
// If the reference node is leaf node,
|
||||
if(reference_table_->node_is_leaf(rnode)) {
|
||||
if(rnode->is_leaf()) {
|
||||
|
||||
// Push both combinations on the back of the trace.
|
||||
trace_.push_back(IteratorArgType(
|
||||
@@ -201,26 +199,28 @@ void DualtreeDfs<ProblemType>::iterator::operator++() {
|
||||
else {
|
||||
|
||||
// Split the reference.
|
||||
TreeType *rnode_left = reference_table_->get_node_left_child(rnode);
|
||||
TreeType *rnode_right =
|
||||
reference_table_->get_node_right_child(rnode);
|
||||
TreeType *rnode_left = rnode->left();
|
||||
TreeType *rnode_right = rnode->right();
|
||||
|
||||
// Prioritize on the left child of the query node.
|
||||
TreeType *rnode_first = NULL, *rnode_second = NULL;
|
||||
core::math::Range squared_distance_range_first;
|
||||
core::math::Range squared_distance_range_second;
|
||||
engine_->Heuristic_(metric_, qnode_left, query_table_, rnode_left,
|
||||
rnode_right, reference_table_,
|
||||
&rnode_first, squared_distance_range_first,
|
||||
&rnode_second, squared_distance_range_second);
|
||||
trace_.push_back(IteratorArgType(
|
||||
metric_, query_table_, qnode_left,
|
||||
reference_table_, rnode_first,
|
||||
squared_distance_range_first));
|
||||
trace_.push_front(IteratorArgType(
|
||||
metric_, query_table_, qnode_left,
|
||||
reference_table_, rnode_second,
|
||||
squared_distance_range_second));
|
||||
engine_->Heuristic_(
|
||||
metric_, qnode_left, query_table_, rnode_left,
|
||||
rnode_right, reference_table_,
|
||||
&rnode_first, squared_distance_range_first,
|
||||
&rnode_second, squared_distance_range_second);
|
||||
trace_.push_back(
|
||||
IteratorArgType(
|
||||
metric_, query_table_, qnode_left,
|
||||
reference_table_, rnode_first,
|
||||
squared_distance_range_first));
|
||||
trace_.push_front(
|
||||
IteratorArgType(
|
||||
metric_, query_table_, qnode_left,
|
||||
reference_table_, rnode_second,
|
||||
squared_distance_range_second));
|
||||
|
||||
// Prioritize on the right child of the query node.
|
||||
engine_->Heuristic_(metric_, qnode_right, query_table_,
|
||||
|
||||
@@ -19,7 +19,9 @@ int main(int argc, char *argv[]) {
|
||||
|
||||
// Parse arguments for Kde.
|
||||
mlpack::kde::KdeArguments<TableType> kde_arguments;
|
||||
mlpack::kde::Kde<TableType>::ParseArguments(argc, argv, &kde_arguments);
|
||||
if(mlpack::kde::Kde<TableType>::ParseArguments(argc, argv, &kde_arguments)) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
// Instantiate a KDE object.
|
||||
mlpack::kde::Kde<TableType> kde_instance;
|
||||
|
||||
@@ -68,11 +68,11 @@ class Kde {
|
||||
const mlpack::kde::KdeArguments<TableType> &arguments_in,
|
||||
ResultType *result_out);
|
||||
|
||||
static void ParseArguments(
|
||||
static bool ParseArguments(
|
||||
const std::vector<std::string> &args,
|
||||
mlpack::kde::KdeArguments<TableType> *arguments_out);
|
||||
|
||||
static void ParseArguments(
|
||||
static bool ParseArguments(
|
||||
int argc,
|
||||
char *argv[],
|
||||
mlpack::kde::KdeArguments<TableType> *arguments_out);
|
||||
|
||||
@@ -39,6 +39,8 @@ class KdeArguments {
|
||||
|
||||
bool normalize_densities_;
|
||||
|
||||
int num_iterations_in_;
|
||||
|
||||
public:
|
||||
|
||||
template<typename GlobalType>
|
||||
@@ -85,6 +87,7 @@ class KdeArguments {
|
||||
metric_ = NULL;
|
||||
tables_are_aliased_ = false;
|
||||
normalize_densities_ = true;
|
||||
num_iterations_in_ = 0;
|
||||
}
|
||||
|
||||
~KdeArguments() {
|
||||
|
||||
@@ -36,12 +36,19 @@ void mlpack::kde::Kde<TableType>::Compute(
|
||||
mlpack::kde::KdeResult< std::vector<double> > *result_out) {
|
||||
|
||||
// Instantiate a dual-tree algorithm of the KDE.
|
||||
core::gnp::DualtreeDfs<mlpack::kde::Kde<TableType> > dualtree_dfs;
|
||||
typedef mlpack::kde::Kde<TableType> ProblemType;
|
||||
core::gnp::DualtreeDfs< ProblemType > dualtree_dfs;
|
||||
dualtree_dfs.Init(*this);
|
||||
|
||||
// Compute the result.
|
||||
dualtree_dfs.Compute(* arguments_in.metric_, result_out);
|
||||
printf("Number of prunes: %d\n", dualtree_dfs.num_deterministic_prunes());
|
||||
if(arguments_in.num_iterations_in_ <= 0) {
|
||||
dualtree_dfs.Compute(* arguments_in.metric_, result_out);
|
||||
printf("Number of prunes: %d\n", dualtree_dfs.num_deterministic_prunes());
|
||||
}
|
||||
else {
|
||||
typename core::gnp::DualtreeDfs<ProblemType>::iterator kde_it =
|
||||
dualtree_dfs.get_iterator(*arguments_in.metric_, result_out);
|
||||
}
|
||||
}
|
||||
|
||||
template<typename TableType>
|
||||
@@ -104,8 +111,11 @@ bool mlpack::kde::Kde<TableType>::ConstructBoostVariableMap_(
|
||||
boost::program_options::value<double>(),
|
||||
"OPTIONAL kernel bandwidth, if you set --bandwidth_selection flag, "
|
||||
"then the --bandwidth will be ignored."
|
||||
)
|
||||
(
|
||||
)(
|
||||
"num_iterations_in",
|
||||
boost::program_options::value<int>()->default_value(0),
|
||||
"The number of iterations to run."
|
||||
)(
|
||||
"probability",
|
||||
boost::program_options::value<double>()->default_value(1.0),
|
||||
"Probability guarantee for the approximation of KDE."
|
||||
@@ -180,7 +190,7 @@ bool mlpack::kde::Kde<TableType>::ConstructBoostVariableMap_(
|
||||
}
|
||||
|
||||
template<typename TableType>
|
||||
void mlpack::kde::Kde<TableType>::ParseArguments(
|
||||
bool mlpack::kde::Kde<TableType>::ParseArguments(
|
||||
const std::vector<std::string> &args,
|
||||
mlpack::kde::KdeArguments<TableType> *arguments_out) {
|
||||
|
||||
@@ -189,7 +199,9 @@ void mlpack::kde::Kde<TableType>::ParseArguments(
|
||||
|
||||
// Construct the Boost variable map.
|
||||
boost::program_options::variables_map vm;
|
||||
ConstructBoostVariableMap_(args, &vm);
|
||||
if(ConstructBoostVariableMap_(args, &vm)) {
|
||||
return true;
|
||||
}
|
||||
|
||||
// Given the constructed boost variable map, parse each argument.
|
||||
|
||||
@@ -247,10 +259,21 @@ void mlpack::kde::Kde<TableType>::ParseArguments(
|
||||
// Parse the kernel type.
|
||||
arguments_out->kernel_ = vm["kernel"].as< std::string >();
|
||||
std::cout << "Using the kernel: " << arguments_out->kernel_ << "\n";
|
||||
|
||||
// Parse the number of iterations.
|
||||
arguments_out->num_iterations_in_ = vm["num_iterations_in"].as<int>();
|
||||
if(arguments_out->num_iterations_in_ > 0) {
|
||||
std::cout << "Running for " << arguments_out->num_iterations_in_ <<
|
||||
" iterations on a progressive mode...\n";
|
||||
}
|
||||
else {
|
||||
std::cout << "Running the algorithm on a non-progressive mode...\n";
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
template<typename TableType>
|
||||
void mlpack::kde::Kde<TableType>::ParseArguments(
|
||||
bool mlpack::kde::Kde<TableType>::ParseArguments(
|
||||
int argc,
|
||||
char *argv[],
|
||||
mlpack::kde::KdeArguments<TableType> *arguments_out) {
|
||||
@@ -258,7 +281,7 @@ void mlpack::kde::Kde<TableType>::ParseArguments(
|
||||
// Convert C input to C++; skip executable name for Boost.
|
||||
std::vector<std::string> args(argv + 1, argv + argc);
|
||||
|
||||
ParseArguments(args, arguments_out);
|
||||
return ParseArguments(args, arguments_out);
|
||||
}
|
||||
|
||||
#endif
|
||||
|
||||
Reference in New Issue
Block a user