TKDE WORKS!

This commit is contained in:
Garry Boyer
2007-05-09 07:46:19 +00:00
parent 4cde7d16bd
commit 14850bbfbe
5 changed files with 57 additions and 24 deletions
+1 -1
View File
@@ -171,7 +171,7 @@ class Allnn {
void FinishReaccumulate(const Param& param,
const QNode& q_node) {
if (distance_sq_hi < 0) abort();
DEBUG_ASSERT_MSG(distance_sq_hi >= 0, "%f", distance_sq_hi);
/* no post-processing steps necessary */
}
+11 -4
View File
@@ -89,10 +89,10 @@ void DualTreeDepthFirst<GNP>::Init(datanode *datanode) {
do_naive_ = fx_param_bool(datanode, "do_naive", 0);
Matrix q_matrix;
data::Load(fx_param_str_req(datanode, "q"), &q_matrix);
ASSERT_PASS(data::Load(fx_param_str_req(datanode, "q"), &q_matrix));
Matrix r_matrix;
data::Load(fx_param_str_req(datanode, "r"), &r_matrix);
ASSERT_PASS(data::Load(fx_param_str_req(datanode, "r"), &r_matrix));
param_.Init(fx_submodule(datanode, "param", "param"),
q_matrix, r_matrix);
@@ -154,7 +154,9 @@ void DualTreeDepthFirst<GNP>::Begin() {
}
fx_timer_stop(datanode_, "execute");
//ot::Print(q_results_);
if (fx_param_bool(datanode_, "print", 0)) {
ot::Print(q_results_);
}
}
template<typename GNP>
@@ -166,7 +168,9 @@ void DualTreeDepthFirst<GNP>::PushDown_(index_t q_node_i) {
for (index_t q_i = q_node->begin(); q_i < q_node->end(); q_i++) {
typename GNP::QResult *q_result = &q_results_[q_i];
typename GNP::Point *q_point = &q_tree_.points()[q_i];
typename GNP::QPointInfo *q_info = &q_tree_.point_info()[q_i];
q_result->ApplyPostponed(param_, q_node_mut->postponed, *q_point);
q_result->Postprocess(param_, *q_point, *q_info, r_tree_.nodes()[0]);
}
} else {
for (index_t k = 0; k < 2; k++) {
@@ -229,6 +233,7 @@ void DualTreeDepthFirst<GNP>::Pair_(index_t q_node_i, index_t r_node_i,
}
// Phase 2: Explore children, and reincorporate their results.
q_node_mut->postponed.Reset(param_);
q_node_mut->mass_result.StartReaccumulate(param_, *q_node);
for (index_t k = 0; k < 2; k++) {
@@ -248,7 +253,6 @@ void DualTreeDepthFirst<GNP>::Pair_(index_t q_node_i, index_t r_node_i,
}
q_node_mut->mass_result.FinishReaccumulate(param_, *q_node);
q_node_mut->postponed.Reset(param_);
} else {
index_t r_child1_i = r_node->child(0);
index_t r_child2_i = r_node->child(1);
@@ -306,6 +310,8 @@ void DualTreeDepthFirst<GNP>::BaseCase_(
visitor.Init(param_);
q_node_mut->mass_result.StartReaccumulate(param_, *q_node);
DEBUG_ASSERT_MSG(q_node->count() != 0, "%d, count = %d, begin = %d",
q_node - qnode_(0), q_node->begin(), q_node->count());
for (index_t q_i = q_node->begin(); q_i < q_node->end(); ++q_i) {
typename GNP::Point *q_point = &q_tree_.points()[q_i];
@@ -331,6 +337,7 @@ void DualTreeDepthFirst<GNP>::BaseCase_(
q_node_mut->mass_result.Accumulate(param_, *q_result);
}
q_node_mut->mass_result.FinishReaccumulate(param_, *q_node);
q_node_mut->postponed.Reset(param_);
}
+2 -2
View File
@@ -173,7 +173,7 @@ void KdTreeMidpointBuilder<TPointInfo, TNode, TParam>::KdTreeMidpointBuilder::Bu
nodes_[node_i].begin(), nodes_[node_i].count(),
&nodes_[left_i].bound(),
&nodes_[right_i].bound());
DEBUG_MSG(3.0,"split (%d,[%d],%d) split_dim %d on %f (between %f, %f)",
nodes_[node_i].begin(), split_col,
nodes_[node_i].begin() + nodes_[node_i].count(), split_dim, split_val,
@@ -183,7 +183,7 @@ void KdTreeMidpointBuilder<TPointInfo, TNode, TParam>::KdTreeMidpointBuilder::Bu
nodes_[left_i].Init(nodes_[node_i].begin(),
split_col - nodes_[node_i].begin());
nodes_[right_i].Init(split_col,
nodes_[node_i].begin() + nodes_[node_i].count() - split_col);
nodes_[node_i].end() - split_col);
// This should never happen if max_width > 0
DEBUG_ASSERT(nodes_[left_i].count() != 0 && nodes_[right_i].count() != 0);
+1 -1
View File
@@ -268,7 +268,7 @@ class SpHrectBound {
double v = (v1 + fabs(v1)) + (v2 + fabs(v2));
sumsq += math::PowAbs<t_pow, 1>(v * v);
sumsq += math::PowAbs<t_pow, 1>(v);
}
return math::Pow<2, t_pow>(sumsq) / 4;
+42 -16
View File
@@ -47,8 +47,10 @@ class Tkde {
kernel.Init(fx_param_double_req(datanode, "h"));
double t = fx_param_double_req(datanode, "threshold");
t = t * kernel.CalcNormConstant(dim);
thresh.lo = t * (1.0 - 1.0e-7);
thresh.hi = t * (1.0 + 1.0e-7);
fx_format_result(datanode, "norm_constant", "%f",
kernel.CalcNormConstant(dim));
thresh.lo = t * (1.0 - 1.0e-4);
thresh.hi = t * (1.0 + 1.0e-4);
// WALDO: Fix me
}
@@ -81,11 +83,13 @@ class Tkde {
double ComputeKernelSum(
double distance_squared,
index_t r_count, const Vector& r_center, double r_sumsq) const {
//q*q - 2qr + rsumsq
//q*q - 2qr + r*r - r*r
double quadratic_term =
(distance_squared - la::Dot(r_center, r_center)) * r_count
+ r_sumsq;
return r_count - quadratic_term * kernel.inv_bandwidth_sq();
return -quadratic_term * kernel.inv_bandwidth_sq() + r_count;
}
};
@@ -273,10 +277,12 @@ class Tkde {
public:
void Init(const Param& param) {
moment_info.Init(param);
label = LAB_UNKNOWN;
}
void Reset(const Param& param) {
moment_info.Reset();
label = LAB_UNKNOWN;
}
void ApplyPostponed(const Param& param, const QPostponed& other) {
@@ -302,10 +308,6 @@ class Tkde {
void Init(const Param& param) {
d_density.Init(0, 0);
}
void ApplyDelta(const Param& param, const Delta& other) {
d_density += other.d_density;
}
};
// rho
@@ -330,13 +332,19 @@ class Tkde {
void Postprocess(const Param& param,
const Vector& q_point, const QPointInfo& q_info,
const RNode& r_root) {
/* nothing special to do */
if (density > param.thresh.hi) {
label |= LAB_HI;
} else if (density < param.thresh.lo) {
label |= LAB_LO;
}
DEBUG_ASSERT(label != LAB_CONFLICT);
}
void ApplyPostponed(const Param& param,
const QPostponed& postponed,
const Vector& q_point) {
label |= postponed.label; /* bitwise OR */
DEBUG_ASSERT_MSG(label >= 0 && label < 3, "%d", label);
if (!postponed.moment_info.is_empty()) {
density += postponed.moment_info.ComputeKernelSum(param, q_point);
@@ -372,7 +380,7 @@ class Tkde {
void Init(const Param& param) {
/* horizontal init */
density.Init(0, 0);
label = 0;
label = LAB_UNKNOWN;
}
void StartReaccumulate(const Param& param, const QNode& q_node) {
@@ -386,12 +394,14 @@ class Tkde {
// but in some cases may require a copy/undo stage
density |= result.density;
label &= result.label;
DEBUG_ASSERT(result.label != LAB_CONFLICT);
}
void Accumulate(const Param& param,
const QMassResult& result, index_t n_points) {
density |= result.density;
density &= result.label;
label &= result.label;
DEBUG_ASSERT(result.label != LAB_CONFLICT);
}
void FinishReaccumulate(const Param& param,
@@ -403,8 +413,9 @@ class Tkde {
void ApplyMassResult(const Param& param,
const QMassResult& mass_result) {
density += mass_result.density;
DEBUG_ASSERT_MSG((label | mass_result.label) != LAB_CONFLICT,
"%d and %d", label, mass_result.label);
label |= mass_result.label;
DEBUG_ASSERT(label != LAB_CONFLICT);
}
void ApplyDelta(const Param& param,
@@ -417,6 +428,7 @@ class Tkde {
bool change_made;
if (unlikely(postponed.label)) {
DEBUG_ASSERT((label | postponed.label) != LAB_CONFLICT);
label = postponed.label;
change_made = true;
} else if (unlikely(!postponed.moment_info.is_empty())) {
@@ -520,24 +532,38 @@ class Tkde {
double distance_sq_lo =
q_node.bound().MinDistanceSqToBound(r_node.bound());
bool need_expansion;
//printf("%f %f %f\n",
// q_node.bound().MinDistanceSqToBound(r_node.bound()),
// q_node.bound().MidDistanceSqToBound(r_node.bound()),
// q_node.bound().MaxDistanceSqToBound(r_node.bound())
// );
if (distance_sq_lo >= param.kernel.bandwidth_sq()) {
if (distance_sq_lo > param.kernel.bandwidth_sq()) {
need_expansion = false;
} else {
double distance_sq_hi =
q_node.bound().MaxDistanceSqToBound(r_node.bound());
if (distance_sq_hi <= param.kernel.bandwidth_sq()) {
if (distance_sq_hi < param.kernel.bandwidth_sq()) {
q_postponed->moment_info.Add(r_node.stat().moment_info);
need_expansion = false;
} else {
delta->d_density = r_node.stat().moment_info.ComputeKernelSumRange(
param, q_node.bound());
#ifdef BIGPRUNE
#endif
//delta->d_density = r_node.stat().moment_info.ComputeKernelSumRange(
// param, q_node.bound());
// we computed the lower bound of the quadratic. if it is positive
// it means we have a better-than-nothing bound; if it is not, then
// we can resort to saying the min contribution is zero.
delta->d_density.lo = max(delta->d_density.lo, 0.0);
// - problem: the upper bound is no good.
//max(delta->d_density.lo, 0.0);
delta->d_density.lo = 0;
delta->d_density.hi = r_node.count() *
param.kernel.EvalUnnormOnSq(distance_sq_lo);
need_expansion = true;
#ifdef BIGPRUNE
#endif
}
}