KDTree查询点远离数据集时性能劣于暴力搜索的原因排查
KDTree查询远离数据集点时性能劣于暴力搜索的原因分析
问题背景
实现了一个基础3D KDTree,当查询点远离数据集时(例如数据集分布在[-10;10]^3,查询点为(100,100,100)),暴力搜索比KDTree快20%,但正常范围内查询时KDTree性能远优于暴力搜索。以下是完整代码与测试结果:
KDTree.h
#pragma once #include <memory> #include <vector> #include <cassert> #include <vector> #include <stdexcept> #include <cmath> #include <stack> /** * @brief K-d tree of 3D float points. * * @see Variable names, and algorithm inspired by https://youtu.be/TrqK-atFfWY?t=2567 */ class KDTree { public: /**** Utility structures ****/ struct Point { float x, y, z; float& operator[](int i) { assert(i < 3 && i >= 0); switch(i) { case 0: return x; case 1: return y; case 2: return z; } throw std::runtime_error("Invalid index"); } const float& operator[](int i) const { return const_cast<Point&>(*this)[i]; } float distance(const Point& other) const { return sqrtf(distanceSquared(other)); } float distanceSquared(const Point& other) const { const float dx = (x - other.x); const float dy = (y - other.y); const float dz = (z - other.z); return dx*dx + dy*dy + dz*dz; } template<typename T> friend T& operator<<(T& lhs, const Point& rhs) { lhs << "(" << rhs.x << ", " << rhs.y << ", " << rhs.z << ")"; return lhs; } }; struct AABB { Point min, max; bool inside(const Point& p) const { return p.x >= min.x && p.y >= min.y && p.z >= min.z && p.x < max.x && p.y < max.y && p.z < max.z; } }; public: /**** Utility (static) functions ****/ /** * @brief Compute the bounding box of a set of points. * @param points * The points to make the boudning box from. * If the set of points is empty, the returned value is undefined. * @return * The bounding box of theses points. * The bounding box will be always of minimum size, i.e. there will be a vertex on each of the 6 faces of the * returned AABB. */ static AABB computeBoundingBox(const std::vector<Point> &points); static float median(std::vector<float> vec); static void store_min(float& current, float newValue) { if(current > newValue) { current = newValue; } } static void store_max(float& current, float newValue) { if(current < newValue) { current = newValue; } } public: KDTree() = default; explicit KDTree(const std::vector<Point>& points); Point computeNearestNeighbor(const Point& pos) const; private: /** * @brief Represent each node of the k-d tree. */ struct Node { /** * @brief * Children. Both of them or neither of them are null. */ std::unique_ptr<Node> left, right; /** * The distance between the origin and the wall to split (for left node), * or the distance from the wall and the end to split (for right node). */ float splitDistance; int splitDim; std::vector<Point> points; /** * @return true if this node is a leaf node (it has no children). */ bool leaf() const { // left or right it doesn't matter return left == nullptr; } }; struct SplitStack { int dim; std::vector<Point> points; Node* node; AABB aabb; }; private: /** * @brief Split the tree during the build. * @param dim The dimension to split. * @param points The list of remaining candidate points for this area, inside the AABB. * @param node The node, allocated, to fill. * @param aabb The bounding box of the node. */ void split(std::stack<SplitStack>& stack); void searchRecursive(const Point& pos, Node* node, float& currentDist, Point& currentNeighbor) const; private: AABB m_rootAABB; std::unique_ptr<Node> m_root; };
KDTree.cpp
#include "KDTree.h" #include <climits> #include <cfloat> #include <algorithm> #include <iostream> #include <cmath> float KDTree::median(std::vector<float> vec) { size_t size = vec.size(); if (size == 0) { return 0; // Undefined, really. } else { std::sort(vec.begin(), vec.end()); if (size % 2 == 0) { return (vec[size / 2 - 1] + vec[size / 2]) / 2; } else { return vec[size / 2]; } } } KDTree::AABB KDTree::computeBoundingBox(const std::vector<Point>& points) { AABB res; if (!points.empty()) { const auto& firstPoint = points.front(); // Initialize the bounding box to a point for (int dim = 0; dim < 3; dim++) { res.min[dim] = firstPoint[dim]; res.max[dim] = firstPoint[dim]; } // Grow the bounding box for each point if needed for (const Point& point: points) { for (int dim = 0; dim < 3; dim++) { store_min(res.min[dim], point[dim]); store_max(res.max[dim], point[dim]); } } } return res; } KDTree::KDTree(const std::vector<Point>& points) { m_rootAABB = computeBoundingBox(points); m_root = std::make_unique<Node>(); std::stack<SplitStack> stack; stack.push(SplitStack{0, points, m_root.get(), m_rootAABB}); split(stack); } void KDTree::split(std::stack<SplitStack>& stack) { // Split recursively in x, y, z, x, y, z... // Split at the center // dim axis --> // 0 --------- aabb[dim].min --------------------------- aabb[dim].max --------- +inf // ------------------|--------------------|-------------------|------------------- // ---------------------- left node ----------- right node ----------------------- // <-----------------> <--------------- // splitDistance: if left if right while (!stack.empty()) { std::vector<Point> points = std::move(stack.top().points); Node& node = *stack.top().node; int dim = stack.top().dim; AABB aabb = stack.top().aabb; stack.pop(); // Stop condition if (points.size() > 100) { node.splitDim = dim; // Absolute position in the dimension of the split node.splitDistance = (aabb.max[dim] + aabb.min[dim]) / 2.0f; AABB leftAABB = aabb; leftAABB.max[dim] = node.splitDistance; AABB rightAABB = aabb; rightAABB.min[dim] = leftAABB.max[dim]; std::vector<Point> leftPoints, rightPoints; for (const Point& p: points) { if (leftAABB.inside(p)) { leftPoints.push_back(p); } else { rightPoints.push_back(p); } } const int nextDim = (dim + 1) % 3; node.right = std::make_unique<Node>(); stack.push(SplitStack{nextDim, std::move(rightPoints), node.right.get(), rightAABB}); node.left = std::make_unique<Node>(); stack.push(SplitStack{nextDim, std::move(leftPoints), node.left.get(), leftAABB}); } else { // Leaf node.points = std::move(points); } } } KDTree::Point KDTree::computeNearestNeighbor(const KDTree::Point& pos) const { // Are we left or right? const Node *node = m_root.get(); AABB aabb = m_rootAABB; float dist = FLT_MAX; Point res; searchRecursive(pos, m_root.get(), dist, res); return res; } void KDTree::searchRecursive(const Point& pos, Node *node, float& currentDist, Point& currentNeighbor) const { // Are we on a leaf? if (node->leaf()) { // We are on a leaf // Search brute force into the leaf node for (const auto& other: node->points) { const float d = other.distanceSquared(pos); if (d < currentDist) { currentDist = d; currentNeighbor = other; } } } else { Node *front, *back; // Are we on the left side? if (pos[node->splitDim] < node->splitDistance) { // Pos is on the left side front = node->left.get(); back = node->right.get(); } else { // Pos is on the right side front = node->right.get(); back = node->left.get(); } searchRecursive(pos, front, currentDist, currentNeighbor); // If the current closest point is closer than the closest point of the back face, no need to search in the back // face because it will be always further. // If not, we save half of the time for the current node const float backDist = fabsf(node->splitDistance - pos[node->splitDim]); // Do not forget all distances all squared if (backDist * backDist <= currentDist) { // If it can be closer, search also in this node searchRecursive(pos, back, currentDist, currentNeighbor); } } }
main.cpp
#include <iostream> #include <vector> #include <random> #include "KDTree.h" #include "viewer.h" #include <chrono> class Timer { public: Timer(const std::string& title) : title_(title), beg_(clock_::now()) {} ~Timer() { std::cout << title_ << " elapsed: " << elapsed() << "s" << std::endl; } void reset() { beg_ = clock_::now(); } double elapsed() const { return std::chrono::duration_cast<second_> (clock_::now() - beg_).count(); } private: std::string title_; typedef std::chrono::high_resolution_clock clock_; typedef std::chrono::duration<double, std::ratio<1>> second_; std::chrono::time_point<clock_> beg_; }; std::vector<KDTree::Point> randomPoints(int size, float bounds = 10.0f) { std::vector<KDTree::Point> points; std::uniform_real_distribution<float> dist(-bounds, bounds); std::mt19937 engine; for(int i = 0; i < size; i++) { KDTree::Point point; for(int dim = 0; dim < 3; dim++) { point[dim] = dist(engine); } points.push_back(point); } return points; } int main() { auto points = randomPoints(1'000'000); for(int i = 0; i < points.size(); i++) { auto& point = points[i]; } auto aabb = KDTree::computeBoundingBox(points); KDTree kdtree; { Timer timer("KDTree build"); kdtree = KDTree(points); } { const int N = 1'000; auto testPts = randomPoints(N, 100.0f); std::vector<KDTree::Point> resKD(N), resBrute(N); { Timer timer("KDTree"); for(int i = 0; i < testPts.size(); i++) { resKD[i] = kdtree.computeNearestNeighbor(testPts[i]); } } { Timer timer("Bruteforce"); for(int i = 0; i < testPts.size(); i++) { const auto& test = testPts[i]; // Brute force KDTree::Point cur; float curDist = FLT_MAX; for(const auto& brute : points) { if(brute.distanceSquared(test) < curDist) { curDist = brute.distanceSquared(test); cur = brute; } } resBrute[i] = cur; } } { float delta = 0.0f; for(int i = 0; i < N; i++) { delta += resKD[i].x - resBrute[i].x; delta += resKD[i].y - resBrute[i].y; delta += resKD[i].z - resBrute[i].z; } std::cout << "delta = " << delta << std::endl; } } return 0; }
测试输出
查询点范围为100.0f时:
KDTree build elapsed: 0.190593s KDTree elapsed: 2.69598s Bruteforce elapsed: 2.34136s delta = 0
查询点范围改为10.0f时:
KDTree build elapsed: 0.195519s KDTree elapsed: 0.000914431s Bruteforce elapsed: 2.35679s delta = 0
核心原因分析
1. 分裂策略导致树结构失衡
当前KDTree采用固定维度循环+AABB中心分裂的方式,而非基于数据分布的中位数分裂。这种分裂方式会导致:
- 树的层级过多,递归遍历的函数调用开销巨大
- 左右子树的点数可能严重失衡,甚至某一子树包含绝大多数数据,使得KDTree退化为线性遍历,同时还额外增加了递归开销
2. 回溯剪枝逻辑完全失效
在searchRecursive中,回溯判断的条件是backDist * backDist <= currentDist。当查询点远离数据集时,初始的currentDist被设为FLT_MAX,这会导致所有节点的回溯条件都成立——KDTree会遍历整个树的所有叶子节点,相当于做了一次带递归开销的暴力搜索,自然比纯暴力遍历慢。
3. 叶子节点阈值设置不合理
当前设置叶子节点最多包含100个点,对于1e6规模的数据集,会生成大量叶子节点。当查询点远离数据集时,每个叶子节点的暴力遍历加上递归调用的开销,累积起来超过了直接遍历整个数据集的开销。
优化方案
1. 改用数据中位数分裂策略
替换AABB中心分裂为当前维度下数据的中位数分裂,确保每次分裂后左右子树点数大致均衡,减少树的层级,同时让回溯剪枝能有效工作。示例修改:
// 在split函数中替换splitDistance计算与点划分逻辑 std::vector<float> dimValues; dimValues.reserve(points.size()); for (const auto& p : points) { dimValues.push_back(p[dim]); } node.splitDistance = median(dimValues); std::vector<Point> leftPoints, rightPoints; for (const Point& p : points) { if (p[dim] < node.splitDistance) { leftPoints.push
相关产品推荐
相关产品推荐

