题解:P15265 [USACO26JAN2] Dynamic Instability P

· · 题解

前言:

感谢 Rain_chr 大神教我一个比当下题解可能更加自然的做法。

思路:

首先考虑 y=1 怎么做,发现我们完全可以用上树上随机游走的那一套来干这件事情,首先先列一下随机游走的标准式子。

设 f_x 表示到达点 1 的期望步数,那么根据定义可以写出以下式子,令 son_x 表示儿子个数,bro(x) 表示 x 的兄弟节点。

可以用树上随机游走的那套从下往上推,只不过此时我们无法表示成 f_x = k_x f_{fa_x} +b_x,但是注意到我们向上递推的时候对于所有祖先的系数全部相同,所以我们可以表示成 f_x = k_x \sum_{u \in anc(x)} f_x + b_x 的形式,最后由于 f_1 = 0,所以我们就可以从下往上递推之后再从上往下推回来了。

好的,那这样我们就完成了 y=1 。

我们再来看 x=1 时怎么做,现在是往下走,我们不妨仿照树上随机游走的做法,设 g_x 表示从 fa_x 走到 x 的期望时间,考虑怎么递推。

此时对于 fa_x 这个点,有两种选择。

那我们不妨设 h_x 表示首次跳出这棵子树的期望步数,这个看似不好求,实际上这个就是你把除了这棵子树之外的节点标记,那么他们的期望时间为 0,然后求走到标记点的期望步数,你会发现这个其实就是我们的 b_x,因为外面的节点都为 0,所以你用上面式子带入乘上 k_x 之后,照样还是 0,那你发现 h_x 就是 b_x。

那么首次跳出子树后,我们知道跳到每个祖先概率是均等的。那么跳到其他儿子再跳回来的贡献实际上就是从跳到的节点一直往下走到 x 的贡献,发现上面点的 g_x,从上往下递推时已经得知,那我们就可以得到走其他儿子的贡献了。

令 s_x 表示 \sum_{u \in anc(x)} g_x,S_x 表示 \sum_{u \in anc(x)} s_x。

我们有:g_x = \sum_{v \in bro(x)} \frac{1}{son_{fa_x}}(b_v + \frac{S_{fa_x}}{dep_{fa_x}} +g_x) + 1,移项即可求出 g。

那么对于任意两点时,其实已经比较清楚了,流程就是先从 x 跳出 lca 在 x 方向上的儿子的子树,然后再从上一直往下跳过去到 y。

那么现在还有一个问题就是求出 x 跳出 lca 在 x 方向上的儿子的子树的期望步数,再观察一下上面求到标记点期望步数的式子,发现其实就是你已经知道了所有的 f_x,你现在要从祖先往下线性递推下来,你只需要对于每个点维护一个矩阵,然后倍增矩乘就行了。具体而言,你发现你每次递推和最后算答案要的是上面的 \sum f_x 和 f_x,还要带一个常数,算 b_x,直接递推这几个东西即可。

复杂度 O(q k^3 \log n ),其中 k=3,可以通过。

代码:

#include<bits/stdc++.h>
using namespace std;
const int MAXN=2e5+10;
const int mod=1e9+7;
int fa[MAXN],n,m;
vector<int>vec[MAXN];
int k[MAXN],b[MAXN],dep[MAXN],inv[MAXN],f[MAXN];
void add(int &x,int y){x+=y;if(x>=mod) x-=mod;}
int Sub(int x,int y){return x-y<0?x-y+mod:x-y;}
int Add(int x,int y){return x+y>=mod?x+y-mod:x+y;}
int ksm(int a,int b){
    int num=1;add(a,mod);
    while(b){
        if(b&1) num=1ll*num*a%mod;
        a=1ll*a*a%mod;b>>=1;
    }return num;
}
int S[MAXN],h[MAXN],g[MAXN],s[MAXN];
int jp[MAXN][18];
int LCA(int x,int y){
    if(dep[x]<dep[y]) swap(x,y);
    for(int i=17;i>=0;i--) 
        if(dep[jp[x][i]]>=dep[y]) x=jp[x][i];
    if(x==y) return x;
    for(int i=17;i>=0;i--){
        if(jp[x][i]^jp[y][i])
            x=jp[x][i],y=jp[y][i];
    }return jp[x][0];
}
void dfs(int now){
    jp[now][0]=fa[now];
    if(!vec[now].size()){
        k[now]=inv[dep[now]-1],b[now]=1;
        return ;
    }int K=0,B=0,cnt=0;
    for(int to:vec[now]){
        dep[to]=dep[now]+1;
        dfs(to);++cnt;
        add(K,k[to]),add(B,b[to]);
    }
    K=1ll*K*inv[cnt]%mod;
    int INV=ksm(1-K,mod-2);
    k[now]=1ll*K*INV%mod;
    b[now]=(1ll*inv[cnt]*B%mod+1)%mod*INV%mod;
}
void dfs2(int now){
    if(!vec[now].size()) return ;
    int sum=0;
    for(int to:vec[now]) add(sum,b[to]);
    int sz=vec[now].size();
    const int INV=inv[dep[now]];
    int v=1ll*(sz-1)%mod*S[now]%mod*INV%mod;
    for(int to:vec[now]){
        g[to]=(1ll*Sub(sum,b[to])+v+sz)%mod;
        S[to]=Add(S[now],1ll*dep[now]*g[to]%mod);
        s[to]=Add(s[now],g[to]);
        dfs2(to);
    }
}
struct Matrix{
    int a[3][3];
    Matrix (){memset(a,0,sizeof(a));}
    friend Matrix operator *(const Matrix &x,const Matrix &y){
        Matrix rs;
        for(int i=0;i<3;i++){
            for(int j=0;j<3;j++){
                if(!x.a[i][j]) continue;
                for(int k=0;k<3;k++){
                    add(rs.a[i][k],1ll*x.a[i][j]*y.a[j][k]%mod);
                }
            }
        }return rs;
    }
};
Matrix st[MAXN][18],I;
signed main(){
    ios::sync_with_stdio(false);
    cin.tie(NULL),cout.tie(NULL);
    cin>>n>>m;inv[1]=1;
    for(int i=0;i<3;i++) I.a[i][i]=1;
    for(int i=2;i<=n;i++){
        cin>>fa[i];
        inv[i]=1ll*(mod-mod/i)*inv[mod%i]%mod;
        vec[fa[i]].push_back(i);
    }dep[1]=1;dfs(1);
    for(int i=1;i<=n;i++){
        Matrix &mat=st[i][0];
        mat.a[0][0]=k[i]+1,mat.a[0][1]=k[i];
        mat.a[2][0]=b[i],mat.a[2][1]=b[i];
        mat.a[2][2]=1;
    }
    for(int j=1;j<=17;j++){
        for(int i=1;i<=n;i++){
            st[i][j]=st[jp[i][j-1]][j-1]*st[i][j-1];
            jp[i][j]=jp[jp[i][j-1]][j-1];
        }
    }dfs2(1);
    while(m--){
        int x,y;cin>>x>>y;
        int lca=LCA(x,y);
        int ans=Sub(s[y],s[lca]);
        if(lca!=x){
            int D=dep[x]-dep[lca];
            add(ans,1ll*S[lca]*inv[dep[lca]]%mod);
            Matrix rs=I,c;c.a[0][2]=1;
            for(int j=0;j<=17;j++){
                if(D>>j&1) rs=st[x][j]*rs,x=jp[x][j];
            }add(ans,(c*rs).a[0][1]);
        }cout<<ans<<'\n';
    }
    return 0;
}