DP高速化は条件を証明してから
遷移を速くする技法にも前提があります。最小化式を整理して、どの部分が状態に依存するかを切り分けます。
言語:Python / 計算量:基準 O(N²)、掲載の離散Li Chao版 O(N log N)
前提:区間DPは短い区間から
考え方
- dp[i]=min_j(dp[j]+(x_i-x_j)^2) を展開。
- x_i² + min_j((-2x_j)x_i + dp[j]+x_j²)。
- 過去の状態が直線になり、直線群の最小値問題へ変換できます。
具体例
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