这是根据大部分题解提到的容斥做法码出的80分代码:
#include<bits/stdc++.h>
using namespace std;
#define N 100005
#define M 200010
#define e(i,x) for(int i=h[x];i;i=nx[i])
#define f(i,a,b) for(int i=a;i<=b;i++)
int n,k,x,y,h[N],v[M],nx[M],cnt,f[N][22][2],c[N],fa[N];
void dfs_fa(int x){
e(i,x){
if(v[i]==fa[x]) continue;
fa[v[i]]=x;dfs_fa(v[i]);
}
}
int dp(int x,int d,int tp){
if(f[x][d][tp]!=-1) return f[x][d][tp];
if(d<0||x==0) return 0;
if(tp==1) return f[x][d][tp]=(x==1?0:dp(fa[x],d-1,0)+dp(fa[x],d-1,1)-dp(x,d-2,0));
f[x][d][0]=c[x];
e(i,x) if(v[i]!=fa[x]) f[x][d][0]+=dp(v[i],d-1,0);
return f[x][d][0];
}
void add(int x,int y){
v[++cnt]=y;nx[cnt]=h[x];h[x]=cnt;
}
int main(){
cin>>n>>k;
f(i,1,n-1) cin>>x>>y,add(x,y),add(y,x);
dfs_fa(1);
f(i,1,n) cin>>c[i];
memset(f,-1,sizeof(f));
f(i,1,n) cout<<dp(i,k,0)+dp(i,k,1)<<'\n';
}
debug时发现构造长为 50 的链, k=20 时,当 i≥22 时, fi,20,1=−1 ,奇奇怪怪的。
然鹅把 f 数组第二维开到 25 就可以过了。
但是这样子,用大部分容斥题解的代码形式:
#include<bits/stdc++.h>
using namespace std;
#define N 100005
#define M 200010
#define ll long long
#define e(i,x) for(int i=h[x];i;i=nx[i])
#define f(i,a,b) for(int i=a;i<=b;i++)
int n,k,x,y,h[N],v[M],nx[M],cnt,f[N][22][2],c[N],fa[N];
void dfs_fa(int x){
e(i,x){
if(v[i]==fa[x]) continue;
fa[v[i]]=x;dfs_fa(v[i]);
}
}
void dfs(int x){
f(i,0,k) f[x][i][0]=c[x];
e(i,x){
if(v[i]==fa[x]) continue;
dfs(v[i]);
f(j,1,k) f[x][j][0]+=f[v[i]][j-1][0];
}
}
void dfs2(int x){
f(i,1,k){
f[x][i][1]=f[fa[x]][i-1][0]+f[fa[x]][i-1][1]-(i>1&&x!=1?f[x][i-2][0]:0);
}
e(i,x) if(v[i]!=fa[x]) dfs2(v[i]);
}
void add(int x,int y){
v[++cnt]=y;nx[cnt]=h[x];h[x]=cnt;
}
int main(){
cin>>n>>k;
f(i,1,n-1) cin>>x>>y,add(x,y),add(y,x);
dfs_fa(1);
f(i,1,n) cin>>c[i];
dfs(1);dfs2(1);
f(i,1,n) cout<<f[i][k][0]+f[i][k][1]<<'\n';
}
f 开到 22 就可以过了。
求问这是为什么。