Official

B - 山岳地帯の雨水シミュレーション / Rainwater Simulation in Mountainous Terrain Editorial by admin

Gemini 3.1 Pro (Thinking)

概要

各地点を標高の高い順に処理し、水が高いところから低いところへ流れる様子をシミュレーションして、最終的な各地点の水量を求める問題です。

考察

水は「標高が高い地点から低い地点へ」しか流れないという性質がこの問題の最大のポイントです。これにより、水路を「標高の高い地点から低い地点へ向かう一方通行の道(有向辺)」とみなすことができます。標高が同じ地点同士を結ぶ水路は水が流れないため、無視して構いません。

もし、実際の時間経過に沿って少しずつ水を流すような素朴なシミュレーションを行うと、水が流れるたびに状態を更新する必要があり、計算時間が膨大になってしまいます(TLEの原因となります)。

しかし、「標高の高い順」に地点を処理していくとどうなるでしょうか。 ある地点 \(v\) を処理する時点では、\(v\) より標高が高いすべての地点の処理はすでに終わっています。つまり、これ以降 \(v\) に新たに水が流れ込んでくることは絶対にありません。したがって、\(v\) を処理するタイミングで \(v\) の水量は完全に確定しており、その水を下流に1回分配するだけで、\(v\) に関する処理を完了させることができます。 これは、グラフ理論における「有向非巡回グラフ(DAG)のトポロジカルソート」と同じ考え方です。

アルゴリズム

  1. グラフの構築: 各水路について、結ばれた2地点の標高を比較します。標高が高い地点から低い地点へ向かう有向辺を張り、隣接リスト(adj)を作成します。標高が同じ場合は辺を張りません。
  2. ダムの記録: ダムが設置されている地点の情報を、真偽値の配列(is_dam)などに記録しておきます。
  3. 処理順序の決定: 地点の番号 \(1 \dots N\) を、標高 \(H\) の降順(高い順)にソートします。
  4. 流出シミュレーション: ソートした順に各地点 \(u\) を取り出し、以下の処理を行います。
    • \(u\) にダムが設置されている場合:水は流出しないため、何もしません。
    • \(u\) にダムがなく、下流隣接地点(有向辺の行き先)が \(d\) 個(\(d \geq 1\))ある場合:\(u\) の水量 \(W_u\)\(d\) 等分し、各下流隣接地点の水量に加算します。その後、\(W_u\)\(0\) にします。
    • \(u\) にダムがなく、下流隣接地点がない場合:水は流出しないため、何もしません。
  5. 結果の出力: すべての処理が終わった後の各地点の水量を出力します。

計算量

  • 時間計算量: \(O(N \log N + M)\) 地点のソートに \(O(N \log N)\) 時間がかかります。グラフの構築と流出シミュレーションでは、各地点と各水路(辺)を1回ずつしか調べないため \(O(N + M)\) 時間で完了します。全体として \(O(N \log N + M)\) となり、制約下で十分に高速です。
  • 空間計算量: \(O(N + M)\) 標高や水量の配列、グラフの隣接リストを保持するために使用します。

実装のポイント

  • 浮動小数点数の扱い: 水量は均等分配によって小数になるため、浮動小数点数(float など)で計算します。出力時には、誤差による不正解を防ぐため、小数点以下の桁数を十分に確保して出力すると安全です(例:Pythonの f"{w:.9f}" など)。

  • 高速な入出力: 入力されるデータ量(\(N, M\))が最大 \(2 \times 10^5\) と大きいため、Pythonなどの一部の言語では、入力を一括で読み込む(sys.stdin.read().split() など)工夫をすると実行時間を短縮できます。

    ソースコード

import sys

def solve():
    input_data = sys.stdin.read().split()
    if not input_data:
        return
    
    N = int(input_data[0])
    M = int(input_data[1])
    K = int(input_data[2])
    
    H = [int(x) for x in input_data[3:N+3]]
    W = [float(x) for x in input_data[N+3:2*N+3]]
    
    idx = 2 * N + 3
    
    adj = [[] for _ in range(N)]
    for _ in range(M):
        u = int(input_data[idx]) - 1
        v = int(input_data[idx+1]) - 1
        idx += 2
        
        if H[u] > H[v]:
            adj[u].append(v)
        elif H[v] > H[u]:
            adj[v].append(u)
            
    is_dam = [False] * N
    if K > 0:
        for _ in range(K):
            s = int(input_data[idx]) - 1
            is_dam[s] = True
            idx += 1
            
    order = sorted(range(N), key=lambda x: H[x], reverse=True)
    
    for u in order:
        if is_dam[u]:
            continue
        d = len(adj[u])
        if d > 0:
            flow = W[u] / d
            for v in adj[u]:
                W[v] += flow
            W[u] = 0.0
            
    print(*(f"{w:.9f}" for w in W))

if __name__ == '__main__':
    solve()

この解説は gemini-3.1-pro-thinking によって生成されました。

posted:
last update: