diff --git a/src/mlpack/core/tree/CMakeLists.txt b/src/mlpack/core/tree/CMakeLists.txt index 6be924c332..0cf7e1bf8c 100644 --- a/src/mlpack/core/tree/CMakeLists.txt +++ b/src/mlpack/core/tree/CMakeLists.txt @@ -21,6 +21,7 @@ set(SOURCES traversers/single_tree_depth_first_traverser.hpp traversers/dual_tree_depth_first_traverser.hpp traversers/dual_tree_breadth_first_traverser.hpp + traversers/dual_cover_tree_traverser.hpp ) # add directory name to sources diff --git a/src/mlpack/core/tree/binary_space_tree.hpp b/src/mlpack/core/tree/binary_space_tree.hpp index fff94f693c..3ca351b021 100644 --- a/src/mlpack/core/tree/binary_space_tree.hpp +++ b/src/mlpack/core/tree/binary_space_tree.hpp @@ -11,6 +11,10 @@ #include "statistic.hpp" #include "traversers/single_tree_depth_first_traverser.hpp" +#include "traversers/dual_tree_depth_first_traverser.hpp" + +// Bad! +#include namespace mlpack { namespace tree /** Trees and tree-building procedures. */ { @@ -74,6 +78,22 @@ class BinarySpaceTree > Type; }; + template + struct PreferredDualTraverser + { + typedef DualTreeDepthFirstTraverser< + BinarySpaceTree, + RuleType + > Type; + }; + + template + struct PreferredRules + { + typedef neighbor::NeighborSearchRules + Type; + }; + /** * Construct this as the root node of a binary space tree using the given * dataset. This will modify the ordering of the points in the dataset! diff --git a/src/mlpack/core/tree/traversers/dual_cover_tree_traverser.hpp b/src/mlpack/core/tree/traversers/dual_cover_tree_traverser.hpp new file mode 100644 index 0000000000..e5635cb618 --- /dev/null +++ b/src/mlpack/core/tree/traversers/dual_cover_tree_traverser.hpp @@ -0,0 +1,91 @@ +/** + * @file dual_cover_tree_traverser.hpp + * @author Ryan Curtin + * + * A dual-tree traverser for the cover tree. + */ +#ifndef __MLPACK_CORE_TREE_DUAL_COVER_TREE_TRAVERSER_HPP +#define __MLPACK_CORE_TREE_DUAL_COVER_TREE_TRAVERSER_HPP + +#include +#include + +namespace mlpack { +namespace tree { + +template +class DualCoverTreeTraverser +{ + public: + DualCoverTreeTraverser(RuleType& rule) : rule(rule), numPrunes(0) { } + + void Traverse(TreeType& queryNode, TreeType& referenceNode) + { + Traverse(queryNode, referenceNode, size_t() - 1); + } + + void Traverse(TreeType& queryNode, TreeType& referenceNode, size_t parent) + { + std::queue referenceQueue; + std::queue referenceParents; + + referenceQueue.push(&referenceNode); + referenceParents.push(parent); + + while (!referenceQueue.empty()) + { + TreeType& reference = *referenceQueue.front(); + referenceQueue.pop(); + + size_t refParent = referenceParents.front(); + referenceParents.pop(); + + // Do the base case, if we need to. + if (refParent != reference.Point()) + rule.BaseCase(queryNode.Point(), reference.Point()); + + if (((queryNode.Scale() < reference.Scale()) && + (reference.NumChildren() != 0)) || + (queryNode.NumChildren() == 0)) + { + // We must descend the reference node. Pruning happens here. + for (size_t i = 0; i < reference.NumChildren(); ++i) + { + // Can we prune? + // if (!rule.CanPrune(queryNode, reference.Child(i))) + { + referenceQueue.push(&(reference.Child(i))); + referenceParents.push(reference.Point()); + } +// else + { + numPrunes++; + } + } + } + else + { + // We must descend the query node. No pruning happens here. For the + // self-child, we trick the recursion into thinking that the base case + // has already been done (which it has). + if (queryNode.NumChildren() >= 1) + Traverse(queryNode.Child(0), reference, reference.Point()); + + for (size_t i = 1; i < queryNode.NumChildren(); ++i) + Traverse(queryNode.Child(i), reference, size_t() - 1); + } + } + } + + size_t NumPrunes() const { return numPrunes; } + size_t& NumPrunes() { return numPrunes; } + + private: + RuleType& rule; + size_t numPrunes; +}; + +}; // namespace tree +}; // namespace mlpack + +#endif