题解:CF2182G Short Garland

· · 题解

挺好的一道。

首先看,对于一个节点,哪些东西是要存储的。

我们发现只需关注这个子树所有 dfn 序的末尾节点的深度即可。进一步的,可以统计每个深度有多少序,令其为 cnt_{u,i}

u 子树 k-2 深度以内的序的数量和为 w_uuson_u 个儿子。枚举序在哪个子树内结束,乘上其他子树的贡献,转移即为:

cnt_{u,i}=(son_u-1)!\sum_{u \to v}{cnt_{v,i-1}\prod_ {u \to v' \wedge v \ne v'}} w_{v'}

这是时空 O(n^2) 的。考虑到转移很简单,所以尝试用启发式合并优化。

选定一个长儿子,整体乘上其转移倍率,然后朴素转移其他儿子的 cnt

其实到这里已经做完了,只是需要规划一下代码:

:::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;
}

:::