Day 026-Q4 — k-d Tree (最近傍 + 矩形範囲)

2026-05-09 赤色 Master / Phase 8+ ★★★★★★★★★ k-d Tree / KNN

問題

2 次元点集合に対して KNN (k-Nearest Neighbor) と矩形範囲カウントを処理。静的 k-d Tree。

制約

$1 \le N \le 2 \times 10^5$
$1 \le Q \le 10^5$
$k \le \min(N, 10)$

入出力例

入力例 1

5
1 2
3 4
5 1
2 5
4 3
3
KNN 3 3 2
RANGE 1 1 4 4
KNN 0 0 1

出力例 1

(4,3) (3,4)
4
(1,2)

ヒント (段階的開示)

ヒント1: 方向性
軸交互に中央値で分割し、各ノードに包含矩形を持たせる。
ヒント2: アプローチ
KNN は最大ヒープで $k$ 個管理、bbox との最小距離で枝刈り。RANGE も bbox 交差判定で枝刈り。
ヒント3: 誘導
bbox-point 距離: $\max(0, lo_x - x, x - hi_x)^2 + \max(0, lo_y - y, y - hi_y)^2$。

模範解答 (Python)

import sys
import heapq
input = sys.stdin.readline

class KDTree:
    def __init__(self, points):
        self.nodes = []
        self.root = self._build(list(points), 0) if points else -1
    def _build(self, pts, depth):
        if not pts: return -1
        axis = depth % 2
        pts.sort(key=lambda p: p[axis])
        mid = len(pts) // 2
        p = pts[mid]
        xs = [q[0] for q in pts]; ys = [q[1] for q in pts]
        bbox = (min(xs), max(xs), min(ys), max(ys))
        idx = len(self.nodes)
        self.nodes.append(None)
        left = self._build(pts[:mid], depth + 1)
        right = self._build(pts[mid+1:], depth + 1)
        self.nodes[idx] = (p, axis, left, right, bbox)
        return idx
    def _dist2(self, p, q):
        return (p[0]-q[0])**2 + (p[1]-q[1])**2
    def _bbox_dist2(self, point, bbox):
        x, y = point
        min_x, max_x, min_y, max_y = bbox
        dx = max(0, min_x - x, x - max_x)
        dy = max(0, min_y - y, y - max_y)
        return dx*dx + dy*dy
    def knn(self, query, k):
        heap = []
        def search(idx):
            if idx == -1: return
            p, axis, left, right, bbox = self.nodes[idx]
            d2 = self._dist2(query, p)
            if len(heap) < k: heapq.heappush(heap, (-d2, p))
            elif d2 < -heap[0][0]: heapq.heapreplace(heap, (-d2, p))
            diff = query[axis] - p[axis]
            near, far = (left, right) if diff <= 0 else (right, left)
            search(near)
            if far != -1:
                far_bbox = self.nodes[far][4]
                if len(heap) < k or self._bbox_dist2(query, far_bbox) < -heap[0][0]:
                    search(far)
        search(self.root)
        result = sorted((-nd, p) for nd, p in heap)
        return [p for _, p in result]
    def range_count(self, x1, y1, x2, y2):
        count = [0]
        def search(idx):
            if idx == -1: return
            p, axis, left, right, bbox = self.nodes[idx]
            min_x, max_x, min_y, max_y = bbox
            if max_x < x1 or x2 < min_x or max_y < y1 or y2 < min_y: return
            if x1 <= p[0] <= x2 and y1 <= p[1] <= y2:
                count[0] += 1
            search(left); search(right)
        search(self.root)
        return count[0]

def solve():
    N = int(input())
    points = []
    for _ in range(N):
        x, y = map(int, input().split())
        points.append((x, y))
    kdt = KDTree(points)
    Q = int(input())
    out = []
    for _ in range(Q):
        line = input().split()
        if line[0] == "KNN":
            qx, qy, k = int(line[1]), int(line[2]), int(line[3])
            result = kdt.knn((qx, qy), k)
            out.append(' '.join(f'({p[0]},{p[1]})' for p in result))
        else:
            x1, y1, x2, y2 = int(line[1]), int(line[2]), int(line[3]), int(line[4])
            out.append(str(kdt.range_count(x1, y1, x2, y2)))
    print('\n'.join(out))

solve()

Step-by-Step 解説

1構築
軸交互 + 中央値分割。各ノードに bbox。
2KNN
最大ヒープで $k$ 個管理。bbox 距離で枝刈り。
3RANGE
bbox 交差判定 → 範囲内なら点をカウント。期待 $O(\sqrt N)$。
4bbox 距離
各軸独立に $\max(0, lo - x, x - hi)$ を計算。

よくあるミス

ミス原因正しい書き方
近い子の選択ミスdiff 符号ミスdiff <= 0 → left
bbox 再計算毎回計算構築時に伝播
ヒープが最小ヒープのまま負化忘れ(-d2, p)

次のステップ

  • 動的 k-d Tree(追加・削除)
  • 高次元近似最近傍 (LSH)

自己評価

自分の回答

気づき・メモ