回退背包学习笔记
回退背包,是背包 DP 的一种技巧,处理的情况是要从所有物品中退掉某一件物品。
主要利用的思想是背包问题与物品的求解顺序无关,因此每一个物品都可以认为是最后一个被求解的,那就可以对着 dp 的转移式子逆向操作,退回这个物品的 dp 状态。
[ICPC 2022 Jinan R] DFS Order 2
可以说是模板的一个题。对于每一个 \(j \in [1,n]\),询问一棵给定的树的每一个节点的 dfs 序是 \(j\) 的方案数。
先考虑一个子树内有几种 dfs 序的生成方案,令其为 \(t_x\),\(t_x=\prod t_v\times num_x!\),\(num_x\) 为 \(x\) 的儿子个数,其阶乘表示排列数。令 \(dp_{x,j}\) 表示节点 \(x\) 在 dfs 序中位置为 \(j\) 的方案数(为了便于转移,不考虑子树内部情况),那最终答案就是每一个 \(dp_{x,j} \times t_x\)。
这个 \(dp\) 的转移肯定是由父亲转移给儿子的,\(dp_{v,j} = \sum dp_{x,i} \times g_{j-i}\),令 \(g_{j-i}\) 表示 \(v\) 的 dfs 序和 \(x\) 相差 \(j-i\) 的方案数。新的问题是 \(g_j\) 怎么计算,和父亲相差 \(j\) 则必然是先去遍历了其他儿子,那我们就只好再设 \(f_{i,j}\) 表示选择了 \(i\) 个 \(x\) 的儿子,它们的子树一共有 \(j\) 个节点的方案数,这可以简单的背包求出来,记得 \(i\) 要从大到小枚举,这样每次用的 \(f_{i-1,j-siz_v}\) 才是上一轮的。
考虑用 \(f\) 怎么算 \(g\),既然 \(g\) 是在 \(v\) 还没被选的情况下设的,那么 \(f\) 就不能包含已经选过 \(v\) 的方案,暴力是每一个 \(v\) 跑一次排除掉他的背包,时间复杂度 \(O(n^4)\),考虑回退背包,每一次都令 \(tmp_{i,j}\) 表示退掉节点 \(v\) 的方案数,正向背包是从大到小枚举,我们逆操作从小到大减去即可。
for(int i=0;i<=cnt;i++){
for(int j=0;j<siz[x];j++){
tmp[i][j]=f[i][j];
g[j]=0;
}
}
for(int i=1;i<=cnt;i++){
for(int j=siz[x]-1;j>=siz[v];j--){
tmp[i][j]=(f[i][j]-tmp[i-1][j-siz[v]]+mod)%mod;
}
}
减去 \(tmp_{i-1,j-siz_v}\) 代表减去包含 \(v\) 的方案。
那么 \(g_j\) 应该由 \(tmp_{i,j-1}\) 转移过来,既然选了 \(i\) 个其他的儿子节点,那就会有顺序,就要乘上 \(i!\times (num_x-i-1)!\)。还有一件事,因为 \(dp\) 是不包含子树内部情况的,所以用来转移出 \(dp\) 的 \(g\) 也不能包含子树内的情况。那就要再乘一个 \(\frac{t_x}{t_v\times num_x!}\)。
#include<bits/stdc++.h>
using namespace std;
#define int long long
#define _int __int128
#define ull unsigned long long
#define pii pair<int,int>
#define fst first
#define scd second
#define pq priority_queue
#define mkp make_pair
#define popcount(x) __builtin_popcount(x)
#define endl '\n'
int n;
const int N = 5e2+10,mod=998244353;
vector<int> e[N];
int fac[N];
int qpow(int a,int b){
int res=1;
while(b){
if(b&1)res=(res*a)%mod;
a=(a*a)%mod;
b>>=1;
}
return res;
}
int t[N],siz[N],f[N][N],dp[N][N],g[N],tmp[N][N],num[N];
void dfs(int x,int fa){
t[x]=1;siz[x]=1;
for(auto v:e[x]){
if(v==fa)continue;
dfs(v,x);
num[x]++;
t[x]=(t[x]%mod*t[v]%mod)%mod;
siz[x]+=siz[v];
}
t[x]=(t[x]*fac[num[x]])%mod;
}
void dfs1(int x,int fa){
int cnt=num[x];
for(int i=0;i<=cnt;i++){
for(int j=0;j<=n;j++){
f[i][j]=0;g[j]=0;
}
}
f[0][0]=1;
for(auto v:e[x]){
if(v==fa)continue;
for(int i=cnt;i>=1;i--){
for(int j=siz[v];j<=siz[x];j++){
(f[i][j]+=f[i-1][j-siz[v]])%=mod;
}
}
}
for(auto v:e[x]){
if(v==fa)continue;
for(int i=0;i<=cnt;i++){
for(int j=0;j<siz[x];j++){
tmp[i][j]=f[i][j];
g[j]=0;
}
}
for(int i=1;i<=cnt;i++){
for(int j=siz[x]-1;j>=siz[v];j--){
tmp[i][j]=(f[i][j]-tmp[i-1][j-siz[v]]+mod)%mod;
}
}
int inv=qpow(t[v],mod-2),invf=qpow(fac[cnt],mod-2);
for(int i=0;i<cnt;i++){
for(int j=1;j<=siz[x];j++){
g[j]=(g[j]+tmp[i][j-1]%mod*fac[i]%mod*fac[cnt-i-1]%mod*t[x]%mod*inv%mod*invf%mod)%mod;
}
}
for(int i=1;i<=n;i++){
for(int j=i+1;j<=n;j++){
dp[v][j]=(dp[v][j]+dp[x][i]%mod*g[j-i]%mod)%mod;
}
}
}
for(auto v:e[x]){
if(v==fa)continue;
dfs1(v,x);
}
}
signed main(){
ios::sync_with_stdio(0);
cin.tie(0),cout.tie(0);
cin>>n;
fac[0]=1;
for(int i=1;i<=n;i++){
fac[i]=(fac[i-1]*i)%mod;
}
for(int i=1;i<n;i++){
int u,v;
cin>>u>>v;
e[u].push_back(v);
e[v].push_back(u);
}
dfs(1,0);
dp[1][1]=1;
dfs1(1,0);
for(int i=1;i<=n;i++){
for(int j=1;j<=n;j++){
cout<<dp[i][j]%mod*t[i]%mod<<" ";
}
cout<<endl;
}
return 0;
}
[COCI 2025/2026 #3] 国家 / Drzava
树上背包问题的回退。\(f_{x,i}\) 表示 \(x\) 的子树内不可选 \(x\),选择 \(i\) 个节点的方案数,\(F_{x,i}\) 则是可选择 \(x\) 的上述方案数,即 \(F_{x,1}=f_{x,1}+1\),且 \(f_x\) 是由 \(F_v\) 所合并出来的。朴素的思路就是枚举根节点,依次计算,复杂度 \(O(n^3)\)。
考虑优化,使用换根 DP,设 \(g_{x,i}\) 表示 \(x\) 的子树外,可选 \(x\),选择 \(i\) 个的方案数。那以 \(x\) 为根的时候的答案就可以用 \(f_{x}\) 和 \(g_{x}\) 合并出来。
考虑怎么计算 \(g_x\),先初始化 \(g_{1,0}=1\),然后考虑每次用父亲节点的信息计算出儿子节点的 \(g_v\)。先把以 \(x\) 为根时的答案用 \(fl=f \times g\) 存下来。发现 \(g_v\) 就是 \(fl\) 剔除掉 \(v\) 的子树内的方案,相当于回退背包把 \(F_v\) 退掉。
对于这种卷积合并我们怎么撤销呢?首先找到 \(F_v\) 的最高项 \(h\),然后从答案的最高项开始到 \(h\) 的每个项 \(i\),我都可以算出 \(g_{v,i-h} = \frac{fl_i}{F_{v,h}}\),但是一项会有不同的合并方式合出来,所以为了避免干扰到计算下一项,我们每算出一个 \(g_{v,i-h}\) 就把 \(fl\) 中它能参与合并的项 \(k\) 全部去除掉 \(g_{v,i-h}\times F_{v,k-i+h}\),不能理解的话列一下 \(a_0x_0+a_1x+1+a_2x_2\) 卷积上 \(b_0x_0+b_1x_1\) 自己模拟一下。
最后记得让 \(g_{v,1}=g_{v,1}+1\),因为它包含了自己。
时间复杂度为 \(O(n^2l)\),\(l\) 为任意两点之间的距离,此处不超过 \(36\)。
#include<bits/stdc++.h>
using namespace std;
#define int long long
#define _int __int128
#define ull unsigned long long
#define pii pair<int,int>
#define fst first
#define scd second
#define pq priority_queue
#define mkp make_pair
#define popcount(x) __builtin_popcount(x)
#define endl '\n'
int n;
const int N = 3e3+10,mod=1e9+7;
vector<int> e[N];
int f[N][N],F[N][N],siz[N],ans[N],tmp[N];
int qpow(int a,int b){
int res=1;
while(b){
if(b&1)res=(res*a)%mod;
a=(a*a)%mod;
b>>=1;
}
return res;
}
void add(int &a,int b){
a+=b;
if(a>=mod)a-=mod;
}
int tot;
void dfs(int x,int fa){
siz[x]=1;
for(auto v:e[x]){
if(v==fa)continue;
dfs(v,x);
siz[x]+=siz[v];
}
}
void dfs1(int x,int fa){
f[x][0]=1;
for(int i=1;i<=n;i++)f[x][i]=0,tmp[i]=0;
int sz=0;
for(auto v:e[x]){
if(v==fa)continue;
dfs1(v,x);
for(int j=0;j<=sz;j++){
for(int k=0;k<=siz[v];k++){
add(tmp[j+k],f[x][j]%mod*F[v][k]%mod);
}
}
sz+=siz[v];
for(int i=0;i<=sz;i++){
f[x][i]=tmp[i];
tmp[i]=0;
}
}
for(int i=0;i<=sz;i++){
F[x][i]=f[x][i];
}
F[x][1]=(f[x][1]+1)%mod;
}
int g[N][N],fl[N];
void dfs2(int x,int fa){
for(int i=0;i<=n;i++)fl[i]=0;
for(int j=0;j<=siz[x];j++){
for(int k=0;k<=n-siz[x];k++){
add(fl[j+k],f[x][j]%mod*g[x][k]%mod);
}
}
for(auto v:e[x]){
if(v==fa)continue;
for(int i=0;i<=n;i++)tmp[i]=fl[i];
int h=siz[v];
while(!F[v][h])h--;
int inv=qpow(F[v][h],mod-2);
for(int i=n;i>=h;i--){
if(tmp[i]){
int val=tmp[i]%mod*inv%mod;
g[v][i-h]=val;
for(int j=0;j<=h;j++){
tmp[i-j]=(tmp[i-j]-(F[v][h-j]%mod*val%mod)+mod)%mod;
}
}
}
add(g[v][1],1);
}
for(auto v:e[x]){
if(v==fa)continue;
dfs2(v,x);
}
}
signed main(){
ios::sync_with_stdio(0);
cin.tie(0),cout.tie(0);
cin>>n;
for(int i=1;i<n;i++){
int u,v;
cin>>u>>v;
e[u].push_back(v);
e[v].push_back(u);
}
dfs(1,0);
dfs1(1,0);
g[1][0]=1;
dfs2(1,0);
for(int x=1;x<=n;x++){
for(int j=0;j<=siz[x];j++){
for(int k=0;k<=n-siz[x];k++){
add(ans[j+k+1],f[x][j]%mod*g[x][k]%mod);
}
}
}
for(int i=1;i<=n;i++){
cout<<ans[i]<<" ";
}
return 0;
}
省选 2026 D1T1,当时场上完全不会。要求每个点到根的轻边数量之和的期望,就求出每条边是轻边的概率,一条边会贡献其儿子端子树内所有点的答案。
考虑怎么计算一条边是轻边的概率,令 \(p_{x,i}\) 表示 \(x\) 所在重链长度为 \(i\) 的概率,设 \(h_j\) 表示 \(x\) 的兄弟节点重链长度之和为 \(j\) 的概率。那么我们要求的 \(x\) 节点的的父边是轻边的概率就是 \(\sum p_{x,i}\times h_j \times \frac{j}{i+j}\)。
考虑怎么计算,\(h_i\) 是类似一个背包 DP,每次要排除一个儿子节点 \(v\),然后计算其他节点 \(z\) 的重链长度之和为 \(i\) 的概率,则 \(h_i=\sum_{j=0}^{\min(i,siz_z)}h'_{i-j}\times p_{z,j}\),注意 \(h'\) 是上一轮的 \(h\),每轮都要清空。\(p_{z,j}\) 在递归儿子时已经算过了。
然后来算 \(p_{x}\),\(p_{x,i+1}=\sum p_{v,i} \times h_{j} \times \frac{i}{i+j}\),值得注意的是,每个点的答案 \(pr_x\) 是在递归 \(fa_x\) 的时候算的。
时间复杂度 \(O(n^3)\),可以得 \(64\) 分,瓶颈在于每一个儿子节点都要单独算一次 \(h\) 的背包。
#include<bits/stdc++.h>
using namespace std;
#define int long long
#define _int __int128
#define ull unsigned long long
#define pii pair<int,int>
#define fst first
#define scd second
#define pq priority_queue
#define mkp make_pair
#define popcount(x) __builtin_popcount(x)
#define endl '\n'
int o,T,n;
const int N = 5e3+10,mod=998244353;
vector<int>e[N];
int inv[2*N];
int qpow(int a,int b){
int res=1;
while(b){
if(b&1)res=(res*a)%mod;
a=(a*a)%mod;
b>>=1;
}
return res;
}
int len[N],siz[N];
int p[N][N],h[N],pr[N];
void dfs(int x,int fa){
siz[x]=1;
for(auto v:e[x]){
if(v==fa)continue;
dfs(v,x);
siz[x]+=siz[v];
}
if(siz[x]==1)p[x][1]=1;
for(auto v:e[x]){
if(v==fa)continue;
for(int i=1;i<=n;i++)h[i]=0;
h[0]=1;
for(auto z:e[x]){
if(z==fa||z==v)continue;
for(int i=siz[x]-siz[v]-1;i>=0;i--){
h[i]=0;
for(int j=0;j<=min(i,siz[z]);j++){
h[i]=(h[i]+h[i-j]%mod*p[z][j]%mod)%mod;
}
}
}
for(int i=1;i<=siz[v];i++){
for(int j=(signed)e[x].size()-2;j<=siz[x]-siz[v];j++){
(p[x][i+1]+=p[v][i]%mod*h[j]%mod*i%mod*inv[i+j]%mod)%=mod;
(pr[v]+=(p[v][i]%mod*h[j]%mod)%mod*j%mod*inv[i+j]%mod)%=mod;
}
}
}
}
signed main(){
ios::sync_with_stdio(0);
cin.tie(0),cout.tie(0);
cin>>o>>T;
inv[1]=qpow(1,mod-2);
for(int i=2;i<=2*N-10;i++){
inv[i]=(mod-(mod/i)*inv[mod%i]%mod)%mod;
}
while(T--){
cin>>n;
for(int i=1;i<n;i++){
int u,v;
cin>>u>>v;
e[u].push_back(v);
e[v].push_back(u);
}
dfs(1,0);
int ans=0;
for(int i=2;i<=n;i++){
//cout<<pr[i]<<endl;
ans=(ans+pr[i]%mod*siz[i]%mod)%mod;
}
cout<<ans<<endl;
for(int i=1;i<=n;i++){
e[i].clear();
pr[i]=0;
}
memset(p,0,sizeof(p));
}
return 0;
}
这时候就要想到回退背包了,先算出总背包,再依次退回单独的儿子。具体的实现和上一题差不多,理解了上一题就会写这个题的代码。时间复杂度 \(O(n^2)\),有点卡常。
#include<bits/stdc++.h>
using namespace std;
#define ll long long
#define _int __int128
#define ull unsigned long long
#define pii pair<int,int>
#define fst first
#define scd second
#define pq priority_queue
#define mkp make_pair
#define popcount(x) __builtin_popcount(x)
#define endl '\n'
int o,T,n;
const int N = 5e3+10,mod=998244353;
vector<int>e[N];
ll inv[2*N];
ll read(){
char c=getchar();
int k=0,f=1;
for(;!isdigit(c);c=getchar())if(c=='-')f=-1;
for(;isdigit(c);c=getchar())k=(k*10+c-'0');
return k*f;
}
void write(ll x){
if(x<10)putchar(x+'0');
else write(x/10),putchar(x%10+'0');
}
ll qpow(ll a,ll b){
int res=1;
while(b){
if(b&1)res=(res*a)%mod;
a=(a*a)%mod;
b>>=1;
}
return res;
}
inline void add(ll &a,ll b){
a+=b;
if(a>=mod)a-=mod;
}
int num[N],siz[N];
ll p[N][N],h[N],pr[N],tmp[N],f[N];
void dfs(int x,int fa){
siz[x]=1;
for(int i=0;i<=n;i++)p[x][i]=0;
for(auto v:e[x]){
if(v==fa)continue;
dfs(v,x);
siz[x]+=siz[v];
num[x]++;
}
if(siz[x]==1)p[x][1]=1;
for(int i=1;i<=n;i++)h[i]=0;
h[0]=1;
int sz=0;
for(auto v:e[x]){
if(v==fa)continue;
for(int i=sz;i>=0;i--){
if(!h[i])continue;
for(int j=1;j<=siz[v];j++){
if(!p[v][j])continue;
add(h[i+j],h[i]*p[v][j]%mod);
}
h[i]=0;
}
sz+=siz[v];
}
for(auto v:e[x]){
if(v==fa)continue;
for(int i=0;i<=n;i++){
tmp[i]=h[i];
}
int h=siz[v];
while(h>=1&&!p[v][h])h--;
ll invx=qpow(p[v][h],mod-2);
for(int i=0;i<=siz[x];i++)f[i]=0;
for(int i=siz[x];i>=h;i--){
if(tmp[i]){
ll val=(tmp[i]*invx)%mod;
f[i-h]=val;
if(!val)continue;
for(int j=0;j<=h;j++){
tmp[i-j]=(tmp[i-j]-(val*p[v][h-j])%mod+mod)%mod;
}
}
}
for(ll i=1;i<=siz[v];i++){
for(ll j=0;j<=siz[x]-siz[v];j++){
int c=(p[v][i]*f[j])%mod*inv[i+j]%mod;
add(p[x][i+1],c*i%mod);
add(pr[v],c*j%mod);
}
}
}
}
signed main(){
ios::sync_with_stdio(0);
cin.tie(0),cout.tie(0);
o=read(),T=read();
inv[1]=qpow(1,mod-2);
int lst=1;
while(T--){
n=read();
if(n>lst){
for(int i=lst+1;i<=n;i++)inv[i]=(mod-(mod/i)*inv[mod%i]%mod)%mod;
lst=n;
}
for(int i=1;i<n;i++){
int u,v;
u=read(),v=read();
e[u].push_back(v);
e[v].push_back(u);
}
dfs(1,0);
ll ans=0;
for(int i=2;i<=n;i++){
//cout<<pr[i]<<endl;
ans=(ans+pr[i]*siz[i])%mod;
}
write(ans);
putchar(endl);
for(int i=1;i<=n;i++){
e[i].clear();
pr[i]=0;
}
}
return 0;
}

浙公网安备 33010602011771号