1 条题解

  • 0
    @ 2026-4-23 17:14:21

    这题不卡三次方我建议评黄。。。

    先拆期望贡献,答案是每条边作为轻边的概率和子树大小的乘积的和。

    pu,lp_{u,l} 表示节点 uu 在长度为 ll 的重链的概率,对于每条边 favvfa_v\rightarrow v,我们计算他为重边的概率。

    gv,lg_{v,l} 表示 vv 的兄弟所在重链长度和为 ll 的概率,则 P=gv,ipv,jji+jP=\sum g_{v,i}p_{v,j}\frac{j}{i+j},直接求是 O(n3)O(n^3) 的,无法通过。

    fu,lf_{u,l} 表示 uu 的所有子节点所在的重链长度和为 ll 的概率,这个可以树上背包做到平方。则 $f_{u,i+j}=\displaystyle\sum_{i+j\leq size_u} g_{v,i}p_{v,j}$。将 gg 视为未知量,这是一个方程组。

    发现 pv,lp_{v,l} 不可能全为 00,记第一个非零的为 pv,kp_{v,k}。则 $g_{v,i}=\displaystyle\frac{f_{u,i+k}-\sum_{k<j\leq i+k-1}g_{v,i+k-j}p_{v,j}}{p_{v,k}}$,这样可以做到平方,后续可转移 pp 及求出答案。

    P.S. 求解的时候,可以先处理逆元,这样就可以做到不带 log。

    #include<iostream>
    #include<algorithm>
    using ll=long long;
    const int sz=5010;
    const ll mod=998244353;
    ll qpow(ll base,ll exp){
      ll ans=1;
      while(exp!=0){
        if(exp&1)ans=ans*base%mod;
        base=base*base%mod,exp>>=1;
      }
      return ans;
    }
    std::basic_string<int>graph[sz];
    ll p[sz][sz],inv[sz],f[sz][sz],g[sz],ans;
    int size[sz];
    void dfs(int u,int fau){
      size[u]=1,f[u][0]=1;
      if(graph[u].size()==1)return p[u][1]=1,void();
      for(int v:graph[u]){
        if(v==fau)continue;
        dfs(v,u);
        std::fill(g,g+size[u]+size[v]+1,0);
        for(int i=0;i<=size[u];i++)
          for(int j=1;j<=size[v];j++)g[i+j]=(g[i+j]+f[u][i]*p[v][j])%mod;
        size[u]+=size[v];
        std::copy(g,g+size[u]+1,f[u]);
      }
      if(graph[u].size()==2){
        int v=graph[u][0]!=fau?graph[u][0]:graph[u][1];
        for(int i=1;i<=size[v];i++)p[u][i+1]=p[v][i];
        return;
      }
      for(int v:graph[u]){
        if(v==fau)continue;
        int K=1;
        while(p[v][K]==0)K++;
        ll invp=qpow(p[v][K],mod-2);
        size[u]-=size[v];
        ll P=0;
        for(int i=1;i<=size[u];i++){
          g[i]=f[u][i+K];
          for(int j=std::min(size[v],i+K-1);j>K;j--)g[i]=(g[i]-g[i+K-j]*p[v][j]%mod+mod)%mod;
          g[i]=g[i]*invp%mod;
          for(int j=1;j<=size[v];j++){
            ll curp=g[i]*p[v][j]%mod*inv[i+j]%mod*j;
            P+=curp,p[u][j+1]+=curp;
          }
        }
        P%=mod;
        for(int j=1;j<=size[v];j++)p[u][j+1]%=mod;
        ans=(ans+size[v]*(mod-P+1))%mod,size[u]+=size[v];
      }
    }
    int main(){
      std::ios::sync_with_stdio(false);
      std::cin.tie(nullptr);
      int c,t;
      std::cin>>c>>t;
      while(t--)[]{
        int n;
        std::cin>>n;
        for(int i=1;i<=n;i++)inv[i]=qpow(i,mod-2);
        for(int i=1;i<=n;i++)graph[i].clear();
        for(int i=1;i<=n;i++)std::fill(f[i],f[i]+n+1,0);
        for(int i=1;i<=n;i++)std::fill(p[i],p[i]+n+1,0);
        for(int i=1,u,v;i<n;i++)std::cin>>u>>v,graph[u]+=v,graph[v]+=u;
        graph[1]+=0;
        ans=0,dfs(1,0);
        std::cout<<ans<<"\n";
      }();
      return 0;
    }
    
    • 1

    信息

    ID
    9673
    时间
    1500ms
    内存
    1024MiB
    难度
    10
    标签
    递交数
    2
    已通过
    1
    上传者