公式

E - 連続区間の選択 / Selection of Contiguous Intervals 解説 by admin

Claude 4.5 Opus

概要

数列から連続する区間を選んだときの「要素の合計 + 区間長 × M」というコストが \(K\) 以下となる区間 \((l, r)\) の個数を数える問題です。

考察

問題の言い換え

区間 \([l, r]\) のコストは以下のようになります: $\(\text{コスト} = \sum_{i=l}^{r} A_i + (r - l + 1) \times M\)$

ここで、\(B_i = A_i + M\) という新しい数列を定義すると: $\(\text{コスト} = \sum_{i=l}^{r} (A_i + M) = \sum_{i=l}^{r} B_i\)$

つまり、\(B\) の連続部分和が \(K\) 以下となる区間を数える」 という問題に帰着できます。

累積和による変形

\(B\) の累積和を \(S\) とします(\(S_0 = 0\), \(S_i = B_0 + B_1 + \cdots + B_{i-1}\))。

区間 \([l, r]\) のコストは \(S_{r+1} - S_l\) と表せるので、求める条件は: $\(S_{r+1} - S_l \leq K\)$

これを変形すると: $\(S_l \geq S_{r+1} - K\)$

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

\(r\) について、条件を満たす \(l\) を全探索すると \(O(N^2)\) となり、\(N \leq 2 \times 10^5\) では TLE になります。

解決策

\(r\) を順に見ていき、「これまでに見た \(S_l\)\(l = 0, 1, \ldots, r\))の中で、\(S_l \geq S_{r+1} - K\) を満たすものの個数」を高速に求めます。

これは 座標圧縮 + Binary Indexed Tree (BIT) を使うことで、各クエリを \(O(\log N)\) で処理できます。

アルゴリズム

  1. 前処理: \(B_i = A_i + M\) を計算し、累積和 \(S\) を求める
  2. 座標圧縮: \(S_0, S_1, \ldots, S_N\) と各閾値 \(S_{r+1} - K\) をまとめてソートし、値を小さい整数に変換
  3. BIT を使ったカウント:
    • 最初に \(S_0\) を BIT に追加
    • \(r = 0, 1, \ldots, N-1\) について:
      • 閾値 \(\text{threshold} = S_{r+1} - K\) を計算
      • BIT から「threshold 以上の値を持つ要素の個数」を取得し、答えに加算
      • \(S_{r+1}\) を BIT に追加(次の \(r\)\(l = r+1\) として使える)

具体例

\(N=3\), \(M=2\), \(K=10\), \(A = [1, 2, 3]\) の場合: - \(B = [3, 4, 5]\) - \(S = [0, 3, 7, 12]\)

\(r=0\) のとき:閾値は \(S_1 - K = 3 - 10 = -7\)\(S_0 = 0 \geq -7\) なので 1 通り。 \(r=1\) のとき:閾値は \(S_2 - K = 7 - 10 = -3\)\(S_0 = 0, S_1 = 3\) ともに \(\geq -3\) なので 2 通り。 \(r=2\) のとき:閾値は \(S_3 - K = 12 - 10 = 2\)\(S_0 = 0 < 2\), \(S_1 = 3 \geq 2\), \(S_2 = 7 \geq 2\) なので 2 通り。

合計:\(1 + 2 + 2 = 5\) 通り。

計算量

  • 時間計算量: \(O(N \log N)\)(座標圧縮のソートと、各 \(r\) についての BIT 操作)
  • 空間計算量: \(O(N)\)(累積和、座標圧縮用配列、BIT)

実装のポイント

  • 座標圧縮: \(S\) の値と閾値 \(S_{r+1} - K\) の両方を圧縮対象に含める必要があります

  • BIT の 1-indexed 化: BIT は通常 1-indexed で実装するため、インデックスの変換に注意

  • オーバーフロー対策: \(K\) が最大 \(10^{18}\) なので、Python では問題ありませんが、他言語では 64 ビット整数を使用する必要があります

  • 範囲クエリ: 「threshold 以上の個数」は「全体の個数 - threshold 未満の個数」として計算します

    ソースコード

import sys
from bisect import bisect_right

def solve():
    input_data = sys.stdin.read().split()
    N = int(input_data[0])
    M = int(input_data[1])
    K = int(input_data[2])
    A = [int(input_data[3 + i]) for i in range(N)]
    
    # B[i] = A[i] + M として、区間[l,r]のコストは sum(B[l..r]) となる
    # コスト = sum(A[l..r]) + (r-l+1)*M = sum(B[l..r])
    B = [A[i] + M for i in range(N)]
    
    # 累積和 S[0] = 0, S[i] = B[0] + B[1] + ... + B[i-1]
    # 区間[l,r]のコスト = S[r+1] - S[l]
    # S[r+1] - S[l] <= K を満たす (l, r) の個数を求める (0 <= l <= r < N)
    
    S = [0] * (N + 1)
    for i in range(N):
        S[i + 1] = S[i] + B[i]
    
    # 各 r について、S[r+1] - S[l] <= K となる l の個数を数える
    # S[l] >= S[r+1] - K となる l (0 <= l <= r) の個数
    
    # セグメント木やBITを使って、S[l]の値でインデックス付けし、
    # S[l] >= threshold となる個数を数える
    
    # 座標圧縮
    # 必要な値: S[0], S[1], ..., S[N] と、各 S[r+1] - K
    all_values = []
    for i in range(N + 1):
        all_values.append(S[i])
    thresholds = []
    for r in range(N):
        thresholds.append(S[r + 1] - K)
        all_values.append(S[r + 1] - K)
    
    # 座標圧縮
    sorted_values = sorted(set(all_values))
    value_to_idx = {v: i for i, v in enumerate(sorted_values)}
    num_values = len(sorted_values)
    
    # BIT (Binary Indexed Tree) for counting
    # BITは1-indexedで使う
    bit = [0] * (num_values + 2)
    
    def bit_add(i, delta=1):
        i += 1  # 1-indexed
        while i <= num_values + 1:
            bit[i] += delta
            i += i & (-i)
    
    def bit_sum(i):
        # [0, i] の合計 (0-indexed の i まで)
        i += 1  # 1-indexed
        s = 0
        while i > 0:
            s += bit[i]
            i -= i & (-i)
        return s
    
    def bit_range_sum(l, r):
        # [l, r] の合計 (0-indexed)
        if l > r:
            return 0
        if l == 0:
            return bit_sum(r)
        return bit_sum(r) - bit_sum(l - 1)
    
    # S[l] >= threshold となる l の個数を数える
    # これは、座標圧縮後のインデックスで threshold 以上の値を持つ要素の個数
    
    ans = 0
    
    # l = 0 を最初に追加
    bit_add(value_to_idx[S[0]])
    
    for r in range(N):
        # r を固定したとき、l は 0 から r まで
        # S[l] >= S[r+1] - K となる l の個数
        threshold = thresholds[r]
        threshold_idx = value_to_idx[threshold]
        
        # threshold_idx 以上の座標圧縮インデックスを持つ要素の個数
        count = bit_range_sum(threshold_idx, num_values - 1)
        ans += count
        
        # S[r+1] を追加 (次の r+1 のときの l として使える)
        bit_add(value_to_idx[S[r + 1]])
    
    print(ans)

solve()

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

投稿日時:
最終更新: