【小结】斜率优化 DP
是单调队列优化 DP 里较难的一种。
HDU3507
如果没有 \(M\),那么每一个单独放一堆更好。
但是有 \(M\),问题变得复杂,考虑 DP。
根据题意,很容易想到一个一维的状态。
\(f_i\) 表示前 \(i\) 个单词全部写完的最小花费。
\(f_i=\min(f_i,f_j+(sum_i-sum_j)^2+M)(1\leq j < i)\)。
时间复杂度 \(\mathcal O(n^2)\)。
然后考虑优化枚举 \(j\) 的过程。
我们将枚举帮助转移的点称为“决策点”。
那么假设前面有 \(j,k\) 两个决策点,\(1\leq j<k<i\)。
不妨设 \(j\) 更优,那么需要满足的条件为:
\(f_j+(sum_i-sum_j)^2+M\leq f_k+(sum_i-sum_k)^2+M\)
\(\to f_j+sum_i^2+sum_j^2-2sum_isum_j+M\leq f_k+sum_i^2+sum_k^2-2sum_isum_k+M\)
\(\to f_j+sum_j^2-(f_k+sum_k^2)\leq 2sum_i(sum_j-sum_k)\)
设 \(y1=f_j+sum_j^2,y2=f_k+sum_k^2,x1=sum_j,x2=sum_k\)。
那么式子就被换元成了这样:
\(y1-y2\leq 2sum_i(x1-x2)\)
\(\to \frac{y1-y2}{x1-x2}\leq 2sum_i\)。
其中,\(\frac{y1-y2}{x1-x2}\) 为类似斜率的东西。所以这种技巧称之为“斜率优化 DP”。
那么考虑接下来怎么干。
如果斜率小于等于 \(2sum_i\),那么前面的点 \(k\) 就可以被踢出。因为 \(k\) 已经可以被 \(j\) 完全替代了。
所以维护一个单调队列,使得斜率不断增大(其实就是下凸壳)每一次在队列中放入节点编号。
当点 \(i\) 查询前面最小值的时候,把队头的所有斜率小于等于 \(2sum_i\) 的点出队(因为已经被替代,不是最优的)。
如此,根据单调性,后面的斜率必定都大于 \(2sum_i\)。
所以后面的一个不如一个优(至少在目前的局面下)。
如果前面的点,已经不如这个点优了,那么在后面 \(2sum_i\) 递增的情况下,永远都不如这个点优。可以直接被踢出了。
踢出之后,队头就是最优决策点。(前面不如这个点优,后面目前也不如)。
当 \(sum_i\) 不断增大时,本质上就是找到最小的一个大于 \(2sum_i\) 的点作为最优决策点。
决策点会不断地往后移动。
计算完答案,就是入队 \(i\)。要保持队列的单调性,也就要把斜率小于等于当前直线斜率的点全部弹走。
这样正确吗?一定正确。因为前面的点比你斜率大,那么前面的点,它干掉别人之前,别人已经把它干掉了。
所以它不可能称为最优决策点。
所以这样就可以正确的维护单调性。
可以使用斜率优化的形式大概是这样的:
\(f_i=\max(f_i,f_j+A_iB_j+C)\)。
斜率优化的使用条件:\(sum_i\) 单调。
注意:
-
删除/添加元素的时候必须保证 \(h+1\leq t\)。也就是队列中至少有 \(2\) 个元素。
-
分数比较大小用交叉相乘。
top 表示分子。down 表示分母。
#include<bits/stdc++.h>
using namespace std;
/*
*/
struct FSI{
template<typename T>
FSI& operator >> (T &res){
res=0;T f=1;char ch=getchar();
while (!isdigit(ch)){if (ch=='-') f=-1;ch=getchar();}
while (isdigit(ch)){res=res*10+ch-'0';ch=getchar();}
res*=f;
return *this;
}
}scan;
typedef long long ll;
const int N=5e6+10;
int n,i,j,h,t;
ll a[N],f[N],q[N],sum[N],m;
ll top(int i,int j)
{
return f[i]+sum[i]*sum[i]-(f[j]+sum[j]*sum[j]);
}
ll down(int i,int j)
{
return sum[i]-sum[j];
}
int main()
{
while (scanf("%d%lld",&n,&m)!=EOF)
{
for (i=1;i<=n;i++) scanf("%lld",&a[i]);
for (i=1;i<=n;i++) sum[i]=sum[i-1]+a[i];
h=1;
t=0;
q[++t]=0;
for (i=1;i<=n;i++)
{
while (h+1<=t&&top(q[h+1],q[h])<=2*sum[i]*down(q[h+1],q[h])) h++;
j=q[h];
f[i]=f[j]+(sum[i]-sum[j])*(sum[i]-sum[j])+m;
while (h+1<=t&&top(q[t],q[t-1])*down(i,q[t])>=top(i,q[t])*down(q[t],q[t-1])) t--;
q[++t]=i;
}
printf("%lld\n",f[n]);
}
return 0;
}

浙公网安备 33010602011771号