Day 032-Q5 — Balanced Partition(均等分割 + DP + WQS二分探索)

2026-05-15 赤色 Master / Phase 8+ ★★★★★★★★★ WQS二分探索・CHT

問題

長さ $N$ の整数列 $a_1, \ldots, a_N$ をちょうど $K$ 個の連続部分列に分割する。各セグメント $[l, r]$ のコストは $\left(\sum_{i=l}^{r} a_i\right)^2$。コストの総和を最小化せよ。

入力形式

N K
a_1 a_2 ... a_N

制約

$1 \le K \le N \le 3 \times 10^5$
$1 \le a_i \le 10^6$
答えは $10^{18}$ 以下

入出力例

入力例 1

6 3
1 2 3 4 5 6

出力例 1

51

ヒント (段階的開示)

ヒント1: 方向性
通常 DP は $O(N^2 K)$ で TLE。WQS二分探索(Aliens Trick)で $O(N \log N \log V)$ に。
ヒント2: アプローチ
$\lambda$ のペナルティをセグメント数に課して制約なし問題化。$\lambda$ 固定 DP は Convex Hull Trick で $O(N)$。
ヒント3: 誘導
dp[i] = min_j (dp[j] + (S[i]-S[j])^2 + lam)
      = min_j (-2*S[j])*S[i] + (dp[j]+S[j]^2+lam) + S[i]^2
# 直線 y = m*x + b, x=S[i], m=-2*S[j], b=dp[j]+S[j]^2+lam
# S[i] 単調増加 + 傾き単調減少 → deque CHT

模範解答 (Python)

import sys
from collections import deque

def main():
    data = sys.stdin.buffer.read().split()
    N, K = int(data[0]), int(data[1])
    a = [int(data[i+2]) for i in range(N)]

    prefix = [0] * (N + 1)
    for i in range(N):
        prefix[i+1] = prefix[i] + a[i]

    S = prefix

    def solve_with_penalty(lam):
        INF = float('inf')
        dp = [INF] * (N + 1)
        cnt = [0] * (N + 1)
        dp[0] = 0
        cnt[0] = 0

        hull = deque()

        def bad(l1, l2, l3):
            m1, b1, _ = l1
            m2, b2, _ = l2
            m3, b3, _ = l3
            return (b2-b1) * (m2-m3) >= (b3-b2) * (m1-m2)

        def add_line(m, b, j):
            line = (m, b, j)
            while len(hull) >= 2 and bad(hull[-2], hull[-1], line):
                hull.pop()
            hull.append(line)

        def query(x):
            while len(hull) >= 2:
                m1, b1, _ = hull[0]
                m2, b2, _ = hull[1]
                if m1*x + b1 >= m2*x + b2:
                    hull.popleft()
                else:
                    break
            m, b, j = hull[0]
            return m*x + b, j

        add_line(-2 * S[0], dp[0] + S[0]*S[0] + lam, 0)

        for i in range(1, N + 1):
            val, prev_j = query(S[i])
            dp[i] = val + S[i]*S[i]
            cnt[i] = cnt[prev_j] + 1

            if i < N:
                add_line(-2 * S[i], dp[i] + S[i]*S[i] + lam, i)

        return dp[N], cnt[N]

    lo, hi = 0, 2 * 10**18

    while lo < hi:
        mid = (lo + hi) // 2
        cost, k = solve_with_penalty(mid)
        if k <= K:
            hi = mid
        else:
            lo = mid + 1

    cost, k = solve_with_penalty(lo)
    ans = cost + lo * (K - k)
    print(ans)

main()

Step-by-Step 解説

1定式化
$\text{dp}[i][k] = \min_j (\text{dp}[j][k-1] + (S[i]-S[j])^2)$。$O(N^2 K)$ で TLE。
2WQS 二分探索の導入
$f(k)$ がセグメント数で凸 → ラグランジュ緩和 $g(\lambda) = \min_k(f(k) + \lambda k)$。適切な $\lambda$ で $k_{\text{opt}} = K$。
3CHT
$\lambda$ 固定 DP を直線形式に変形。$S_i$ 単調増 + 傾き単調減 → deque CHT $O(N)$。
4WQS 実装
$\lambda$ を二分探索し $k_\lambda \le K$ を満たす最小値を求める。最終答え: $\text{cost} + \lambda(K - k)$。

よくあるミス

ミス原因正しい書き方
WQS 補正忘れcost をそのまま出力cost + lambda * (K - k)
CHT の bad 関数の符号最大/最小で条件が逆最小: (b2-b1)*(m2-m3) >= (b3-b2)*(m1-m2)
$\lambda$ の上限が小さい最大コスト差を考慮しない$\text{hi} \sim (\sum a)^2$
セグメント数の単調性過信厳密単調でない場合あり<=> を慎重に

次のステップ

  • セグメントコストが「長さの3乗」の場合の WQS + CHT 適用可否

自己評価

自分の回答

気づき・メモ