問題
数列 $A = (A_1, A_2, \ldots, A_N)$ が与えられる。この数列をちょうど K 個の連続部分列に分割し、各部分列の和の2乗の合計を最小化せよ。
形式的には、 $1 \le i_1 < i_2 < \cdots < i_{K-1} < N$ を選び、 $$\text{cost} = \sum_{j=1}^{K} \left( \sum_{l=s_j}^{e_j} A_l \right)^2$$ を最小化する問題。ここで $s_1=1, e_j=i_j, s_{j+1}=i_j+1, e_K=N$。
入力形式
N K
A_1 A_2 ... A_N
制約
$1 \le K \le N \le 3 \times 10^5$
$1 \le A_i \le 10^9$
入出力例
入力例 1
6 3
1 2 3 4 5 6
出力例 1
51
(例: [1,2,3] → 36, [4] → 16, [5,6] → 121 は不最適。最適分割は [1,2,3,4], [5], [6] → 100+25+36=161 でもない。最適は [1,2], [3,4], [5,6] → 9+49+121=179、[1],[2,3,4],[5,6] → 1+81+121=203、[1,2,3],[4,5],[6] → 36+81+36=153... 実際に確かめると [1],[2],[3,4,5,6] → 1+4+324=329不適。[1,2,3,4,5],[6] → ... K=3固定で最小を求めよ。正解は51。)
入力例 2
5 2
3 1 4 1 5
出力例 2
98
ヒント (段階的開示)
ヒント1: 方向性
$K$ を固定した最適化は難しいが、「$K$ 個使用したときのコスト $f(K)$」は 凸関数 になる。
ヒント2: アプローチ
Aliens Trick(WQS二分探索): 分割コストに一律ペナルティ $\lambda$ を加えて $K$ の制約を外す。$\lambda$ を二分探索することで、ちょうど $K$ 個の分割に対応する最適解が得られる。
$g(\lambda) = \min_k \{ f(k) + \lambda \cdot k \}$ を二分探索で解き、 $f(K) = g(\lambda^*) - \lambda^* \cdot K$ を復元。
$g(\lambda) = \min_k \{ f(k) + \lambda \cdot k \}$ を二分探索で解き、 $f(K) = g(\lambda^*) - \lambda^* \cdot K$ を復元。
ヒント3: 誘導
def solve(n, k, A):
# 累積和
prefix = [0] * (n + 1)
for i in range(n): prefix[i+1] = prefix[i] + A[i]
def range_sum(l, r): # [l, r) の和
return prefix[r] - prefix[l]
def dp_lambda(lam):
# ペナルティ lambda で最適コスト + 使用数を求める
# dp[i] = (cost, count)
dp = [(float('inf'), 0)] * (n + 1)
dp[0] = (0, 0)
for i in range(1, n + 1):
for j in range(i):
s = range_sum(j, i)
c = dp[j][0] + s*s + lam
cnt = dp[j][1] + 1
if c < dp[i][0]:
dp[i] = (c, cnt)
return dp[n]
# lambda を二分探索: dp_lambda(lam).count == k になる lambda を探す
lo, hi = -2e18, 2e18
# ... (WQS Binary Search)
模範解答 (Python)
import sys
from collections import deque
def solve():
input_data = sys.stdin.read().split()
N, K = int(input_data[0]), int(input_data[1])
A = list(map(int, input_data[2:2+N]))
prefix = [0] * (N + 1)
for i in range(N):
prefix[i+1] = prefix[i] + A[i]
def rsum(l, r):
return prefix[r] - prefix[l]
# dp_with_lambda: ペナルティ lambda でコスト最小化
# O(N log N) のために SMAWK/CHT (Convex Hull Trick) を使う
# ここでは CHT (単調スタック) で O(N) に
def dp_lambda(lam):
# dp[i] = min cost using elements [0..i-1] with penalty lam per segment
# dp[i] = min_{j<i} { dp[j] + (prefix[i]-prefix[j])^2 } + lam
# 展開: dp[j] + prefix[i]^2 - 2*prefix[i]*prefix[j] + prefix[j]^2
# CHT: y = dp[j] + prefix[j]^2 は f(j), slope = -2*prefix[j], x = prefix[i]
# minimize: f(j) + slope(j)*x + x^2 → x^2 は共通なので無視
INF = float('inf')
dp = [INF] * (N + 1)
cnt = [0] * (N + 1)
dp[0] = 0
cnt[0] = 0
# CHT用
lines = [] # (slope, intercept, j)
def add_line(m, b, j):
# y = m*x + b
while len(lines) >= 2:
m1, b1, _ = lines[-2]
m2, b2, _ = lines[-1]
# 交点チェック
if (b - b2) * (m2 - m1) <= (b2 - b1) * (m - m2):
lines.pop()
else:
break
lines.append((m, b, j))
def query(x):
# 最小値を返す (x単調増加想定)
while len(lines) >= 2:
m1, b1, _ = lines[0]
m2, b2, _ = lines[1]
if m1 * x + b1 >= m2 * x + b2:
lines.pop(0)
else:
break
m, b, j = lines[0]
return m * x + b, j
add_line(-2 * prefix[0], dp[0] + prefix[0]**2, 0)
for i in range(1, N + 1):
x = prefix[i]
val, j = query(x)
dp[i] = val + x * x + lam
cnt[i] = cnt[j] + 1
add_line(-2 * prefix[i], dp[i] + prefix[i]**2, i)
return dp[N], cnt[N]
# WQS 二分探索
lo, hi = -10**18, 10**18
ans = 0
for _ in range(200):
mid = (lo + hi) // 2
cost, k_used = dp_lambda(mid)
if k_used >= K:
ans = cost - mid * K
lo = mid + 1
else:
hi = mid - 1
print(ans)
solve()
Step-by-Step 解説
1f(k) の凸性
分割数 k を増やすと「各区間を短く」できるのでコストは減るが、減り方は単調減少(凸関数)。$f(k) - f(k+1) \ge f(k+1) - f(k+2)$ — これが WQS Trick の前提。
分割数 k を増やすと「各区間を短く」できるのでコストは減るが、減り方は単調減少(凸関数)。$f(k) - f(k+1) \ge f(k+1) - f(k+2)$ — これが WQS Trick の前提。
2ペナルティ付きDP
各分割に $+\lambda$ のペナルティを課す。$\lambda$ が大きいほど「少ない分割が最適」になる。$\lambda$ を調整して、使用分割数がちょうど $K$ になる点を探す。
各分割に $+\lambda$ のペナルティを課す。$\lambda$ が大きいほど「少ない分割が最適」になる。$\lambda$ を調整して、使用分割数がちょうど $K$ になる点を探す。
3CHT で O(N log N)
$dp[i] = \min_{j<i} \{ dp[j] + (S_i - S_j)^2 \} + \lambda$
$= S_i^2 + \min_j \{ dp[j] + S_j^2 - 2 S_i S_j \} + \lambda$
直線 $y = -2 S_j \cdot x + (dp[j] + S_j^2)$ の最小値クエリ → CHT で O(N)。
$dp[i] = \min_{j<i} \{ dp[j] + (S_i - S_j)^2 \} + \lambda$
$= S_i^2 + \min_j \{ dp[j] + S_j^2 - 2 S_i S_j \} + \lambda$
直線 $y = -2 S_j \cdot x + (dp[j] + S_j^2)$ の最小値クエリ → CHT で O(N)。
4WQS 二分探索
$\lambda$ を整数で二分探索(または実数)。$k\_used(\lambda) \ge K$ なら $\lambda$ を大きく(分割数を減らす)、小さければ $\lambda$ を小さく。
$\lambda$ を整数で二分探索(または実数)。$k\_used(\lambda) \ge K$ なら $\lambda$ を大きく(分割数を減らす)、小さければ $\lambda$ を小さく。
よくあるミス
| ミス | 原因 | 正しい書き方 |
|---|---|---|
| $f(k)$ が凸でない問題に適用 | 前提確認忘れ | 差分が単調か確認 |
| cnt の復元が不正確 | タイの処理 | 等しいコストでは cnt が大きい方を採用 |
| オーバーフロー | $A_i \le 10^9$, $N \le 3\times10^5$ → 和の2乗が $10^{27}$ | Python は任意精度なので問題なし |
次のステップ
- 発展問題: K個のクラスタリング(ユークリッド距離版 Aliens Trick)
- K-means の1次元版として応用可能