题解:P17249 【Gensokyo OI Round 2】长夜梦终觉

· · 题解

题意简述

一棵树按输入顺序依次断边。对每个得到的森林,加入最少数量的边使其重新成为一棵树。求最大直径及达到最大值的连边方案数。

解题思路

设当前森林有 m 个连通块。对第 i 个连通块,记直径长度为 l_i。能作为某条直径端点的顶点数记为 c_i,有序直径端点对数记为 f_i。单点的两个连接位置可以取同一个点,故规定其 f_i=1

一条路径在每个连通块内至多经过 l_i 条边。在连通块之间,它至多经过 m-1 条新边。因此,最大直径为:

D=m-1+\sum_{i=1}^m l_i

把所有连通块排成一条链,并依次连接它们的直径端点,可以达到这个上界。

考虑一种带方向的连通块排列。两端的连通块各选择一个直径端点,中间每个连通块选择一个有序直径端点对。反转整条链不会改变新增边集,且每种方案恰好被两个方向计算。定义:

P(x)=\prod_{i=1}^m(f_i+c_ix)

m\ge2 时,方案数为:

C=(m-2)![x^2]P(x)

m=1 时不需要加边,方案数为 1。 只保留 P(x) 的零至二次项,并用迭代乘积树维护全部因子。 两个连通块合并时,把较小连通块的因子改成 1, 再更新合并后连通块的因子,每次只需 O(\log n)

下面计算单个连通块的 cf。若直径为偶数 2r,它有唯一中心点。设第 j 个相邻分支内有 b_j 个点到中心的距离为 r,则:

\begin{aligned} c & =\sum_j b_j \\ f & =c^2-\sum_j b_j^2 \end{aligned}

两个端点必须来自不同分支。若直径为奇数 2r+1,中心是一条边。设中心边两侧分别有 c_0,c_1 个最远点,则:

\begin{aligned} c & =c_0+c_1 \\ f & =2c_0c_1 \end{aligned}

正向删边不便维护,故离线倒序加边。合并两个连通块时,从连接点开始遍历较小连通块。将其中的点依次作为叶子加入较大连通块。按大小合并后,每个点至多移动 O(\log n) 次。

考虑加入叶子 x。若当前直径为 2r,中心为 z,记 d=\operatorname{dist}(x,z)

若当前直径为 2r+1,设 x 属于中心边的第 s 侧。记它到该侧中心点的距离为 d

新点与当前连通块相邻,以上距离不会超过对应半径加一,因此这些情况完整覆盖每次转移。分支计数用版本号复用全局数组,不需要清空旧中心的全部分支。

距离和路径第一步都在原树上查询。将原树定根并求深度优先搜索(Depth-First Search,DFS)序。对 DFS 序区间维护深度最小值,深度相同时取序号较大的点。这样可以 O(1) 求最近公共祖先和路径第一步。预处理、按大小合并和线段树维护的时间复杂度均为 O(n\log n)。空间复杂度为 O(n\log n)

正确性证明

任意重连方案的最长路径在每个原连通块内至多走过其直径。它还至多经过全部 m-1 条新边,所以长度不超过 D。把连通块排成链并连接直径端点后,所得路径长度恰为 D。故算法计算的最大直径正确。

达到 D 时,最长路径必须经过每个连通块的整条直径和每条新边。因此,连通块之间形成的树必须是一条链。链端各贡献 c_i 种选择,中间连通块各贡献 f_i 种有序选择。带方向排列除以反转产生的两次计数后,恰为 (m-2)![x^2]P(x)。故线段树维护的方案数不重不漏。

树的中心性质保证:偶直径的端点必须位于中心的不同最深分支。奇直径的端点必须位于中心边两侧。由此得到两种 f 的公式。加入一个叶子时,直径至多增长一。上述距离分类覆盖直径不变、增加同层端点和中心移动三种情况,并准确更新对应最远点集合。因此,每次合并后维护的这些量都与当前连通块一致。

倒序加入前 i 条边所得森林,正是删去输入前若干条边后的对应状态。算法在每次合并后记录答案,所以输出的全部状态均正确。

参考代码

#include <bits/stdc++.h>
using namespace std;

using ll=long long;
const int N=300005;
const int K=20;
const int S=1<<19;
const int mod=998244353;
struct edge
{
    int u,v;
}e[N];
struct component
{
    int siz,len,sum,val,ver;
    int md[2],cnt[2];
    bool od;
}t[N];
struct poly
{
    int a[3];
}p[S<<1];
vector<int> G[N];
int n,tim;
int dfn[N],dep[N],fa[N],lg[N],st[N][K],tag[N],stk[N],br[N],bv[N];
int ans_d[N],ans_c[N],fac[N];
int mn(int x,int y)
{
    if(dep[x]!=dep[y])return dep[x]<dep[y]?x:y;
    return dfn[x]>dfn[y]?x:y;
}
int qry(int l,int r)
{
    int k=lg[r-l+1];
    return mn(st[l][k],st[r-(1<<k)+1][k]);
}
int dis(int x,int y)
{
    if(x==y)return 0;
    if(dfn[x]>dfn[y])swap(x,y);
    return dep[x]+dep[y]-2*dep[fa[qry(dfn[x]+1,dfn[y])]];
}
int nxt(int x,int y)
{
    if(dfn[x]<dfn[y])return fa[y];
    int z=qry(dfn[y]+1,dfn[x]);
    return dep[z]>dep[y]?z:fa[y];
}
int get(int id,int x)
{
    return bv[x]==t[id].ver?br[x]:0;
}
void refresh(int id)
{
    if(t[id].len==0)t[id].val=1;
    else if(t[id].od)t[id].val=int((ll)t[id].cnt[0]*t[id].cnt[1]%mod*2%mod);
    else t[id].val=int(((ll)t[id].cnt[0]*t[id].cnt[0]-t[id].sum+mod)%mod);
}
void add(int id,int x)
{
    t[id].siz++;
    if(!t[id].od)
    {
        int d=dis(x,t[id].md[0]),r=t[id].len/2;
        if(d==r)
        {
            int y=nxt(x,t[id].md[0]),c=get(id,y);
            t[id].cnt[0]++;
            t[id].sum=int((t[id].sum+(ll)c*2+1)%mod);
            br[y]=c+1;
            bv[y]=t[id].ver;
        }
        else if(d==r+1)
        {
            int y=nxt(x,t[id].md[0]);
            t[id].cnt[0]-=get(id,y);
            t[id].cnt[1]=1;
            t[id].md[1]=y;
            t[id].len++;
            t[id].od=1;
        }
    }
    else
    {
        int d[2]={dis(x,t[id].md[0]),dis(x,t[id].md[1])};
        int s=d[0]<d[1]?0:1,r=t[id].len/2;
        if(d[s]==r)t[id].cnt[s]++;
        else if(d[s]==r+1)
        {
            int z=t[id].md[s],o=t[id].md[s^1],y=nxt(x,z),c=t[id].cnt[s^1];
            t[id].md[0]=z;
            t[id].cnt[0]=c+1;
            t[id].cnt[1]=0;
            t[id].sum=int(((ll)c*c+1)%mod);
            t[id].len++;
            t[id].od=0;
            t[id].ver=++tim;
            br[o]=c;
            bv[o]=t[id].ver;
            br[y]=1;
            bv[y]=t[id].ver;
        }
    }
    refresh(id);
}
poly mul(poly x,poly y)
{
    poly z={};
    for(int i=0;i<=2;i++)for(int j=0;i+j<=2;j++)z.a[i+j]=int((z.a[i+j]+(ll)x.a[i]*y.a[j])%mod);
    return z;
}
poly make(int id)
{
    poly res={};
    if(!id)res.a[0]=1;
    else
    {
        res.a[0]=t[id].val;
        res.a[1]=(t[id].cnt[0]+t[id].cnt[1])%mod;
    }
    return res;
}
void upd(int x,int id)
{
    int u=S+x-1;
    p[u]=make(id);
    for(u>>=1;u;u>>=1)p[u]=mul(p[u<<1],p[u<<1|1]);
}
int main()
{
    ios::sync_with_stdio(false);
    cin.tie(nullptr);
    cin>>n;
    for(int i=1;i<n;i++)
    {
        cin>>e[i].u>>e[i].v;
        G[e[i].u].push_back(e[i].v);
        G[e[i].v].push_back(e[i].u);
    }
    int top=0;
    stk[++top]=1;
    int cnt=0;
    while(top)
    {
        int u=stk[top--];
        dfn[u]=++cnt;
        st[cnt][0]=u;
        for(auto v:G[u])if(v!=fa[u])
        {
            fa[v]=u;
            dep[v]=dep[u]+1;
            stk[++top]=v;
        }
    }
    for(int i=2;i<=n;i++)lg[i]=lg[i>>1]+1;
    for(int j=1;j<K;j++)for(int i=1;i+(1<<j)-1<=n;i++)st[i][j]=mn(st[i][j-1],st[i+(1<<(j-1))][j-1]);
    for(int i=1;i<=n;i++)G[i].clear();
    fac[0]=1;
    for(int i=1;i<=n;i++)fac[i]=int((ll)fac[i-1]*i%mod);
    for(int i=1;i<=n;i++)
    {
        tag[i]=i;
        t[i].siz=t[i].cnt[0]=t[i].val=1;
        t[i].md[0]=t[i].md[1]=i;
        t[i].ver=++tim;
        p[S+i-1]=make(i);
    }
    for(int i=n+1;i<=S;i++)p[S+i-1]=make(0);
    for(int i=S-1;i;i--)p[i]=mul(p[i<<1],p[i<<1|1]);
    int sum=0,m=n;
    ans_d[n]=n-1;
    ans_c[n]=m==1?1:int((ll)fac[m-2]*p[1].a[2]%mod);
    for(int i=n-1;i>=1;i--)
    {
        int u=e[i].u,v=e[i].v,x=tag[u],y=tag[v];
        if(t[x].siz>t[y].siz){swap(x,y);swap(u,v);}
        sum-=t[x].len+t[y].len;
        G[u].push_back(v);
        G[v].push_back(u);
        top=0;
        stk[++top]=u;
        tag[u]=y;
        while(top)
        {
            int z=stk[top--];
            add(y,z);
            for(auto w:G[z])if(tag[w]!=y)
            {
                tag[w]=y;
                stk[++top]=w;
            }
        }
        upd(x,0);
        upd(y,y);
        sum+=t[y].len;
        m--;
        ans_d[i]=sum+m-1;
        ans_c[i]=m==1?1:int((ll)fac[m-2]*p[1].a[2]%mod);
    }
    for(int i=1;i<=n;i++)cout<<ans_d[i]<<' '<<ans_c[i]<<'\n';
    return 0;
}