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

Java实现KD树的M近邻查找方法问题求助

KD树M近邻查找实现(Java)

核心思路

你的最近邻算法是维护单个最优节点和最小距离,M近邻则需要维护一个容量为m的最大堆——堆顶始终是当前已找到的m个邻居中距离最远的那个。这样做的好处是:

  • 堆未满时,直接加入新节点;
  • 堆已满时,只要新节点距离比堆顶小,就替换堆顶并调整堆;
  • 遍历分支的条件从「当前维度差平方 ≥ 最小距离」改为「当前维度差平方 < 堆顶距离」(堆未满时直接检查另一侧),确保不会漏掉可能的更近节点。

因为你不想用现成类库,下面先实现一个极简的最大堆,再基于你已有的最近邻逻辑修改M近邻递归方法。


第一步:实现自定义最大堆

用来存储已找到的近邻节点,堆顶为距离最远的节点:

class MaxDistanceHeap {
    private static class HeapNode {
        KDNode node;
        float distanceSq; // 用距离平方避免开根号,提升效率

        HeapNode(KDNode node, float distanceSq) {
            this.node = node;
            this.distanceSq = distanceSq;
        }
    }

    private HeapNode[] heap;
    private int size;
    private int maxSize;

    public MaxDistanceHeap(int m) {
        maxSize = m;
        heap = new HeapNode[m];
        size = 0;
    }

    // 获取堆顶的距离平方(当前最远的距离)
    public float getTopDistanceSq() {
        return size == 0 ? Float.MAX_VALUE : heap[0].distanceSq;
    }

    // 插入节点,维护最大堆性质
    public void insert(KDNode node, float distanceSq) {
        if (size < maxSize) {
            // 堆未满,直接加入并上浮
            heap[size] = new HeapNode(node, distanceSq);
            bubbleUp(size);
            size++;
        } else {
            // 堆已满,只有当新节点距离更近时才替换堆顶
            if (distanceSq < heap[0].distanceSq) {
                heap[0] = new HeapNode(node, distanceSq);
                bubbleDown(0);
            }
        }
    }

    // 上浮调整堆
    private void bubbleUp(int idx) {
        while (idx > 0) {
            int parentIdx = (idx - 1) / 2;
            if (heap[idx].distanceSq > heap[parentIdx].distanceSq) {
                swap(idx, parentIdx);
                idx = parentIdx;
            } else {
                break;
            }
        }
    }

    // 下沉调整堆
    private void bubbleDown(int idx) {
        while (true) {
            int leftChild = 2 * idx + 1;
            int rightChild = 2 * idx + 2;
            int largest = idx;

            if (leftChild < size && heap[leftChild].distanceSq > heap[largest].distanceSq) {
                largest = leftChild;
            }
            if (rightChild < size && heap[rightChild].distanceSq > heap[largest].distanceSq) {
                largest = rightChild;
            }

            if (largest != idx) {
                swap(idx, largest);
                idx = largest;
            } else {
                break;
            }
        }
    }

    private void swap(int i, int j) {
        HeapNode temp = heap[i];
        heap[i] = heap[j];
        heap[j] = temp;
    }

    // 获取所有堆中的节点坐标
    public float[][] getResult() {
        float[][] result = new float[size][];
        for (int i = 0; i < size; i++) {
            result[i] = heap[i].node.getCoordinates();
        }
        return result;
    }
}

第二步:实现M近邻递归方法

基于你已有的nearest方法修改,核心是用堆替换单个最优节点:

// 类内变量(和你原来的visited类似)
private int visited;

private void findMNearestRecursive(KDNode root, KDNode target, int index, MaxDistanceHeap heap, int k) {
    if (root == null) {
        return;
    }
    visited++;

    // 计算当前节点与目标的距离平方(避免开根号)
    float distanceSq = root.distanceSq(target); // 建议新增这个方法,计算坐标差的平方和

    // 尝试将当前节点加入堆
    heap.insert(root, distanceSq);

    // 如果已经找到m个距离为0的节点,直接返回(无需继续搜索)
    if (heap.getTopDistanceSq() == 0) {
        return;
    }

    // 确定优先遍历的子树(和最近邻逻辑一致)
    float dx = root.getCoordinates()[index] - target.getCoordinates()[index];
    int nextIndex = (index + 1) % k;
    findMNearestRecursive(dx > 0 ? root.getLeft() : root.getRight(), target, nextIndex, heap, k);

    // 检查是否需要遍历另一侧子树:如果当前维度差的平方 < 堆顶的距离平方,说明另一侧可能有更近的节点
    if (dx * dx < heap.getTopDistanceSq()) {
        findMNearestRecursive(dx > 0 ? root.getRight() : root.getLeft(), target, nextIndex, heap, k);
    }
}

注:需要给你的KDNode新增一个distanceSq方法,计算欧氏距离的平方,避免开根号的开销:

public float distanceSq(KDNode target) {
    float sum = 0;
    float[] coords = getCoordinates();
    float[] targetCoords = target.getCoordinates();
    for (int i = 0; i < coords.length; i++) {
        float diff = coords[i] - targetCoords[i];
        sum += diff * diff;
    }
    return sum;
}

第三步:封装成要求的findMNearest方法

public float[][] findMNearest(float[] point, int m) {
    if (m <= 0 || point == null) {
        return new float[0][];
    }
    // 将目标点转为KDNode(复用你已有的节点结构)
    KDNode target = new KDNode(point);
    MaxDistanceHeap heap = new MaxDistanceHeap(m);
    visited = 0; // 重置访问计数
    // 假设k是你的KD树维度,需为类内变量或传入
    findMNearestRecursive(root, target, 0, heap, k);
    return heap.getResult();
}

关键差异说明

和你原来的最近邻方法对比,核心变化:

  1. 用最大堆替代单个best节点和bestDistance,同时跟踪m个近邻;
  2. 遍历另一侧分支的条件从「dx*dx >= bestDistance」改为「dx*dx < 堆顶距离平方」——因为堆顶是当前最远的近邻,只要分割线到目标点的距离比它小,另一侧就可能存在更近的节点;
  3. 用距离平方替代实际距离,减少计算开销,且不影响大小比较。

内容的提问来源于stack exchange,提问作者Iman barca

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 04:05:19