Official

E - 宇宙ステーションへの移動 / Traveling to the Space Station Editorial by admin

Claude Opus 5 (High)

概要

平面上の点(デブリ)を足場として、\((0,0)\) から \((0,W)\) まで「1 回の移動距離 \(\le D\)」で移動するときの最小ジャンプ回数を求める問題です。辺の重みがすべて 1 のグラフの最短路なので BFS ですが、辺が最大 \(O(N^2)\) 本あるため、空間データ構造で「近傍点の列挙+取り出し」を高速に行う必要があります。

考察

グラフとしての定式化

頂点を「母船 \((0,0)\)」「各デブリ \((x_i,y_i)\)」「ステーション \((0,W)\)」とし、距離が \(D\) 以下の 2 点間に辺を張ります。ただし「母船 → ステーション」の辺だけは張りません(必ずデブリを 1 つ以上経由する必要があるため)。

すべての辺のコストが 1 なので、最小ジャンプ回数 = 辺数最小の最短路 = BFSで求まります。

素朴な方法の問題点

素朴に隣接関係を全部作ると、\(N \le 10^5\) に対して点のペアが約 \(5\times10^9\) 通りあり、辺の構築だけで TLE・MLE します。つまり 隣接リストを明示的に作ってはいけない のがこの問題の核心です。

重要な観察:BFS では各点は 1 度しか使わない

BFS では、いったん距離が確定した頂点は二度と更新されません。したがって

「現在地 \((cx,cy)\) から距離 \(D\) 以内にある まだ未訪問の 点をすべて取り出す」

という操作が高速にできれば十分です。しかも取り出した点は構造から削除してよいので、削除込みの近傍クエリを合計 \(N\) 回やれば BFS が完了します。削除があるため「取り出される点の総数」は \(N\) 個で抑えられ、出力サイズの合計が小さいのがポイントです。

近傍クエリの実現

そこで、点集合を四分木(quadtree, k-d tree でもよい)に載せます。各ノードに

  • そのノードに含まれる(未削除の)点の個数 cnt
  • そのノードに含まれる点のバウンディングボックス \([\text{minx},\text{maxx}]\times[\text{miny},\text{maxy}]\)

を持たせておくと、クエリ時に

  • cnt == 0 のノードは即スキップ
  • 現在地からバウンディングボックスまでの最短距離が \(D\) を超えるノードは即スキップ

という枝刈りができ、円内の点だけを効率よく回収できます。回収した点は葉のリストから消し、祖先の cnt を減らします。

(別解として「一辺 \(D/\sqrt{2}\) のグリッドにバケット分割し、周囲 \(5\times5\) 程度のセルだけ見る」方法もあります。こちらは座標を辞書でハッシュして実装します。四分木は座標分布に依らず安定に動くのが利点です。)

アルゴリズム

  1. 入力を読み、全デブリを四分木に構築する。
    • 各ノードは、点数が閾値(例: 16)以下になるまで 4 分割する。
    • 各ノードに点数 cnt と実際の点のバウンディングボックスを持たせる。
  2. 母船 \((0,0)\) から距離 \(D\) 以内のデブリを query(0,0) で全部取り出し、距離 1 として BFS キューに入れる(同時に木から削除)。
  3. キューから点 \(u=(x_u,y_u)\) を取り出すたびに
    • まず ステーションに届くか を判定:\(x_u^2+(y_u-W)^2 \le D^2\) なら答えは dist[u]+1 で終了(BFS なので最初に見つかったものが最小)。
    • そうでなければ query(x_u,y_u) で距離 \(D\) 以内の未訪問デブリを取り出し、dist = dist[u]+1 を付けてキューに追加。
  4. キューが空になってもステーションに届かなければ -1。

BFS は距離の小さい順に頂点を処理するので、「最初にステーションへ届いたときの dist[u]+1」が答えになります。

計算量

\(N\) 個の点、四分木の高さを \(h\)(おおむね \(O(\log N)\)、最悪は座標範囲に依存)とすると

  • 木の構築: \(O(N \log N)\)

  • 点の取り出し: 各点はちょうど 1 回だけ取り出され、祖先の cnt 更新に \(O(h)\) → 合計 \(O(N\log N)\)

  • クエリの探索コスト: 枝刈りにより「円と交差するノード」だけを辿る。理論上の最悪は \(O(\sqrt N)\) 程度/クエリで、全体 \(O(N\sqrt N)\) だが、削除により構造がどんどん小さくなるため実際には非常に高速

  • 時間計算量: \(O(N\sqrt{N})\)(最悪見積り、実用的には \(O(N\log N)\) 程度)

  • 空間計算量: \(O(N)\)

実装のポイント

  • 平方根を取らない:距離比較は \(dx^2+dy^2 \le D^2\) と二乗のまま行い、浮動小数点誤差を完全に回避します。\(|x|\le 10^9,\ D \le 10^9\) なので、C++ なら long long(最大でも約 \(8\times10^{18}\) で収まる)を使います。

  • 母船→ステーションの直行を許さない:BFS の開始時に \((0,0)\) から \((0,W)\) の判定を「行わない」ことで自然に満たされます。ステーション到達判定はデブリを取り出したときにのみ行います。

  • バウンディングボックスまでの距離は軸ごとに独立に \(dx=\max(\text{minx}-cx,\ 0,\ cx-\text{maxx})\)、\(dy\) も同様として \(dx^2+dy^2\) で計算します。

  • 点を削除したら、祖先の cnt を減らし、必要なら葉のバウンディングボックスを再計算して枝刈りをより強くします(cnt == 0 のノードは以後一切触らない)。

  • 再帰ではなく明示的なスタックで木の構築・探索を行うと、Python でも再帰上限やオーバーヘッドを避けられます。入力は sys.stdin.buffer.read().split() で一括読み込みしましょう。

  • BFS キューはリストと添字 head で管理すると deque より軽く動きます(距離ごとに厳密な層分けは不要で、単調性は自動的に保たれます)。

    ソースコード

