hi
This commit is contained in:
@@ -0,0 +1,353 @@
|
||||
// Copyright 2007 Georgia Institute of Technology. All rights reserved.
|
||||
// ABSOLUTELY NOT FOR DISTRIBUTION
|
||||
/**
|
||||
* @param bounds.h
|
||||
*
|
||||
* Bounds that are useful for binary space partitioning trees.
|
||||
*
|
||||
* TODO: Come up with a better design so you can do plug-and-play distance
|
||||
* metrics.
|
||||
*/
|
||||
|
||||
#ifndef TREE_BOUNDS_H
|
||||
#define TREE_BOUNDS_H
|
||||
|
||||
#include "la/matrix.h"
|
||||
#include "la/la.h"
|
||||
|
||||
/**
|
||||
* Simple real-valued range.
|
||||
*/
|
||||
struct DBound {
|
||||
public:
|
||||
double lo;
|
||||
double hi;
|
||||
|
||||
public:
|
||||
DBound() {}
|
||||
|
||||
void Init() {
|
||||
lo = DBL_MAX;
|
||||
hi = -DBL_MAX;
|
||||
}
|
||||
|
||||
double width() const {
|
||||
return hi - lo;
|
||||
}
|
||||
|
||||
double mid() const {
|
||||
return (hi + lo) / 2;
|
||||
}
|
||||
};
|
||||
|
||||
/**
|
||||
* Hyper-rectangle bound.
|
||||
*/
|
||||
class DHrectBound {
|
||||
private:
|
||||
DBound *bounds_;
|
||||
//double diagonal_sq_;
|
||||
index_t dim_;
|
||||
|
||||
public:
|
||||
DHrectBound() {
|
||||
DEBUG_POISON_PTR(bounds_);
|
||||
DEBUG_ONLY(dim_ = BIG_BAD_NUMBER);
|
||||
}
|
||||
|
||||
~DHrectBound() {
|
||||
mem::Free(bounds_);
|
||||
}
|
||||
|
||||
template<typename Deserializer>
|
||||
void Deserialize(Deserializer *s) {
|
||||
DEBUG_ASSERT_MSG(dim_ == BIG_BAD_NUMBER, "Already initialized");
|
||||
|
||||
s->Get(&dim_);
|
||||
bounds_ = mem::Alloc<DBound>(dim_);
|
||||
s->Get(bounds_, dim_);
|
||||
|
||||
//ComputeDiagonal_();
|
||||
}
|
||||
|
||||
template<typename Serializer>
|
||||
void Serialize(Serializer *s) const {
|
||||
s->Put(dim_);
|
||||
s->Put(bounds_, dim_);
|
||||
}
|
||||
|
||||
void Init(index_t dimension) {
|
||||
DEBUG_ASSERT_MSG(dim_ == BIG_BAD_NUMBER, "Already initialized");
|
||||
|
||||
bounds_ = mem::Alloc<DBound>(dimension);
|
||||
|
||||
for (index_t i = 0; i < dimension; i++) {
|
||||
bounds_[i].Init();
|
||||
}
|
||||
|
||||
dim_ = dimension;
|
||||
|
||||
//ComputeDiagonal_();
|
||||
}
|
||||
|
||||
bool Belongs(const Vector& point) const {
|
||||
for (index_t i = 0; i < point.length(); i++) {
|
||||
const DBound *bound = &bounds_[i];
|
||||
if (point[i] > bound->hi || point[i] < bound->lo) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
double MinDistanceSqToInstance(const Vector& point) const {
|
||||
DEBUG_ASSERT(point.length() == dim_);
|
||||
return MinDistanceSqToInstance(point.ptr());
|
||||
}
|
||||
|
||||
double MinDistanceSqToInstance(const double *mpoint) const {
|
||||
double sumsq = 0;
|
||||
//index_t mdim = dim_;
|
||||
const DBound *mbound = bounds_;
|
||||
|
||||
index_t d = dim_;
|
||||
|
||||
do {
|
||||
double v = *mpoint;
|
||||
double v1 = mbound->lo - v;
|
||||
double v2 = v - mbound->hi;
|
||||
|
||||
v = (v1 + fabs(v1)) + (v2 + fabs(v2));
|
||||
|
||||
mbound++;
|
||||
mpoint++;
|
||||
|
||||
sumsq += v * v;
|
||||
} while (--d);
|
||||
|
||||
return sumsq / 4;
|
||||
}
|
||||
|
||||
double MaxDistanceSqToInstance(const Vector& point) const {
|
||||
double sumsq = 0;
|
||||
|
||||
DEBUG_ASSERT(point.length() == dim_);
|
||||
|
||||
for (index_t d = 0; d < dim_; d++) {
|
||||
double v = max(point[d] - bounds_[d].lo,
|
||||
bounds_[d].hi - point[d]);
|
||||
|
||||
sumsq += v * v;
|
||||
}
|
||||
|
||||
return sumsq;
|
||||
}
|
||||
|
||||
double MinDistanceSqToBound(const DHrectBound& other) const {
|
||||
double sumsq = 0;
|
||||
const DBound *a = this->bounds_;
|
||||
const DBound *b = other.bounds_;
|
||||
index_t mdim = dim_;
|
||||
|
||||
DEBUG_ASSERT(dim_ == other.dim_);
|
||||
|
||||
// We invoke the following:
|
||||
// x + fabs(x) = max(x * 2, 0)
|
||||
// (x * 2)^2 / 4 = x^2
|
||||
|
||||
for (index_t d = 0; d < mdim; d++) {
|
||||
#if 0
|
||||
double v = b[d].lo - a[d].hi;
|
||||
|
||||
if (v < 0) {
|
||||
v = a[d].lo - b[d].hi;
|
||||
}
|
||||
|
||||
if (likely(v > 0)) {
|
||||
sumsq += v * v;
|
||||
}
|
||||
#else
|
||||
double v1 = b[d].lo - a[d].hi;
|
||||
double v2 = a[d].lo - b[d].hi;
|
||||
|
||||
double v = (v1 + fabs(v1)) + (v2 + fabs(v2));
|
||||
|
||||
sumsq += v * v;
|
||||
#endif
|
||||
}
|
||||
|
||||
return sumsq / 4;
|
||||
}
|
||||
|
||||
double MinDistanceSqToBoundFarEnd(const DHrectBound& other) const {
|
||||
double sumsq = 0;
|
||||
const DBound *a = this->bounds_;
|
||||
const DBound *b = other.bounds_;
|
||||
index_t mdim = dim_;
|
||||
|
||||
DEBUG_ASSERT(dim_ == other.dim_);
|
||||
|
||||
for (index_t d = 0; d < mdim; d++) {
|
||||
double v1 = b[d].hi - a[d].hi;
|
||||
double v2 = a[d].lo - b[d].lo;
|
||||
|
||||
double v = max(v1, v2);
|
||||
v = (v + fabs(v)); /* truncate negative */
|
||||
|
||||
sumsq += v * v;
|
||||
}
|
||||
|
||||
return sumsq / 4;
|
||||
}
|
||||
|
||||
double MaxDistanceSqToBound(const DHrectBound& other) const {
|
||||
double sumsq = 0;
|
||||
const DBound *a = this->bounds_;
|
||||
const DBound *b = other.bounds_;
|
||||
|
||||
DEBUG_ASSERT(dim_ == other.dim_);
|
||||
|
||||
for (index_t d = 0; d < dim_; d++) {
|
||||
double v = max(b[d].hi - a[d].lo, a[d].hi - b[d].lo);
|
||||
|
||||
sumsq += v * v;
|
||||
}
|
||||
|
||||
return sumsq;
|
||||
}
|
||||
|
||||
double MidDistanceSqToBound(const DHrectBound& other) const {
|
||||
double sumsq = 0;
|
||||
const DBound *a = this->bounds_;
|
||||
const DBound *b = other.bounds_;
|
||||
|
||||
DEBUG_ASSERT(dim_ == other.dim_);
|
||||
|
||||
for (index_t d = 0; d < dim_; d++) {
|
||||
double v = (a[d].hi + a[d].lo - b[d].hi - b[d].lo) * 0.5;
|
||||
|
||||
sumsq += v * v;
|
||||
}
|
||||
|
||||
return sumsq;
|
||||
}
|
||||
|
||||
void Update(const Vector& vector) {
|
||||
DEBUG_ASSERT(vector.length() == dim_);
|
||||
|
||||
for (index_t i = 0; i < dim_; i++) {
|
||||
DBound* bound = &bounds_[i];
|
||||
double d = vector[i];
|
||||
|
||||
if (unlikely(d > bound->hi)) {
|
||||
bound->hi = d;
|
||||
}
|
||||
if (unlikely(d < bound->lo)) {
|
||||
bound->lo = d;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const DBound& get(index_t i) const {
|
||||
return bounds_[i];
|
||||
}
|
||||
|
||||
//double diagonal_sq() const {
|
||||
// return diagonal_sq_;
|
||||
//}
|
||||
|
||||
FORBID_COPY(DHrectBound);
|
||||
|
||||
private:
|
||||
//void ComputeDiagonal_() {
|
||||
// diagonal_sq_ = 0;
|
||||
// for (index_t d = 0; d < dim_; d++) {
|
||||
// double v = bounds_[d].lo - bounds_[d].hi;
|
||||
// diagonal_sq_ += v*v;
|
||||
// }
|
||||
//}
|
||||
};
|
||||
|
||||
/**
|
||||
* Euclidean metric for use with ball bounds.
|
||||
*/
|
||||
class DEuclideanMetric {
|
||||
public:
|
||||
static double CalculateMetric(const Vector& a, const Vector& b) {
|
||||
return sqrt(la::DistanceSqEuclidean(a.ptr(), b.ptr(), a.length()));
|
||||
}
|
||||
};
|
||||
|
||||
/**
|
||||
* Bound of a ball tree.
|
||||
*/
|
||||
template<class TInstance, class TMetric>
|
||||
class BallBound {
|
||||
FORBID_COPY(BallBound);
|
||||
|
||||
public:
|
||||
typedef TMetric Metric;
|
||||
typedef TInstance Instance;
|
||||
|
||||
private:
|
||||
Instance center_;
|
||||
double radius_;
|
||||
|
||||
public:
|
||||
BallBound() {}
|
||||
|
||||
const Instance& center() const {
|
||||
return center;
|
||||
}
|
||||
|
||||
Instance& center() {
|
||||
return center;
|
||||
}
|
||||
|
||||
double radius() const {
|
||||
return radius;
|
||||
}
|
||||
|
||||
void set_radius(double d) {
|
||||
radius = d;
|
||||
}
|
||||
|
||||
double DistanceToCenter(const Instance& point) {
|
||||
return Metric::CalculateMetric(point, center_);
|
||||
}
|
||||
|
||||
bool Belongs(const Instance& point) {
|
||||
return DistanceToCenter(point) <= radius_;
|
||||
}
|
||||
|
||||
double MinDistanceToInstance(const Instance& point) {
|
||||
return max(0.0, DistanceToCenter(point) - radius_);
|
||||
}
|
||||
|
||||
double MaxDistanceToInstance(const Instance& point) {
|
||||
return DistanceToCenter(point) + radius_;
|
||||
}
|
||||
|
||||
double MinDistanceToBound(const BallBound& ball) {
|
||||
return max(0,
|
||||
DistanceToCenter(ball.center_) - (radius_ + ball.radius_));
|
||||
}
|
||||
|
||||
double MaxDistanceToBound(const BallBound& ball) {
|
||||
return DistanceToCenter(ball.center_) + (radius_ + ball.radius_);
|
||||
}
|
||||
|
||||
double MidDistanceToBound(const BallBound& other) {
|
||||
return DistanceToCenter(other.center_);
|
||||
}
|
||||
|
||||
double MidDistanceToInstance(const Instance& point) {
|
||||
return DistanceToCenter(point);
|
||||
}
|
||||
};
|
||||
|
||||
typedef BallBound<Vector, DEuclideanMetric> DEuclideanBallBound;
|
||||
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,10 @@
|
||||
|
||||
librule(
|
||||
sources = [],
|
||||
headers = [
|
||||
"kdtree.h", "bounds.h", "spacetree.h", "statistic.h"
|
||||
],
|
||||
deplibs = ["base:base", "la:la", "col:col",
|
||||
"file:file"]
|
||||
)
|
||||
|
||||
@@ -0,0 +1,253 @@
|
||||
// Copyright 2007 Georgia Institute of Technology. All rights reserved.
|
||||
// ABSOLUTELY NOT FOR DISTRIBUTION
|
||||
/**
|
||||
* @file kdtree.h
|
||||
*
|
||||
* Tools for kd-trees.
|
||||
*
|
||||
* Eventually we hope to support KD trees with non-L2 (Euclidean)
|
||||
* metrics, like Manhattan distance.
|
||||
*/
|
||||
|
||||
#ifndef TREE_KDTREE_H
|
||||
#define TREE_KDTREE_H
|
||||
|
||||
#include "spacetree.h"
|
||||
#include "bounds.h"
|
||||
|
||||
#include "base/common.h"
|
||||
#include "col/arraylist.h"
|
||||
#include "file/serialize.h"
|
||||
|
||||
/* Implementation */
|
||||
|
||||
namespace tree_kdtree_private {
|
||||
|
||||
template<typename TBound>
|
||||
void FindBoundFromMatrix(const Matrix& matrix,
|
||||
index_t first, index_t count, TBound *bounds) {
|
||||
index_t end = first + count;
|
||||
for (index_t i = first; i < end; i++) {
|
||||
Vector col;
|
||||
|
||||
matrix.MakeColumnVector(i, &col);
|
||||
bounds->Update(col);
|
||||
}
|
||||
}
|
||||
|
||||
template<typename TBound>
|
||||
index_t MatrixPartition(
|
||||
Matrix& matrix, index_t dim, double splitvalue,
|
||||
index_t first, index_t count,
|
||||
TBound* left_bound, TBound* right_bound,
|
||||
index_t *old_from_new) {
|
||||
index_t left = first;
|
||||
index_t right = first + count - 1;
|
||||
|
||||
/* At any point:
|
||||
*
|
||||
* everything < left is correct
|
||||
* everything > right is correct
|
||||
*/
|
||||
for (;;) {
|
||||
while (matrix.get(dim, left) < splitvalue && likely(left <= right)) {
|
||||
Vector left_vector;
|
||||
matrix.MakeColumnVector(left, &left_vector);
|
||||
left_bound->Update(left_vector);
|
||||
left++;
|
||||
}
|
||||
|
||||
while (matrix.get(dim, right) >= splitvalue && likely(left <= right)) {
|
||||
Vector right_vector;
|
||||
matrix.MakeColumnVector(right, &right_vector);
|
||||
right_bound->Update(right_vector);
|
||||
right--;
|
||||
}
|
||||
|
||||
if (unlikely(left > right)) {
|
||||
/* left == right + 1 */
|
||||
break;
|
||||
}
|
||||
|
||||
Vector left_vector;
|
||||
Vector right_vector;
|
||||
|
||||
matrix.MakeColumnVector(left, &left_vector);
|
||||
matrix.MakeColumnVector(right, &right_vector);
|
||||
|
||||
left_vector.SwapValues(&right_vector);
|
||||
|
||||
left_bound->Update(left_vector);
|
||||
right_bound->Update(right_vector);
|
||||
|
||||
if (old_from_new) {
|
||||
index_t t = old_from_new[left];
|
||||
old_from_new[left] = old_from_new[right];
|
||||
old_from_new[right] = t;
|
||||
}
|
||||
|
||||
DEBUG_ASSERT(left <= right);
|
||||
right--;
|
||||
|
||||
// this conditional is always true, I belueve
|
||||
//if (likely(left <= right)) {
|
||||
// right--;
|
||||
//}
|
||||
}
|
||||
|
||||
DEBUG_ASSERT(left == right + 1);
|
||||
|
||||
return left;
|
||||
}
|
||||
|
||||
template<typename TKdTree>
|
||||
void SplitKdTreeMidpoint(Matrix& matrix,
|
||||
TKdTree *node, index_t leaf_size, index_t *old_from_new) {
|
||||
TKdTree *left = NULL;
|
||||
TKdTree *right = NULL;
|
||||
|
||||
//FindBoundFromMatrix(matrix, node->begin(), node->count(),
|
||||
// &node->bound());
|
||||
|
||||
if (node->count() > leaf_size) {
|
||||
index_t split_dim = BIG_BAD_NUMBER;
|
||||
double max_width = -1;
|
||||
|
||||
for (index_t d = 0; d < matrix.n_rows(); d++) {
|
||||
double w = node->bound().get(d).width();
|
||||
|
||||
if (unlikely(w > max_width)) {
|
||||
max_width = w;
|
||||
split_dim = d;
|
||||
}
|
||||
}
|
||||
|
||||
double split_val = node->bound().get(split_dim).mid();
|
||||
|
||||
if (max_width == 0) {
|
||||
// Okay, we can't do any splitting, because all these points are the
|
||||
// same. We have to give up.
|
||||
} else {
|
||||
left = new TKdTree();
|
||||
left->bound().Init(matrix.n_rows());
|
||||
|
||||
right = new TKdTree();
|
||||
right->bound().Init(matrix.n_rows());
|
||||
|
||||
index_t split_col = MatrixPartition(matrix, split_dim, split_val,
|
||||
node->begin(), node->count(),
|
||||
&left->bound(), &right->bound(),
|
||||
old_from_new);
|
||||
|
||||
DEBUG_MSG(3.0,"split (%d,[%d],%d) dim %d on %f (between %f, %f)",
|
||||
node->begin(), split_col,
|
||||
node->begin() + node->count(), split_dim, split_val,
|
||||
node->bound().get(split_dim).lo,
|
||||
node->bound().get(split_dim).hi);
|
||||
|
||||
left->Init(node->begin(), split_col - node->begin());
|
||||
right->Init(split_col, node->begin() + node->count() - split_col);
|
||||
|
||||
// This should never happen if max_width > 0
|
||||
DEBUG_ASSERT(left->count() != 0 && right->count() != 0);
|
||||
|
||||
SplitKdTreeMidpoint(matrix, left, leaf_size, old_from_new);
|
||||
SplitKdTreeMidpoint(matrix, right, leaf_size, old_from_new);
|
||||
}
|
||||
}
|
||||
|
||||
node->set_children(matrix, left, right);
|
||||
}
|
||||
};
|
||||
|
||||
namespace tree {
|
||||
/**
|
||||
* Creates a KD tree from data, splitting on the midpoint.
|
||||
*
|
||||
* This requires you to pass in two unitialized ArrayLists which will contain
|
||||
* index mappings so you can account for the re-ordering of the matrix.
|
||||
* (By unitialized I mean don't call Init on it)
|
||||
*
|
||||
* @param matrix data where each column is a point, WHICH WILL BE RE-ORDERED
|
||||
* @param old_from_new pointer to an unitialized arraylist; it will map
|
||||
* original indexes to new indices
|
||||
* @param old_from_new pointer to an unitialized arraylist; it will map
|
||||
* new indices to original
|
||||
*/
|
||||
template<typename TKdTree>
|
||||
TKdTree *MakeKdTreeMidpoint(Matrix& matrix, index_t leaf_size,
|
||||
ArrayList<index_t> *old_from_new = NULL,
|
||||
ArrayList<index_t> *new_from_old = NULL) {
|
||||
TKdTree *node = new TKdTree();
|
||||
index_t *old_from_new_ptr;
|
||||
|
||||
if (old_from_new) {
|
||||
old_from_new->Init(matrix.n_cols());
|
||||
|
||||
for (index_t i = 0; i < matrix.n_cols(); i++) {
|
||||
(*old_from_new)[i] = i;
|
||||
}
|
||||
|
||||
old_from_new_ptr = old_from_new->begin();
|
||||
} else {
|
||||
old_from_new_ptr = NULL;
|
||||
}
|
||||
|
||||
node->Init(0, matrix.n_cols());
|
||||
node->bound().Init(matrix.n_rows());
|
||||
tree_kdtree_private::FindBoundFromMatrix(matrix,
|
||||
0, matrix.n_cols(), &node->bound());
|
||||
|
||||
tree_kdtree_private::SplitKdTreeMidpoint(matrix, node, leaf_size,
|
||||
old_from_new_ptr);
|
||||
|
||||
if (new_from_old) {
|
||||
new_from_old->Init(matrix.n_cols());
|
||||
for (index_t i = 0; i < matrix.n_cols(); i++) {
|
||||
(*new_from_old)[(*old_from_new)[i]] = i;
|
||||
}
|
||||
}
|
||||
|
||||
return node;
|
||||
}
|
||||
|
||||
// TODO: Perhaps move this into a "util.h" file
|
||||
|
||||
template<typename TKdTree, typename Serializer>
|
||||
void SerializeKdTree(const TKdTree *tree,
|
||||
const Matrix& matrix,
|
||||
const ArrayList<index_t>& old_from_new,
|
||||
Serializer *s) {
|
||||
s->PutMagic(file::CreateMagic("kdtree"));
|
||||
tree->SerializeAll(matrix, s);
|
||||
old_from_new.Serialize(s);
|
||||
}
|
||||
|
||||
template<typename TKdTree, typename Deserializer>
|
||||
void DeserializeKdTree(TKdTree *uninit_tree,
|
||||
Matrix* uninit_matrix,
|
||||
ArrayList<index_t>* uninit_old_from_new,
|
||||
Deserializer *s) {
|
||||
s->AssertMagic(file::CreateMagic("kdtree"));
|
||||
uninit_tree->DeserializeAll(uninit_matrix, s);
|
||||
if (uninit_old_from_new) {
|
||||
uninit_old_from_new->Deserialize(s);
|
||||
}
|
||||
}
|
||||
|
||||
template<typename TKdTree>
|
||||
void ReadKdTreeFromFile(const char *fname,
|
||||
TKdTree *uninit_tree,
|
||||
Matrix* uninit_matrix,
|
||||
ArrayList<index_t>* uninit_old_from_new) {
|
||||
NativeFileDeserializer ds;
|
||||
|
||||
ASSERT_PASS(ds.Init(fname));
|
||||
DeserializeKdTree(uninit_tree, uninit_matrix, uninit_old_from_new,
|
||||
&ds);
|
||||
}
|
||||
};
|
||||
|
||||
typedef BinarySpaceTree<DHrectBound, Matrix> BasicKdTree;
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,279 @@
|
||||
// Copyright 2007 Georgia Institute of Technology. All rights reserved.
|
||||
// ABSOLUTELY NOT FOR DISTRIBUTION
|
||||
/**
|
||||
* @file spacetree.h
|
||||
*
|
||||
* Generalized space partitioning tree.
|
||||
*/
|
||||
|
||||
#ifndef TREE_SPACETREE_H
|
||||
#define TREE_SPACETREE_H
|
||||
|
||||
#include "base/cc.h"
|
||||
#include "file/serialize.h"
|
||||
#include "statistic.h"
|
||||
|
||||
/**
|
||||
* A binary space partitioning tree, such as KD or ball tree.
|
||||
*
|
||||
* This particular tree forbids you from having more children.
|
||||
*
|
||||
* @param TBound the bounding type of each child (TODO explain interface)
|
||||
* @param TDataset the data set type
|
||||
* @param TStatistic extra data in the node
|
||||
*/
|
||||
template<class TBound,
|
||||
class TDataset,
|
||||
class TStatistic = EmptyStatistic<TDataset> >
|
||||
class BinarySpaceTree {
|
||||
public:
|
||||
typedef TBound Bound;
|
||||
typedef TDataset Dataset;
|
||||
typedef TStatistic Statistic;
|
||||
|
||||
private:
|
||||
Bound bound_;
|
||||
BinarySpaceTree *left_;
|
||||
BinarySpaceTree *right_;
|
||||
index_t begin_;
|
||||
index_t count_;
|
||||
Statistic stat_;
|
||||
|
||||
public:
|
||||
BinarySpaceTree() {
|
||||
DEBUG_ONLY(begin_ = BIG_BAD_NUMBER);
|
||||
DEBUG_ONLY(count_ = BIG_BAD_NUMBER);
|
||||
DEBUG_POISON_PTR(left_);
|
||||
DEBUG_POISON_PTR(right_);
|
||||
}
|
||||
|
||||
|
||||
~BinarySpaceTree() {
|
||||
if (!is_leaf()) {
|
||||
delete left_;
|
||||
delete right_;
|
||||
}
|
||||
DEBUG_ONLY(begin_ = BIG_BAD_NUMBER);
|
||||
DEBUG_ONLY(count_ = BIG_BAD_NUMBER);
|
||||
DEBUG_POISON_PTR(left_);
|
||||
DEBUG_POISON_PTR(right_);
|
||||
}
|
||||
|
||||
void Init(index_t begin_in, index_t count_in) {
|
||||
DEBUG_ASSERT(begin_ == BIG_BAD_NUMBER);
|
||||
DEBUG_POISON_PTR(left_);
|
||||
DEBUG_POISON_PTR(right_);
|
||||
begin_ = begin_in;
|
||||
count_ = count_in;
|
||||
}
|
||||
|
||||
/**
|
||||
* Find a node in this tree by its begin and count.
|
||||
*
|
||||
* Every node is uniquely identified by these two numbers.
|
||||
* This is useful for communicating position over the network,
|
||||
* when pointers would be invalid.
|
||||
*
|
||||
* @param begin_q the begin() of the node to find
|
||||
* @param count_q the count() of the node to find
|
||||
* @return the found node, or NULL
|
||||
*/
|
||||
const BinarySpaceTree* FindByBeginCount(
|
||||
index_t begin_q, index_t count_q) const {
|
||||
DEBUG_ASSERT(begin_q >= begin_);
|
||||
DEBUG_ASSERT(count_q <= count_);
|
||||
if (begin_ == begin_q && count_ == count_q) {
|
||||
return this;
|
||||
} else if (unlikely(is_leaf())) {
|
||||
return NULL;
|
||||
} else if (begin_q < right_->begin_) {
|
||||
return left_->FindByBeginCount(begin_q, count_q);
|
||||
} else {
|
||||
return right_->FindByBeginCount(begin_q, count_q);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Find a node in this tree by its begin and count (const).
|
||||
*
|
||||
* Every node is uniquely identified by these two numbers.
|
||||
* This is useful for communicating position over the network,
|
||||
* when pointers would be invalid.
|
||||
*
|
||||
* @param begin_q the begin() of the node to find
|
||||
* @param count_q the count() of the node to find
|
||||
* @return the found node, or NULL
|
||||
*/
|
||||
BinarySpaceTree* FindByBeginCount(
|
||||
index_t begin_q, index_t count_q) {
|
||||
DEBUG_ASSERT(begin_q >= begin_);
|
||||
DEBUG_ASSERT(count_q <= count_);
|
||||
if (begin_ == begin_q && count_ == count_q) {
|
||||
return this;
|
||||
} else if (unlikely(is_leaf())) {
|
||||
return NULL;
|
||||
} else if (begin_q < right_->begin_) {
|
||||
return left_->FindByBeginCount(begin_q, count_q);
|
||||
} else {
|
||||
return right_->FindByBeginCount(begin_q, count_q);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Serializes the tree _structure_ only.
|
||||
* Statistics are not stored (this allows you to re-load the tree
|
||||
* for problems that require different statistics).
|
||||
*/
|
||||
template<typename Serializer>
|
||||
void Serialize(Serializer *s) const {
|
||||
bound_.Serialize(s);
|
||||
s->Put(begin_);
|
||||
s->Put(count_);
|
||||
|
||||
bool children = !is_leaf();
|
||||
s->Put(children);
|
||||
|
||||
if (children) {
|
||||
left_->Serialize(s);
|
||||
right_->Serialize(s);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Deserializes the tree from its structure, and re-calculates bottom-up
|
||||
* statistics.
|
||||
*/
|
||||
template<typename Deserializer>
|
||||
void Deserialize(const Dataset& data, Deserializer *s) {
|
||||
DEBUG_ASSERT(begin_ == BIG_BAD_NUMBER);
|
||||
|
||||
bound_.Deserialize(s);
|
||||
s->Get(&begin_);
|
||||
s->Get(&count_);
|
||||
|
||||
bool children;
|
||||
|
||||
s->Get(&children);
|
||||
|
||||
BinarySpaceTree *l = NULL;
|
||||
BinarySpaceTree *r = NULL;
|
||||
|
||||
if (children) {
|
||||
l = new BinarySpaceTree();
|
||||
l->Deserialize(data, s);
|
||||
r = new BinarySpaceTree();
|
||||
r->Deserialize(data, s);
|
||||
}
|
||||
|
||||
set_children(data, l, r);
|
||||
}
|
||||
|
||||
template<typename Serializer>
|
||||
void SerializeAll(const Dataset& data, Serializer *s) const {
|
||||
// can't use BinarySpaceTree as a magic number, because we want to be
|
||||
// able to deserialize this class with another statistic
|
||||
s->PutMagic(file::CreateMagic("spacetree")
|
||||
+ MAGIC_NUMBER(TDataset) + MAGIC_NUMBER(TBound));
|
||||
data.Serialize(s);
|
||||
Serialize(s);
|
||||
}
|
||||
|
||||
template<typename Deserializer>
|
||||
void DeserializeAll(Dataset* data, Deserializer *s) {
|
||||
// can't use BinarySpaceTree as a magic number, because we want to be
|
||||
// able to deserialize this class with another statistic
|
||||
s->AssertMagic(file::CreateMagic("spacetree")
|
||||
+ MAGIC_NUMBER(TDataset) + MAGIC_NUMBER(TBound));
|
||||
data->Deserialize(s);
|
||||
Deserialize(*data, s);
|
||||
}
|
||||
|
||||
// TODO: Not const correct
|
||||
|
||||
/**
|
||||
* Used only when constructing the tree.
|
||||
*/
|
||||
void set_children(const Dataset& data,
|
||||
BinarySpaceTree *left_in, BinarySpaceTree *right_in) {
|
||||
left_ = left_in;
|
||||
right_ = right_in;
|
||||
if (!is_leaf()) {
|
||||
stat_.Init(data, begin_, count_, left_->stat_, right_->stat_);
|
||||
DEBUG_ASSERT(count_ == left_->count_ + right_->count_);
|
||||
DEBUG_ASSERT(left_->begin_ == begin_);
|
||||
DEBUG_ASSERT(right_->begin_ == begin_ + left_->count_);
|
||||
} else {
|
||||
stat_.Init(data, begin_, count_);
|
||||
}
|
||||
}
|
||||
|
||||
const Bound& bound() const {
|
||||
return bound_;
|
||||
}
|
||||
|
||||
Bound& bound() {
|
||||
return bound_;
|
||||
}
|
||||
|
||||
const Statistic& stat() const {
|
||||
return stat_;
|
||||
}
|
||||
|
||||
Statistic& stat() {
|
||||
return stat_;
|
||||
}
|
||||
|
||||
bool is_leaf() const {
|
||||
return !left_;
|
||||
}
|
||||
|
||||
/**
|
||||
* Gets the left branch of the tree.
|
||||
*/
|
||||
BinarySpaceTree *left() const {
|
||||
// TODO: Const correctness
|
||||
return left_;
|
||||
}
|
||||
|
||||
/**
|
||||
* Gets the right branch.
|
||||
*/
|
||||
BinarySpaceTree *right() const {
|
||||
// TODO: Const correctness
|
||||
return right_;
|
||||
}
|
||||
|
||||
/**
|
||||
* Gets the index of the begin point of this subset.
|
||||
*/
|
||||
index_t begin() const {
|
||||
return begin_;
|
||||
}
|
||||
|
||||
/**
|
||||
* Gets the index one beyond the last index in the series.
|
||||
*/
|
||||
index_t end() const {
|
||||
return begin_ + count_;
|
||||
}
|
||||
|
||||
/**
|
||||
* Gets the number of points in this subset.
|
||||
*/
|
||||
index_t count() const {
|
||||
return count_;
|
||||
}
|
||||
|
||||
void Print() const {
|
||||
printf("node: %d to %d: %d points total\n",
|
||||
begin_, begin_ + count_ - 1, count_);
|
||||
if (!is_leaf()) {
|
||||
left_->Print();
|
||||
right_->Print();
|
||||
}
|
||||
}
|
||||
|
||||
FORBID_COPY(BinarySpaceTree);
|
||||
};
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,40 @@
|
||||
// Copyright 2007 Georgia Institute of Technology. All rights reserved.
|
||||
// ABSOLUTELY NOT FOR DISTRIBUTION
|
||||
/**
|
||||
* @file statistic.h
|
||||
*
|
||||
* Home for the concept of tree statistics.
|
||||
*
|
||||
* You should define your own statistic that looks like EmptyStatistic.
|
||||
*/
|
||||
|
||||
#ifndef TREE_STATISTIC_H
|
||||
#define TREE_STATISTIC_H
|
||||
|
||||
/**
|
||||
* Empty statistic if you are not interested in storing statistics in your
|
||||
* tree. Use this as a template for your own.
|
||||
*/
|
||||
template<class TDataset>
|
||||
class EmptyStatistic {
|
||||
public:
|
||||
EmptyStatistic() {}
|
||||
~EmptyStatistic() {}
|
||||
|
||||
/**
|
||||
* Initializes by taking statistics on raw data.
|
||||
*/
|
||||
void Init(const TDataset& dataset, index_t start, index_t count) {
|
||||
}
|
||||
|
||||
/**
|
||||
* Initializes by combining statistics of two partitions.
|
||||
*
|
||||
* This lets you build fast bottom-up statistics when building trees.
|
||||
*/
|
||||
void Init(const TDataset& dataset, index_t start, index_t count,
|
||||
const EmptyStatistic& left_stat, const EmptyStatistic& right_stat) {
|
||||
}
|
||||
};
|
||||
|
||||
#endif
|
||||
Reference in New Issue
Block a user