Day 017-Q1 — Aliens Trick(WQS二分探索・λ最適化DP)

2026-04-30 赤色 Master / Phase 8+ ★★★★★★★★★ Aliens Trick

問題

数列 $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$ を復元。
ヒント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 の前提。
2ペナルティ付きDP
各分割に $+\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)。
4WQS 二分探索
$\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次元版として応用可能

自己評価

自分の回答

気づき・メモ