线段树(或树状数组)优化 DP

有时候 DP 转移时会涉及到区修区查,此时我们可以用线段树(或树状数组)来维护需要的值。

展示了一个最基本的和一个较难的题,其他的我觉得不难就不放了。懒得写

例题

因该是伟大的Wendyasif的原创题:严格上升子序列的数量

题目大意:求出长度为 \(n\) 的序列 \(A\) 中严格上升子序列的数量,对 \(10^9+7\) 取模。

\(1 \le n \le 10^5\)

还是老方法,先考虑暴力 DP 转移,在看能不能优化。

设 \(dp[i]\) 表示以 \(i\) 结尾的上升子序列的数量,则转移为:

\[dp[i]=\sum_{j=1,a[j]<a[i]}^{i-1}dp[j] \]

时间复杂度 \(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\) 组可以获得的最大价值,那么转移为:

\[dp[i,j]=\max_{x \in [j-1,i-1]}(dp[i-x,j-1]+c(x+1,i)) \]

\(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 使用,能用数组用数组。

然后我就抢了两个最优解。

posted @ 2026-02-25 19:28  vivid/stasis  阅读(34)  评论(0)    收藏  举报