题解:CF2182G Short Garland
挺好的一道。
首先看,对于一个节点,哪些东西是要存储的。
我们发现只需关注这个子树所有 dfn 序的末尾节点的深度即可。进一步的,可以统计每个深度有多少序,令其为
令
这是时空
选定一个长儿子,整体乘上其转移倍率,然后朴素转移其他儿子的
其实到这里已经做完了,只是需要规划一下代码:
- 用长链剖分给节点赋 dfn 序,这样长链就变成了区间,用线段树维护区间乘。而且可以解决“位移
1 格”的问题。 - 转移其他儿子时,先下传整个区间的懒标记,然后直接改变线段树子节点的值,可以做到
O(n\log n) 。 - 维护转移的最后一项不能乘逆元,因为可能有
cnt 的值是998244353 的倍数。可以用前缀积和后缀积。
:::success[代码]
#include<bits/stdc++.h>
#define ll long long
using namespace std;
const int N=3e5+7,mod=998244353;
int T,n,m,k,u,v,head[N],to[N],last[N];
int sz[N],son[N],f[N],cnt[N],y[N];
int idx,tot,rt[N],dep[N],p[N],p_[N],dfn[N],top[N];
ll mul[N];
struct node{
int l,r;
ll tm,ta,c;
}t[N*2];
void add(int u,int v){
++idx;
to[idx]=v,last[idx]=head[u],head[u]=idx;
}
int init(int l,int r){
int mid=l+r>>1,pos=++idx;
t[pos].ta=t[pos].c=0,t[pos].tm=1;
if(l==r){cnt[l]=pos;return pos;}
t[pos].l=init(l,mid),t[pos].r=init(mid+1,r);
return pos;
}
void calc(int pos,ll k,ll k_){
t[pos].c*=k,t[pos].c+=k_,t[pos].c%=mod;
t[pos].ta*=k,t[pos].ta+=k_,t[pos].ta%=mod;
t[pos].tm*=k,t[pos].tm%=mod;
}
void down(int pos){
ll k=t[pos].tm,k_=t[pos].ta;
calc(t[pos].l,k,k_),calc(t[pos].r,k,k_);
t[pos].tm=1,t[pos].ta=0;
}
void up(int pos){
t[pos].c=(t[t[pos].l].c+t[t[pos].r].c)%mod;
}
void add(int l,int r,int x,int y,ll k,int pos,bool op){
int mid=l+r>>1;
if(l==x&&r==y){
if(op) calc(pos,1,k);
else calc(pos,k,0);
return;
}
down(pos);
if(y<=mid) add(l,mid,x,y,k,t[pos].l,op);
else if(x>mid) add(mid+1,r,x,y,k,t[pos].r,op);
else add(l,mid,x,mid,k,t[pos].l,op),add(mid+1,r,mid+1,y,k,t[pos].r,op);
up(pos);
}
ll ask(int l,int r,int x,int y,int pos){
if(x>y) return 0;
int mid=l+r>>1;
if(l==x&&r==y) return t[pos].c;
down(pos);
if(y<=mid) return ask(l,mid,x,y,t[pos].l);
else if(x>mid) return ask(mid+1,r,x,y,t[pos].r);
return (ask(l,mid,x,mid,t[pos].l)+ask(mid+1,r,mid+1,y,t[pos].r))%mod;
}
void dfs(int u,int fa){
f[u]=fa;
for(int i=head[u];i;i=last[i]){
if(to[i]==fa) continue;
dfs(to[i],u),y[u]++;
if(dep[to[i]]>dep[son[u]]) son[u]=to[i];
}
dep[u]=dep[son[u]]+1;
}
void dfs(int u,int fa,bool op){
dfn[u]=++tot,top[u]=(op?top[fa]:u);
if(son[u]) dfs(son[u],u,1);
for(int i=head[u];i;i=last[i]){
if(to[i]!=fa&&to[i]!=son[u]) dfs(to[i],u,0);
}
}
void func(int l,int r,int x,int y,int pos){
if(x>y) return;
int mid=l+r>>1;
if(l==r) return;
down(pos);
if(y<=mid) func(l,mid,x,y,t[pos].l);
else if(x>mid) func(mid+1,r,x,y,t[pos].r);
else func(l,mid,x,mid,t[pos].l),
func(mid+1,r,mid+1,y,t[pos].r);
up(pos);
}
void dfs_(int u,int fa){
if(!son[u]){
add(1,n,dfn[u],dfn[u],1,1,1);
return;
}
vector<ll> pre(y[u]+2),nxt(y[u]+2),res(y[u]+2);
pre[0]=1;
int num=0,d=1,x=0,mx=0;
dfs_(son[u],u),++num;
res[num]=ask(1,n,dfn[son[u]],dfn[u]-1+min(dep[u],k),1);
pre[num]=res[num];
for(int i=head[u];i;i=last[i]){
if(to[i]==fa||to[i]==son[u]) continue;
dfs_(to[i],u),++num,x=dfn[to[i]];
res[num]=ask(1,n,x,x-1+min(dep[to[i]],k-1),1);
pre[num]=(pre[num-1]*res[num])%mod;
mx=max(mx,dep[to[i]]);
}
nxt[num+1]=1;
for(int i=num;i>=1;i--) nxt[i]=(nxt[i+1]*res[i])%mod;
add(1,n,dfn[u]+1,dfn[u]+dep[u]-1,nxt[2]*mul[num-1]%mod,1,0);
func(1,n,dfn[u]+1,dfn[u]+mx,1);
for(int i=head[u];i;i=last[i]){
if(to[i]==fa||to[i]==son[u]) continue;
++d,x=dfn[to[i]];
func(1,n,x,x+dep[to[i]]-1,1);
ll w=pre[d-1]*nxt[d+1]%mod*mul[num-1]%mod;
for(int j=x;j<x+dep[to[i]];j++){
(t[cnt[j-x+dfn[son[u]]]].c
+=t[cnt[j]].c*w%mod)%=mod;
}
}
func(1,n,dfn[u]+1,dfn[u]+mx,1);
}
int main(){
ios::sync_with_stdio(0);
cin.tie(0),cout.tie(0);
mul[0]=1;
for(int i=1;i<N;i++) mul[i]=(mul[i-1]*i)%mod;
cin>>T;
while(T--){
cin>>n>>k,init(1,n),idx=0;
for(int i=2;i<=n;i++) cin>>u,add(u,i);
tot=0,dfs(1,0),dfs(1,0,0),dfs_(1,0);
cout<<ask(1,n,dfn[1],dfn[1]+dep[1]-1,1)<<'\n';
for(int i=1;i<=n;i++){
dep[i]=sz[i]=head[i]=son[i]=f[i]=0;
}
for(int i=1;i<=idx;i++) to[i]=last[i]=0;
idx=0;
}
return 0;
}
:::