import sys

def main():
    data = sys.stdin.buffer.read().split()
    n = int(data[0]); W = int(data[1]); D = int(data[2])
    D2 = D * D
    xs = [0] * n; ys = [0] * n
    idx = 3
    for i in range(n):
        xs[i] = int(data[idx]); ys[i] = int(data[idx + 1]); idx += 2

    minx = min(xs); maxx = max(xs); miny = min(ys); maxy = max(ys)
    span = max(maxx - minx, maxy - miny) + 1
    size = 1
    while size < span:
        size <<= 1

    LEAF = 16
    cnt = [0]; par = [-1]
    bminx = [0]; bmaxx = [0]; bminy = [0]; bmaxy = [0]
    ch = [0, 0, 0, 0]
    leafpts = [None]

    stack = [(0, list(range(n)), minx, miny, size)]
    while stack:
        v, idxs, x0, y0, sz = stack.pop()
        cnt[v] = len(idxs)
        ax = bx = xs[idxs[0]]; ay = by = ys[idxs[0]]
        for i in idxs:
            x = xs[i]; y = ys[i]
            if x < ax: ax = x
            elif x > bx: bx = x
            if y < ay: ay = y
            elif y > by: by = y
        bminx[v] = ax; bmaxx[v] = bx; bminy[v] = ay; bmaxy[v] = by
        if len(idxs) <= LEAF or sz <= 1:
            leafpts[v] = idxs
            continue
        half = sz >> 1
        mx = x0 + half; my = y0 + half
        q0 = []; q1 = []; q2 = []; q3 = []
        for i in idxs:
            if xs[i] < mx:
                if ys[i] < my: q0.append(i)
                else: q1.append(i)
            else:
                if ys[i] < my: q2.append(i)
                else: q3.append(i)
        b4 = 4 * v
        k = 0
        for q, nx0, ny0 in ((q0, x0, y0), (q1, x0, my), (q2, mx, y0), (q3, mx, my)):
            if q:
                u = len(cnt)
                cnt.append(0); par.append(v)
                bminx.append(0); bmaxx.append(0); bminy.append(0); bmaxy.append(0)
                ch.append(0); ch.append(0); ch.append(0); ch.append(0)
                leafpts.append(None)
                ch[b4 + k] = u
                stack.append((u, q, nx0, ny0, half))
            k += 1

    def query(cx, cy, cnt=cnt, par=par, bminx=bminx, bmaxx=bmaxx, bminy=bminy,
              bmaxy=bmaxy, ch=ch, leafpts=leafpts, xs=xs, ys=ys, D2=D2):
        res = []
        st = [0]
        while st:
            v = st.pop()
            if cnt[v] == 0: continue
            t = bminx[v]
            if cx < t:
                dx = t - cx
            else:
                t = bmaxx[v]
                dx = cx - t if cx > t else 0
            t = bminy[v]
            if cy < t:
                dy = t - cy
            else:
                t = bmaxy[v]
                dy = cy - t if cy > t else 0
            if dx * dx + dy * dy > D2: continue
            lp = leafpts[v]
            if lp is None:
                b = 4 * v
                u = ch[b]
                if u: st.append(u)
                u = ch[b + 1]
                if u: st.append(u)
                u = ch[b + 2]
                if u: st.append(u)
                u = ch[b + 3]
                if u: st.append(u)
            else:
                keep = []
                nf = 0
                for i in lp:
                    ddx = xs[i] - cx; ddy = ys[i] - cy
                    if ddx * ddx + ddy * ddy <= D2:
                        res.append(i); nf += 1
                    else:
                        keep.append(i)
                if nf:
                    leafpts[v] = keep
                    u = v
                    while u >= 0:
                        cnt[u] -= nf
                        u = par[u]
                    if keep:
                        ax = bx = xs[keep[0]]; ay = by = ys[keep[0]]
                        for i in keep:
                            x = xs[i]; y = ys[i]
                            if x < ax: ax = x
                            elif x > bx: bx = x
                            if y < ay: ay = y
                            elif y > by: by = y
                        bminx[v] = ax; bmaxx[v] = bx; bminy[v] = ay; bmaxy[v] = by
        return res

    dist = [0] * n
    ans = -1
    q = query(0, 0)
    for i in q:
        dist[i] = 1
    head = 0
    while head < len(q):
        u = q[head]; head += 1
        xu = xs[u]; yu = ys[u]
        t = yu - W
        if xu * xu + t * t <= D2:
            ans = dist[u] + 1
            break
        nd = dist[u] + 1
        for i in query(xu, yu):
            dist[i] = nd
            q.append(i)
    sys.stdout.write(str(ans) + "\n")

main()

この解説は claude5opus-high によって生成されました。

posted:
last update: