线段树(或树状数组)优化 DP
有时候 DP 转移时会涉及到区修区查,此时我们可以用线段树(或树状数组)来维护需要的值。
展示了一个最基本的和一个较难的题,其他的我觉得不难就不放了。懒得写
例题
因该是伟大的Wendyasif的原创题:严格上升子序列的数量
题目大意:求出长度为 \(n\) 的序列 \(A\) 中严格上升子序列的数量,对 \(10^9+7\) 取模。
\(1 \le n \le 10^5\)
还是老方法,先考虑暴力 DP 转移,在看能不能优化。
设 \(dp[i]\) 表示以 \(i\) 结尾的上升子序列的数量,则转移为:
时间复杂度 \(O(n^2)\) 会炸。
注意到比 \(a[i]\) 小的 \(a[j]\) 的 \(dp[j]\) 和可以用权值线段树维护,那么我们遍历到 \(i\) 时,查询出 \(dp[i]\),然后将线段树中 \(a[i]\) 的值加上 \(dp[i]\) 就可以了。
线段树优化代码
#include<bits/stdc++.h>
#define ls p<<1
#define rs p<<1|1
using namespace std;
using ll=long long;
const int mod=1e9+7;
struct stree{
int l,r;
ll sum;
}t[400005];
inline void pushup(int p){
t[p].sum=t[ls].sum+t[rs].sum;
t[p].sum%=mod;
}
void build(int p,int l,int r){
t[p].l=l,t[p].r=r;
if(l==r) return;
int mid=l+r>>1;
build(ls,l,mid);
build(rs,mid+1,r);
}
void update(int p,int x,int c){
if(t[p].l==x&&t[p].r==x){
t[p].sum=(t[p].sum+c)%mod;
return;
}
int mid=t[p].l+t[p].r>>1;
if(x<=mid) update(ls,x,c);
if(x>mid) update(rs,x,c);
pushup(p);
}
int query(int p,int l,int r){
if(l<=t[p].l&&t[p].r<=r) return t[p].sum;
ll s=0;
int mid=t[p].l+t[p].r>>1;
if(l<=mid) s=(s+query(ls,l,r))%mod;
if(r>mid) s=(s+query(rs,l,r))%mod;
return s;
}
ll dp[100005];
int a[100005],sat[100005];
int main(){
ios::sync_with_stdio(0);
cin.tie(0),cout.tie(0);
int n;
cin>>n;
for(int i=1;i<=n;i++){
cin>>a[i];
sat[i]=a[i];
dp[i]=1;
}
sort(sat+1,sat+1+n);
int nn=unique(sat+1,sat+1+n)-sat-1;
for(int i=1;i<=n;i++) a[i]=lower_bound(sat+1,sat+1+nn,a[i])-sat;
build(1,1,n);
for(int i=1;i<=n;i++){
int sum=0;
if(a[i]-1>0) sum=query(1,1,a[i]-1)%mod;
dp[i]=(dp[i]+sum)%mod;
update(1,a[i],dp[i]);
}
ll ans=0;
for(int i=1;i<=n;i++) ans=(ans+dp[i])%mod;
cout<<ans;
return 0;
}
/*
*/
树状数组优化代码
#include<bits/stdc++.h>
#define lowbit(x) x&-x;
using namespace std;
using ll=long long;
const int mod=1e9+7;
template<typename T=int>
inline void MOD(T &x){x=(x>mod?x-mod:x);}
int n;
int t[200005];
inline int path(int x){
int ans=0;
while(x){
ans+=t[x],MOD(ans);
x-=lowbit(x);
}
return ans;
}
inline void update(int x,int add){
while(x<n){
t[x]+=add,MOD(t[x]);
x+=lowbit(x);
}
}
inline int query(int l,int r){
return (path(r)-path(l-1)+mod)%mod;
}
int dp[100005];
int a[100005],sat[100005];
int main(){
ios::sync_with_stdio(0);
cin.tie(0),cout.tie(0);
cin>>n;
for(int i=1;i<=n;i++){
cin>>a[i];
sat[i]=a[i];
dp[i]=1;
}
sort(sat+1,sat+1+n);
int nn=unique(sat+1,sat+1+n)-sat-1;
for(int i=1;i<=n;i++) a[i]=lower_bound(sat+1,sat+1+nn,a[i])-sat;
for(int i=1;i<=n;i++){
int sum=0;
if(a[i]-1>0) sum=query(1,a[i]-1);
dp[i]+=sum,MOD(dp[i]);
update(a[i],dp[i]);
}
int ans=0;
for(int i=1;i<=n;i++) ans+=dp[i],MOD(ans);
cout<<ans;
return 0;
}
CF833B 自动糕蛋接贝场打包机(?)
题目大意:给定一个长度为 \(n\) 的序列 ,要求将其分成 \(k\) 段连续的⼦序列。每段的价值定义为该段中不同数字的个数,求最大价值和。
\(1 \le n \le 3 \times 10^4\),\(1 \le k \le \min(n,50)\)。
依旧是那个套路,先考虑暴力 DP,再考虑优化。
设 \(dp[i,j]\) 为前 \(i\) 个糕蛋分成 \(j\) 组可以获得的最大价值,那么转移为:
\(c(i,j)\) 表示 \([i,j]\) 的价值。
可以用 set 预处理出 \(c(i,j)\),DP 中要枚举 \(i\),\(j\),\(x\),时间复杂度为 \(O(n^2k)\),会炸。
看看 \(val(k+1,i)\) 是否具有规律。
序列:\([3,2,1,3,2,4]\)
列出 \(i,k\) 以及 \(val(k+1,i)\) 的表格:
| \(i\) 的取值 | 0 | 1 | 2 | 3 | 4 | 5 |
|---|---|---|---|---|---|---|
| 1 | 1 | |||||
| 2 | 2 | 1 | ||||
| 3 | 3 | 2 | 1 | |||
| 4 | 3 | 3 | 2 | 1 | ||
| 5 | 3 | 3 | 3 | 2 | 1 | |
| 6 | 4 | 4 | 4 | 3 | 2 | 1 |
表格来自 STY,我觉得挺好理解的就挂着了。
注意到:
- 每一次最多增加 1。
- 每一次增加的区间连续。
再注意到:每一次增加的范围是 \(a[i]\) 上一次出现的位置到 \(i-1\)。
\(a[i]\) 上一次出现的位置可以直接在 DP 结束时,用一个数组 \(pre\) 记录,$ pre[a[i]]=i $
这下就找到连续的区间了。
若我们先枚举 \(j\),则 \(j\) 对于内部为定值,所以我们可以线段树找到任意 \(i\) 时 \(c(k+1,i)\) 的最大值。
也可以直接开多棵线段树维护。封装,保留我的码风,不要有注释
线段树优化代码
#include<bits/stdc++.h>
#define MOD(x) x=(x>mod?x-mod:x)
using namespace std;
int a[120005],pre[120005];
int dp[120005][55];
struct SEG{
#define ls p<<1
#define rs p<<1|1
struct stree{
int l,r;
int max,add;
}t[160005];
inline void pushup(int p){
t[p].max=max(t[ls].max,t[rs].max);
}
inline void pushdown(int p){
if(t[p].add){
t[ls].max+=t[p].add;
t[rs].max+=t[p].add;
t[ls].add+=t[p].add;
t[rs].add+=t[p].add;
t[p].add=0;
}
}
void build(int p,int l,int r){
t[p].l=l,t[p].r=r;
if(l==r) return;
int mid=(l+r)>>1;
build(ls,l,mid);
build(rs,mid+1,r);
}
void update(int p,int l,int r,int add){
if(l<=t[p].l&&t[p].r<=r){
t[p].max+=add;
t[p].add+=add;
return;
}
int mid=(t[p].l+t[p].r)>>1;
pushdown(p);
if(l<=mid) update(ls,l,r,add);
if(r>mid) update(rs,l,r,add);
pushup(p);
}
int query(int p,int l,int r){
if(l<=t[p].l&&t[p].r<=r) return t[p].max;
int mid=(t[p].l+t[p].r)>>1,res=0;
pushdown(p);
if(l<=mid) res=max(res,query(ls,l,r));
if(r>mid) res=max(res,query(rs,l,r));
return res;
}
}tr[55];
signed main(){
ios::sync_with_stdio(0);
cin.tie(0),cout.tie(0);
int n,k;
cin>>n>>k;
for(int i=1;i<=n;i++) cin>>a[i];
for(int i=1;i<=k;i++) tr[i].build(1,0,n);
for(int i=1;i<=n;i++){
for(int j=1;j<=k;j++){
tr[j].update(1,i-1,i-1,dp[i-1][j-1]);
tr[j].update(1,pre[a[i]],i-1,1);
dp[i][j]=tr[j].query(1,j-1,i-1);
// cout<<dp[i][j]<<" ";
}
// cout<<"\n";
pre[a[i]]=i;
}
cout<<dp[n][k];
return 0;
}
树状数组优化代码
``` 咕咕咕 ```总的来说,线段树优化 DP 要注意以下几点:
-
线段树是否打对,可能对着正确的 DP 代码调很久,然后发现是线段树打挂了。
别问我怎么知道的 -
区间修改或查询的范围是否正确。如 CF833B,修改的范围不包括 \(i\)。
-
常数是否爆炸,如果是可以卡卡常或者改成树状数组。
所以今天我不仅学到了线段树优化 DP,还学到了几个卡常技巧。
卡常技巧
1.取模的时候,如果只有加法且两个加数都小于 mod,那么可以使用 x>mod?x-mod:x 大法。990ms -> 430ms
2.使用三目运算符代替 STL 的 max/min 函数不一定更快。替换前 AC,替换后 TLE60tps。但有些时候就是更快(玄学卡常,慎用)
3.减少递归调用,能写循环写循环,比如树状数组和并查集,至少快 5 倍。
4.减少 STL 使用,能用数组用数组。
然后我就抢了两个最优解。

线段树优化 DP 和小卡常技巧
浙公网安备 33010602011771号