問題
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。
軸交互 + 中央値分割。各ノードに bbox。
2KNN
最大ヒープで $k$ 個管理。bbox 距離で枝刈り。
最大ヒープで $k$ 個管理。bbox 距離で枝刈り。
3RANGE
bbox 交差判定 → 範囲内なら点をカウント。期待 $O(\sqrt N)$。
bbox 交差判定 → 範囲内なら点をカウント。期待 $O(\sqrt N)$。
4bbox 距離
各軸独立に $\max(0, lo - x, x - hi)$ を計算。
各軸独立に $\max(0, lo - x, x - hi)$ を計算。
よくあるミス
| ミス | 原因 | 正しい書き方 |
|---|---|---|
| 近い子の選択ミス | diff 符号ミス | diff <= 0 → left |
| bbox 再計算 | 毎回計算 | 構築時に伝播 |
| ヒープが最小ヒープのまま | 負化忘れ | (-d2, p) |
次のステップ
- 動的 k-d Tree(追加・削除)
- 高次元近似最近傍 (LSH)