Merge pull request #646 from MarcosPividori/traversal-info

Remove duplicated code for traversal info.
This commit is contained in:
Ryan Curtin
2016-05-24 17:45:20 -04:00
10 changed files with 18 additions and 86 deletions
+6
View File
@@ -9,6 +9,9 @@
#ifndef MLPACK_CORE_TREE_TRAVERSAL_INFO_HPP
#define MLPACK_CORE_TREE_TRAVERSAL_INFO_HPP
namespace mlpack {
namespace tree {
/**
* The TraversalInfo class holds traversal information which is used in
* dual-tree (and single-tree) traversals. A traversal should be updating the
@@ -82,4 +85,7 @@ class TraversalInfo
double lastBaseCase;
};
} // namespace tree
} // namespace mlpack
#endif
+2 -2
View File
@@ -9,7 +9,7 @@
#include <mlpack/core.hpp>
#include "../neighbor_search/ns_traversal_info.hpp"
#include <mlpack/core/tree/traversal_info.hpp>
namespace mlpack {
namespace emst {
@@ -105,7 +105,7 @@ class DTBRules
TreeType& referenceNode,
const double oldScore) const;
typedef neighbor::NeighborSearchTraversalInfo<TreeType> TraversalInfoType;
typedef typename tree::TraversalInfo<TreeType> TraversalInfoType;
const TraversalInfoType& TraversalInfo() const { return traversalInfo; }
TraversalInfoType& TraversalInfo() { return traversalInfo; }
+2 -3
View File
@@ -9,8 +9,7 @@
#include <mlpack/core.hpp>
#include <mlpack/core/tree/cover_tree/cover_tree.hpp>
#include "../neighbor_search/ns_traversal_info.hpp"
#include <mlpack/core/tree/traversal_info.hpp>
namespace mlpack {
namespace fastmks {
@@ -91,7 +90,7 @@ class FastMKSRules
//! Modify the number of times Score() was called.
size_t& Scores() { return scores; }
typedef neighbor::NeighborSearchTraversalInfo<TreeType> TraversalInfoType;
typedef typename tree::TraversalInfo<TreeType> TraversalInfoType;
const TraversalInfoType& TraversalInfo() const { return traversalInfo; }
TraversalInfoType& TraversalInfo() { return traversalInfo; }
@@ -9,7 +9,7 @@
#ifndef MLPACK_METHODS_KMEANS_DUAL_TREE_KMEANS_RULES_HPP
#define MLPACK_METHODS_KMEANS_DUAL_TREE_KMEANS_RULES_HPP
#include <mlpack/methods/neighbor_search/ns_traversal_info.hpp>
#include <mlpack/core/tree/traversal_info.hpp>
namespace mlpack {
namespace kmeans {
@@ -39,7 +39,7 @@ class DualTreeKMeansRules
TreeType& referenceNode,
const double oldScore);
typedef neighbor::NeighborSearchTraversalInfo<TreeType> TraversalInfoType;
typedef typename tree::TraversalInfo<TreeType> TraversalInfoType;
TraversalInfoType& TraversalInfo() { return traversalInfo; }
const TraversalInfoType& TraversalInfo() const { return traversalInfo; }
@@ -9,8 +9,6 @@
#ifndef MLPACK_METHODS_KMEANS_PELLEG_MOORE_KMEANS_RULES_HPP
#define MLPACK_METHODS_KMEANS_PELLEG_MOORE_KMEANS_RULES_HPP
#include <mlpack/methods/neighbor_search/ns_traversal_info.hpp>
namespace mlpack {
namespace kmeans {
@@ -8,7 +8,6 @@ set(SOURCES
neighbor_search_stat.hpp
ns_model.hpp
ns_model_impl.hpp
ns_traversal_info.hpp
sort_policies/nearest_neighbor_sort.hpp
sort_policies/nearest_neighbor_sort.cpp
sort_policies/nearest_neighbor_sort_impl.hpp
@@ -8,7 +8,7 @@
#ifndef MLPACK_METHODS_NEIGHBOR_SEARCH_NEIGHBOR_SEARCH_RULES_HPP
#define MLPACK_METHODS_NEIGHBOR_SEARCH_NEIGHBOR_SEARCH_RULES_HPP
#include "ns_traversal_info.hpp"
#include <mlpack/core/tree/traversal_info.hpp>
namespace mlpack {
namespace neighbor {
@@ -94,7 +94,7 @@ class NeighborSearchRules
size_t& Scores() { return scores; }
//! Convenience typedef.
typedef NeighborSearchTraversalInfo<TreeType> TraversalInfoType;
typedef typename tree::TraversalInfo<TreeType> TraversalInfoType;
//! Get the traversal info.
const TraversalInfoType& TraversalInfo() const { return traversalInfo; }
@@ -1,70 +0,0 @@
/**
* @file ns_traversal_info.hpp
* @author Ryan Curtin
*
* This class holds traversal information for dual-tree traversals that are
* using the NeighborSearchRules RuleType.
*/
#ifndef MLPACK_METHODS_NEIGHBOR_SEARCH_TRAVERSAL_INFO_HPP
#define MLPACK_METHODS_NEIGHBOR_SEARCH_TRAVERSAL_INFO_HPP
namespace mlpack {
namespace neighbor {
/**
* Traversal information for NeighborSearch. This information is used to make
* parent-child prunes or parent-parent prunes in Score() without needing to
* evaluate the distance between two nodes.
*
* The information held by this class is the last node combination visited
* before the current node combination was recursed into and the distance
* between the node centroids.
*/
template<typename TreeType>
class NeighborSearchTraversalInfo
{
public:
/**
* Create the TraversalInfo object and initialize the pointers to NULL.
*/
NeighborSearchTraversalInfo() :
lastQueryNode(NULL),
lastReferenceNode(NULL),
lastScore(0.0),
lastBaseCase(0.0) { /* Nothing to do. */ }
//! Get the last query node.
TreeType* LastQueryNode() const { return lastQueryNode; }
//! Modify the last query node.
TreeType*& LastQueryNode() { return lastQueryNode; }
//! Get the last reference node.
TreeType* LastReferenceNode() const { return lastReferenceNode; }
//! Modify the last reference node.
TreeType*& LastReferenceNode() { return lastReferenceNode; }
//! Get the score associated with the last query and reference nodes.
double LastScore() const { return lastScore; }
//! Modify the score associated with the last query and reference nodes.
double& LastScore() { return lastScore; }
//! Get the base case associated with the last node combination.
double LastBaseCase() const { return lastBaseCase; }
//! Modify the base case associated with the last node combination.
double& LastBaseCase() { return lastBaseCase; }
private:
//! The last query node.
TreeType* lastQueryNode;
//! The last reference node.
TreeType* lastReferenceNode;
//! The last distance.
double lastScore;
//! The last base case.
double lastBaseCase;
};
} // namespace neighbor
} // namespace mlpack
#endif
@@ -7,7 +7,7 @@
#ifndef MLPACK_METHODS_RANGE_SEARCH_RANGE_SEARCH_RULES_HPP
#define MLPACK_METHODS_RANGE_SEARCH_RANGE_SEARCH_RULES_HPP
#include "../neighbor_search/ns_traversal_info.hpp"
#include <mlpack/core/tree/traversal_info.hpp>
namespace mlpack {
namespace range {
@@ -96,7 +96,7 @@ class RangeSearchRules
TreeType& referenceNode,
const double oldScore) const;
typedef neighbor::NeighborSearchTraversalInfo<TreeType> TraversalInfoType;
typedef typename tree::TraversalInfo<TreeType> TraversalInfoType;
const TraversalInfoType& TraversalInfo() const { return traversalInfo; }
TraversalInfoType& TraversalInfo() { return traversalInfo; }
+2 -2
View File
@@ -9,7 +9,7 @@
#ifndef MLPACK_METHODS_RANN_RA_SEARCH_RULES_HPP
#define MLPACK_METHODS_RANN_RA_SEARCH_RULES_HPP
#include "../neighbor_search/ns_traversal_info.hpp"
#include <mlpack/core/tree/traversal_info.hpp>
namespace mlpack {
namespace neighbor {
@@ -185,7 +185,7 @@ class RASearchRules
return arma::sum(numSamplesMade);
}
typedef neighbor::NeighborSearchTraversalInfo<TreeType> TraversalInfoType;
typedef typename tree::TraversalInfo<TreeType> TraversalInfoType;
const TraversalInfoType& TraversalInfo() const { return traversalInfo; }
TraversalInfoType& TraversalInfo() { return traversalInfo; }