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)\) で処理できます。
アルゴリズム
- 前処理: \(B_i = A_i + M\) を計算し、累積和 \(S\) を求める
- 座標圧縮: \(S_0, S_1, \ldots, S_N\) と各閾値 \(S_{r+1} - K\) をまとめてソートし、値を小さい整数に変換
- 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 によって生成されました。
投稿日時:
最終更新: