Day 079-Q3 — 動的平面最近傍探索(KD-Tree + √N 再構築)

2026-07-01 赤色 Master / Phase 8+ ★★★★★★★★★ KD-Tree・最近傍探索・平方根分解・動的点追加

問題

2次元平面上の点集合に対し、以下の2種類のクエリを処理せよ。

  • add x y: 点 $(x, y)$ を集合に追加する
  • query qx qy: 現在の集合の中から点 $(q_x, q_y)$ に最もユークリッド距離が近い点を答えよ

制約

パラメータ範囲備考
$N$(追加数)$1 \le N \le 2 \times 10^5$add 回数
$Q$(クエリ数)$1 \le Q \le 2 \times 10^5$query 回数
座標$0 \le x, y, qx, qy \le 10^9$整数

入出力例

入力例1

8 5
add 3 7
add 1 4
add 8 2
add 5 5
query 4 6
add 7 1
add 2 9
add 6 4
query 6 5
query 0 0
add 9 8
query 5 5
query 8 8

出力例1

3 7
6 4
1 4
5 5
9 8

クエリ 4 6 時点では {(3,7),(1,4),(8,2),(5,5)} の4点が存在。 (3,7) と (5,5) が距離² = 2 で同距離。最初に見つかった (3,7) を返す。

概念図: KD-Tree + √N バッファ平方根分解

動的最近傍探索: KD-Tree + √N バッファ(平方根分解) 2次元平面 x y (3,7) (1,4) (8,2) (5,5) (7,1) (2,9) (6,4) Q(4,6) nearest KD-split x=5 KD-Tree に組み込まれた点 バッファの点(未構築) √N 平方根分解の仕組み add(x, y) → buffer.append((x,y)) O(1) len(buffer) ≥ B=√N → KDTree 再構築 query(qx, qy) ① KDTree.nearest(q) O(√N) 期待 ② buffer 線形スキャン O(B)=O(√N) → min(①, ②) を返す 計算量 add (バッファ追加): O(1) add (再構築): O(N log N) / √N 回 → 合計 O(N^{3/2}) query: O(√N) 期待

青い点は KD-Tree に組み込まれた点、オレンジの点は未組み込みバッファ点。 クエリ時は両方を探索し最近傍を返す。バッファが $B = O(\sqrt{N})$ に達したら全点で KD-Tree を再構築。

ヒント

ヒント1(方向性)

静的KD-Treeは構築 $O(N \log N)$、クエリ $O(\sqrt{N})$ 期待値。動的追加をサポートするには「Sqrt Decomposition」を使う。バッファサイズ $B = \sqrt{N}$ 点を超えたら全点で KD-Tree を再構築する。

ヒント2(アプローチ)

新しい点はまずバッファ(リスト)に追加。バッファが $B$ 点に達したら全点で KD-Tree を再構築。クエリ時: (1) バッファの全点を線形スキャン $O(B)$、(2) 静的 KD-Tree を探索 $O(\sqrt{N})$、両者の最小を返す。

ヒント3(ほぼ答え)

Sqrt Decomposition の骨格:

B = isqrt(N) + 1  # バッファサイズ
all_pts, buffer = [], []
kdt = KDTree([])

# add 操作
buffer.append(pt); all_pts.append(pt)
if len(buffer) >= B:
    kdt = KDTree(all_pts)  # 全点で再構築
    buffer.clear()

# query 操作
best_p, best_d = kdt.nearest(query)   # KD-Tree
for pt in buffer:                      # バッファ線形スキャン
    d = dist2(query, pt)
    if d < best_d: best_d = d; best_p = pt

模範解答

import sys
from math import isqrt

def dist2(a, b):
    return (a[0]-b[0])**2 + (a[1]-b[1])**2

class KDTree:
    def __init__(self, points):
        self.nodes = []
        if points:
            self._build(list(points), 0)

    def _build(self, pts, depth):
        if not pts: return -1
        axis = depth % 2
        pts.sort(key=lambda p: p[axis])
        mid = len(pts) // 2
        idx = len(self.nodes)
        self.nodes.append(None)
        left_idx  = self._build(pts[:mid],   depth + 1)
        right_idx = self._build(pts[mid+1:], depth + 1)
        self.nodes[idx] = (pts[mid], left_idx, right_idx, axis)
        return idx

    def nearest(self, query):
        if not self.nodes:
            return None, float('inf')
        best_d = [float('inf')]
        best_p = [None]

        def _search(node_idx):
            if node_idx == -1: return
            point, left, right, axis = self.nodes[node_idx]
            d = dist2(query, point)
            if d < best_d[0]:
                best_d[0] = d; best_p[0] = point
            diff = query[axis] - point[axis]
            near, far = (left, right) if diff <= 0 else (right, left)
            _search(near)
            if diff * diff < best_d[0]:
                _search(far)

        _search(0)
        return best_p[0], best_d[0]

def main():
    data = sys.stdin.buffer.read().split()
    idx = 0
    N = int(data[idx]); idx += 1
    Q = int(data[idx]); idx += 1

    B = max(1, isqrt(N) + 1)
    all_points = []
    buffer = []
    kdt = KDTree([])

    out = []
    total_ops = N + Q

    for _ in range(total_ops):
        op = data[idx]; idx += 1
        if op == b'add':
            x = int(data[idx]); idx += 1
            y = int(data[idx]); idx += 1
            pt = (x, y)
            buffer.append(pt)
            all_points.append(pt)
            if len(buffer) >= B:
                kdt = KDTree(all_points)
                buffer.clear()
        else:
            qx = int(data[idx]); idx += 1
            qy = int(data[idx]); idx += 1
            query = (qx, qy)
            best_p, best_d = kdt.nearest(query)
            for pt in buffer:
                d = dist2(query, pt)
                if d < best_d:
                    best_d = d; best_p = pt
            out.append(f"{best_p[0]} {best_p[1]}")

    sys.stdout.write('\n'.join(out) + '\n')

main()

Step-by-Step 解説

Step 1: KD-Tree の構造

各ノードに「点」「左右サブツリー」「バウンディングボックス(bbox)」「分割軸」を持つ。X軸とY軸を交互に使って2分割。

Step 2: 最近傍探索の枝刈り

クエリ点 $(qx, qy)$ からノードの bbox への最短距離² が現在の最良値以上なら探索不要。これにより期待 $O(\sqrt{N})$ で動作。

Step 3: Sqrt Decomposition

部分サイズクエリコスト
静的 KD-Tree最大 $N$ 点$O(\sqrt{N})$
バッファ最大 $B = O(\sqrt{N})$ 点$O(B) = O(\sqrt{N})$

Step 4: 計算量まとめ

操作計算量
add(バッファ追加)$O(1)$
add(再構築トリガー)$O(N \log N)$ / 回、合計 $O(N^{3/2})$
query$O(\sqrt{N})$ 期待

よくあるミス

ミス原因正しい書き方
再構築時に buffer だけ渡すall_points を渡さないと既存点が失われるKDTree(all_points) で全点を使う
バッファクリアを忘れるバッファが永遠に成長して線形スキャンが $O(N)$再構築後に buffer.clear()
KD-Tree が空のとき nearest を呼ぶ最初の数クエリで参照エラー空チェック後に return None, float('inf')

次のステップ

発展問題: 点の 削除 も扱う動的 KD-Tree(Scapegoat 型の再構築)または Logarithmic Rebuilding(バイナリ法でサイズ 1, 2, 4, ... の KD-Tree を複数維持)

自己評価