From f673bbb8a22a1d84e7d8f5e6ef5533e37cff8c9d Mon Sep 17 00:00:00 2001 From: nslagle Date: Thu, 3 Nov 2011 17:56:35 +0000 Subject: [PATCH] mlpack/contrib/nslagle: add the multi-tree base function --- src/contrib/nslagle/myKDE/kde_dual_tree.hpp | 5 +- .../nslagle/myKDE/kde_dual_tree_impl.hpp | 66 ++++++++++++++++++- 2 files changed, 67 insertions(+), 4 deletions(-) diff --git a/src/contrib/nslagle/myKDE/kde_dual_tree.hpp b/src/contrib/nslagle/myKDE/kde_dual_tree.hpp index 668d3d518a..74383afcbe 100644 --- a/src/contrib/nslagle/myKDE/kde_dual_tree.hpp +++ b/src/contrib/nslagle/myKDE/kde_dual_tree.hpp @@ -46,6 +46,7 @@ template remainingBandwidths); double GetPriority(TTree* nodeQ, TTree* nodeT) { diff --git a/src/contrib/nslagle/myKDE/kde_dual_tree_impl.hpp b/src/contrib/nslagle/myKDE/kde_dual_tree_impl.hpp index 33fbacdbea..62ef5e9132 100644 --- a/src/contrib/nslagle/myKDE/kde_dual_tree_impl.hpp +++ b/src/contrib/nslagle/myKDE/kde_dual_tree_impl.hpp @@ -48,22 +48,84 @@ KdeDualTree::SetDefaults() BandwidthRange(0.01, 100.0); bandwidthCount = 10; delta = epsilon = 0.05; + kernel = TKernel(1.0); } template void KdeDualTree::MultiBandwidthDualTreeBase(TTree* Q, - TTree* T, + TTree* T, size_t QIndex, std::set remainingBandwidths) { + size_t sizeOfTNode = T->count(); + size_t sizeOfQNode = Q->count(); for (size_t q = Q->begin(); q < Q->end(); ++q) { + arma::vec queryPoint = queryData.unsafe_col(q); for (size_t t = T->begin(); t < T->end(); ++t) { + arma::vec diff = queryPoint - referenceData.unsafe_col(t); + double distSquared = arma::dot(diff, diff); + std::set::iterator bIt = remainingBandwidths.end(); + size_t bandwidthIndex = bandwidthCount; + while (bIt != remainingBandwidths.begin()) + { + --bIt; + --bandwidthIndex; + double bandwidth = *bIt; + double scaledProduct = distSquared / (bandwidth * bandwidth); + /* TODO: determine the power of the incoming argument */ + double contribution = kernel(scaledProduct); + if (contribution > DBL_EPSILON) + { + upperBoundQPointByBandwidth(q, bandwidthIndex) += contribution; + lowerBoundQPointByBandwidth(q, bandwidthIndex) += contribution; + } + else + { + break; + } + } + } + for (size_t bIndex = bandwidthCount - remainingBandwidths.size(); bIndex < remainingBandwidths.size(); ++bIndex) + { + upperBoundQPointByBandwidth(q, bIndex) -= sizeOfTNode; } } + size_t levelOfQ = GetLevelOfNode(Q); + for (size_t bIndex = bandwidthCount - remainingBandwidths.size(); bIndex < remainingBandwidths.size(); ++bIndex) + { + /* subtract out the current log-likelihood amount for this Q node so we can readjust + * the Q node bounds by current bandwidth */ + upperBoundLevelByBandwidth(levelOfQ, bIndex) -= + sizeOfQNode * log(upperBoundQNodeByBandwidth(QIndex, bIndex)); + lowerBoundLevelByBandwidth(levelOfQ, bIndex) -= + sizeOfQNode * log(lowerBoundQNodeByBandwidth(QIndex, bIndex)); + arma::vec upperBound = upperBoundQPointByBandwidth.unsafe_col(bIndex); + arma::vec lowerBound = lowerBoundQPointByBandwidth.unsafe_col(bIndex); + double minimumLower = lowerBoundQPointByBandwidth(Q->begin(), bIndex); + double maximumUpper = upperBoundQPointByBandwidth(Q->begin(), bIndex); + for (size_t q = Q->begin(); q < Q->end(); ++q) + { + if (lowerBoundQPointByBandwidth(q,bIndex) < minimumLower) + { + minimumLower = lowerBoundQPointByBandwidth(q,bIndex); + } + if (upperBoundQPointByBandwidth(q,bIndex) > maximumUpper) + { + maximumUpper = upperBoundQPointByBandwidth(q,bIndex); + } + } + /* adjust Q node bounds, then add the new quantities to the level by bandwidth + * log-likelihood bounds */ + lowerBoundQNodeByBandwidth(QIndex, bIndex) = minimumLower; + upperBoundQNodeByBandwidth(QIndex, bIndex) = maximumUpper - sizeOfTNode; + upperBoundLevelByBandwidth(levelOfQ, bIndex) += + sizeOfQNode * log(upperBoundQNodeByBandwidth(QIndex, bIndex)); + lowerBoundLevelByBandwidth(levelOfQ, bIndex) += + sizeOfQNode * log(lowerBoundQNodeByBandwidth(QIndex, bIndex)); + } } - }; };