This commit is contained in:
Garry Boyer
2007-03-14 17:18:28 +00:00
parent 6726d3ce4e
commit 24ee655e68
15 changed files with 2129 additions and 0 deletions
+353
View File
@@ -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
+10
View File
@@ -0,0 +1,10 @@
librule(
sources = [],
headers = [
"kdtree.h", "bounds.h", "spacetree.h", "statistic.h"
],
deplibs = ["base:base", "la:la", "col:col",
"file:file"]
)
+253
View 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
+279
View File
@@ -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
+40
View File
@@ -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