Official

E - ボールの転送 / Ball Transfer Editorial by admin

DeepSeek V3

概要

各人が右方向に投げたボールが、自分より厳密に身長が高い最も左の人にキャッチされるとき、その転送コスト(区間最小身長値)の合計を求める問題です。

考察

素朴なアプローチでは、各人 \(i\) について右側を線形探索して次に大きい身長の人 \(j\) を見つけ、区間 \([i, j]\) の最小値を求めることになります。この方法は \(O(N^2)\) 時間かかり、\(N\) が最大 \(2 \times 10^5\) であるため実行時間制限に間に合いません。

効率的に解くためには、次の2つの処理を高速化する必要があります: 1. 各人 \(i\) について、次に大きい身長の人 \(j\) を見つける(Next Greater Element) 2. 任意の区間 \([l, r]\) の最小値を高速に求める(Range Minimum Query)

アルゴリズム

  1. Next Greater Elementの計算:

    • スタックを用いて、各要素について次に大きい要素のインデックスを \(O(N)\) 時間で計算します。スタックにはインデックスを保持し、現在の要素がスタックトップの要素より大きい場合、スタックトップの要素の次に大きい要素は現在の要素であると判定します。
  2. Range Minimum Query (RMQ) の前処理:

    • スパーステーブルを構築します。スパーステーブルは \(O(N \log N)\) 時間の前処理を行い、任意の区間の最小値を \(O(1)\) 時間で求められるようにします。
    • 前処理では、長さ \(1, 2, 4, ..., 2^k\) の区間の最小値を事前に計算します。
  3. コストの計算:

    • 各人 \(i\) について、次に大きい身長の人 \(j\) が存在する場合、区間 \([i, j]\) の最小値をRMQで求め、合計に加算します。

計算量

  • 時間計算量: \(O(N \log N)\)
    • Next Greater Elementの計算: \(O(N)\)
    • スパーステーブルの構築: \(O(N \log N)\)
    • クエリ処理: \(O(N)\) (各クエリは \(O(1)\))
  • 空間計算量: \(O(N \log N)\)
    • スパーステーブルの保存に \(O(N \log N)\) のメモリを使用

実装のポイント

  • スタックを使用する際、未処理のインデックスを残さないように、ループ終了後に残ったインデックスのNext Greater Elementを \(N\)(存在しないことを表す)に設定します。

  • スパーステーブルの前処理では、2の冪乗の長さの区間を効率的に処理するため、対数テーブルを事前に計算しておきます。

  • クエリ処理時には、区間 \([l, r]\) を被覆する2つの事前計算済み区間の最小値を比較することで、\(O(1)\) 時間で最小値を求めます。

    ソースコード

import sys

def main():
    data = sys.stdin.read().split()
    n = int(data[0])
    H = list(map(int, data[1:1+n]))
    
    stack = []
    next_greater = [n] * n
    for i in range(n):
        while stack and H[stack[-1]] < H[i]:
            idx = stack.pop()
            next_greater[idx] = i
        stack.append(i)
    
    log_table = [0] * (n + 1)
    for i in range(2, n + 1):
        log_table[i] = log_table[i // 2] + 1
    
    k = log_table[n] + 1
    st = [[0] * n for _ in range(k)]
    st[0] = H[:]
    
    for j in range(1, k):
        step = 1 << (j - 1)
        for i in range(n - (1 << j) + 1):
            st[j][i] = min(st[j - 1][i], st[j - 1][i + step])
    
    def query_min(l, r):
        j = log_table[r - l + 1]
        return min(st[j][l], st[j][r - (1 << j) + 1])
    
    total = 0
    for i in range(n):
        j = next_greater[i]
        if j == n:
            continue
        min_val = query_min(i, j)
        total += min_val
        
    print(total)

if __name__ == "__main__":
    main()

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

posted:
last update: