Official

C - お菓子の詰め合わせ / Assortment of Sweets Editorial by admin

Claude 4.6 Opus (Thinking)

概要

連続する部分列の重さの合計が、袋の耐荷重の最小値 \(C_{\min}\) 以下となるような選び方 \((l, r)\) の個数を求める問題です。尺取り法(Two Pointers)を用いて効率的に数え上げます。

考察

問題の本質

袋の耐荷重は \(M\) 個ありますが、使うのは最小値 \(C_{\min} = \min(C_1, \ldots, C_M)\) だけです。よって問題は次のように言い換えられます:

\(W_l + W_{l+1} + \cdots + W_r \leq C_{\min}\) を満たす \((l, r)\)\(1 \leq l \leq r \leq N\))の個数を求めよ。

素朴なアプローチとその問題点

すべての \((l, r)\) の組を全探索すると、組の数は \(O(N^2)\) 個あり、各組について合計を計算すると最悪 \(O(N^3)\) かかります。累積和を使えば \(O(N^2)\) になりますが、\(N \leq 5 \times 10^5\) では \(N^2 = 2.5 \times 10^{11}\) となり、到底間に合いません。

重要な観察:単調性

すべてのお菓子の重さは正(\(W_i \geq 1\))なので、次の性質が成り立ちます:

  • \(r\) を固定したとき\(l\) を小さくするほど合計は大きくなる(お菓子が増えるため)。
  • ある \(l\) で合計が \(C_{\min}\) を超えるなら、それより小さい \(l\) でも必ず超える。

この単調性があるため、尺取り法が使えます。

アルゴリズム

尺取り法(しゃくとり法 / Two Pointers)で解きます。

  1. \(C_{\min} = \min(C_1, \ldots, C_M)\) を計算する。
  2. left = 0(0-indexed), current_sum = 0, ans = 0 で初期化する。
  3. right\(0\) から \(N-1\) まで動かしながら以下を行う:
    • current_sum\(W[\text{right}]\) を加える。
    • current_sum > C_{\min} である限り、current_sum から \(W[\text{left}]\) を引いて left\(1\) 増やす。
    • この時点で、\([\text{left}, \text{right}]\) から \([\text{right}, \text{right}]\) までのすべての区間が条件を満たす。その個数は \(\text{right} - \text{left} + 1\) なので、ans に加える。

具体例

\(N=4\), \(W = [3, 1, 4, 1]\), \(C_{\min} = 5\) の場合:

right 追加 current_sum left の移動 有効な区間 個数
0 3 3 なし [0,0] 1
1 1 4 なし [0,1],[1,1] 2
2 4 8→5 left: 0→1 [1,2],[2,2] 2
3 1 6→2 left: 1→2 [2,3],[3,3] 2

答えは \(1+2+2+2 = 7\) 通りです。

計算量

  • 時間計算量: \(O(N + M)\)
    • \(C_{\min}\) の計算に \(O(M)\)、尺取り法に \(O(N)\)leftright はそれぞれ最大 \(N\) 回しか進まない)。
  • 空間計算量: \(O(N + M)\)
    • 配列 \(W\)\(C\) の格納に必要。

実装のポイント

  • \(C_j\) の値が最大 \(10^{18}\) と非常に大きいため、累積和が 64 ビット整数の範囲に収まるか注意が必要です。Python は多倍長整数なので問題ありませんが、C++ などでは long long を使う必要があります。

  • 高速入力: \(N, M\) が最大 \(5 \times 10^5\) と大きいため、Python では sys.stdin.buffer.read() でまとめて読み込むことで入力を高速化しています。

  • 尺取り法では leftright を超えることはありません。すべての \(W_i \geq 1\) かつ \(C_{\min} \geq 1\) なので、少なくとも長さ 1 の区間(お菓子 1 個)は必ず条件を満たすことが保証されます。

    ソースコード

import sys

def main():
    input_data = sys.stdin.buffer.read().split()
    idx = 0
    N = int(input_data[idx]); idx += 1
    M = int(input_data[idx]); idx += 1
    W = [int(input_data[idx + i]) for i in range(N)]; idx += N
    C = [int(input_data[idx + i]) for i in range(M)]; idx += M
    
    C_min = min(C)
    
    ans = 0
    current_sum = 0
    left = 0
    for right in range(N):
        current_sum += W[right]
        while current_sum > C_min:
            current_sum -= W[left]
            left += 1
        ans += right - left + 1
    
    print(ans)

main()

この解説は claude4.6opus-thinking によって生成されました。

posted:
last update: