DP高速化は条件を証明してから

遷移を速くする技法にも前提があります。最小化式を整理して、どの部分が状態に依存するかを切り分けます。

言語:Python / 計算量:基準 O(N²)、掲載の離散Li Chao版 O(N log N)

前提:区間DPは短い区間から

考え方

  1. dp[i]=min_j(dp[j]+(x_i-x_j)^2) を展開。
  2. x_i² + min_j((-2x_j)x_i + dp[j]+x_j²)。
  3. 過去の状態が直線になり、直線群の最小値問題へ変換できます。

具体例

jを固定すると傾き -2x_j、切片 dp[j]+x_j²。x_i はクエリの座標です。単調な傾き・質問順がなければ単純なdeque型CHTは使えません。

実装

x=[0,2,5];dp=[0]*3
for i in range(1,3):
    dp[i]=min(dp[j]+(x[i]-x[j])**2 for j in range(i))
print(dp[-1])

注意する条件

Divide and Conquer最適化、Knuth最適化、CHTは別の条件を要求します。式の見た目だけで置換しないでください。

確認問題

x=[0,2,5] の例で dp[2] は?

解答と理由

13

0から直接なら25、2を経由なら4+9=13です。

実装課題

N と厳密増加配列 X(X_0=0)。dp[0]=0、dp[i]=min_{j<i}(dp[j]+(X_i-X_j)^2)。N≤500。基準実装で dp[N-1] を求める。

入力:
3
0 2 5
出力:
13
参考実装
n=int(input());x=list(map(int,input().split()));dp=[0]*n
for i in range(1,n):
    dp[i]=min(dp[j]+(x[i]-x[j])**2 for j in range(i))
print(dp[-1])

Li Chao treeで直線の最小値を管理

直線2本の大小関係は高々1回しか入れ替わりません。区間中央で小さい方を節点に残すと、もう一方が勝つ可能性のある側だけへ降ろせます。これが O(log M) で挿入できる理由です。

下の実装は質問するxを事前にすべて知っている離散座標版です。直線の追加順と質問順は任意ですが、xsに存在しないxは質問できません。Mは異なる質問座標数。構築O(M log M)、追加・取得O(log M)です。

dp[i]=min_{j<i}(dp[j]+(x_i-x_j)^2) なら、状態jの計算後に (-2*x_j, dp[j]+x_j*x_j) の直線を入れます。質問の答えに x_i*x_i を足してdp[i]にします。自分自身を候補に入れないよう、質問→追加の順にします。

from bisect import bisect_left

class LiChao:
    def __init__(self, xs):
        self.xs = sorted(set(xs))
        assert self.xs
        self.lines = [None] * (4 * len(self.xs))

    @staticmethod
    def value(line, x):
        return line[0]*x + line[1]

    def add(self, a, b):
        new = (a, b)
        k, l, r = 1, 0, len(self.xs)
        while True:
            if self.lines[k] is None:
                self.lines[k] = new
                return
            m = (l+r)//2
            cur = self.lines[k]
            left = self.value(new,self.xs[l]) < self.value(cur,self.xs[l])
            middle = self.value(new,self.xs[m]) < self.value(cur,self.xs[m])
            if middle:
                self.lines[k], new = new, cur
            if r-l == 1:
                return
            if left != middle:
                k, r = 2*k, m
            else:
                k, l = 2*k+1, m

    def query(self, x):
        i = bisect_left(self.xs, x)
        assert i < len(self.xs) and self.xs[i] == x
        k, l, r = 1, 0, len(self.xs)
        ans = float("inf")
        while True:
            if self.lines[k] is not None:
                ans = min(ans, self.value(self.lines[k], x))
            if r-l == 1:
                return ans
            m = (l+r)//2
            if i < m:
                k, r = 2*k, m
            else:
                k, l = 2*k+1, m

x = [0,2,5]
hull = LiChao(x)
hull.add(0,0)
for v in x[1:]:
    dp = v*v + hull.query(v)
    hull.add(-2*v, dp+v*v)
print(dp)  # 13

読了の記録・下書き・メモへ

関連する公式資料・課題