You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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
相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.12 08:45:35