Another checkpoint, getting there.

This commit is contained in:
Dongryeol Lee
2010-12-29 17:55:51 +00:00
parent 77769ed20b
commit 83f39cb181
5 changed files with 75 additions and 47 deletions
@@ -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