Merge pull request #646 from MarcosPividori/traversal-info
Remove duplicated code for traversal info.
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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; }
|
||||
|
||||
@@ -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; }
|
||||
|
||||
@@ -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; }
|
||||
|
||||
Reference in New Issue
Block a user