题解:P17323 [ICPC 2018 Nanjing R] Frank

· · 题解

题意简述

给定一张强连通有向多重图。每到一个城市,等概率选择一条出边。每次询问给出城市序列,求它成为访问序列的子序列前,经过边数的期望。

解题思路

H_{i,j} 为从 i 出发首次到达 j 的期望步数。每匹配一个城市后,之后的随机游走与此前路径无关。因此,序列 c_0,c_1,\dots,c_{k-1} 的答案为:

\sum_{i=0}^{k-2}H_{c_i,c_{i+1}}

只需预处理询问中作为终点出现的 H_{i,j}

先考虑固定终点 t。设 x_i=H_{i,t},转移概率为 p_{i,j}。到达 t 时过程结束,所以 x_t=0。对于 i\ne t,首步分析得到:

(1-p_{i,i})x_i=1+\sum_{j\ne i}p_{i,j}x_j

重边通过 p_{i,j} 合并,自环已经在等式两侧抵消。

维护当前未消元点集 S。每个 i\in S 的方程写作:

d_ix_i=b_i+\sum_{j\in S\setminus\{i\}}w_{i,j}x_j

初始时 b_i=1w_{i,j}=p_{i,j}。矩阵每行的系数和为 0,所以始终有:

d_i=\sum_{j\in S\setminus\{i\}}w_{i,j}

消去点 p 时,先计算:

v=\sum_{j\in S\setminus\{p\}}w_{p,j}

强连通保证 v>0。将第 p 行除以 v,得到用于回代的关系:

x_p=b_p+\sum_{j\in S\setminus\{p\}}w_{p,j}x_j

这里的 b_p,w_{p,j} 均指除以 v 后的值。对于其余点 u,令 a=w_{u,p},代入上式并更新:

\begin{aligned} b_u & \leftarrow b_u+ab_p \\ w_{u,j} & \leftarrow w_{u,j}+aw_{p,j} \end{aligned}

第二个式子只处理 j\ne u,p。产生的 x_u 项并入等式左侧。消元结束后令 x_t=0,按照相反顺序回代,即可得到所有 H_{i,t}

直接对每个终点独立消元需要 O(n^4)。设询问使用的不同终点构成集合 S。不在 S 中的点对所有终点都能直接消去。此后将 S 分成两半:求左半终点时,先统一消去右半;求右半终点时,恢复现场并统一消去左半。递归到单个终点后再回代。

设当前集合大小为 s,递归两支的总消元量满足 T(s)=2T(s/2)+O(s^3),所以总复杂度仍为 O(n^3)

该消元过程只进行非负数的加法、乘法和除法,不会出现两个巨大特解相减的消减误差。任意两点间存在长度至多为 n-1 的简单路径,每条存在的边被选择的概率至少为 1/m。因此:

H_{i,t}\le(n-1)m^{n-1}<10^{2238}

需要保留的非零路径概率也大于 10^{-2236}。x86 扩展精度 long double 的指数范围足够。其他平台未必提供相同语义,且扩展精度运算较慢。

代码使用结构体 num 显式保存二进制尾数与指数。非零数表示为 x\cdot2^e,其中 2^{63}\le x<2^{64}。加法先对齐指数。乘除法使用 __uint128_t 保存中间结果,再将尾数归一化至 64 位。指数使用 int,远大于上述范围所需。

所有运算均为非负数,相加时不会发生消减。每次运算的相对截断误差小于 2^{-63}。误差沿依赖链至多累计 O(n^2) 层。因此,精度足以满足误差限制。

总时间复杂度为 O(n^3+m+\sum k),空间复杂度为 O(n^2\log n+\sum k)

参考代码

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

using ull=unsigned long long;
using u128=__uint128_t;
struct num
{
    ull x;
    int e;
    num(double v=0)
    {
        if(v==0)
        {
            x=e=0;
            return;
        }
        int d;
        double z=frexp(v,&d);
        x=(ull)ldexp(z,64);
        e=d-64;
    }
    num(ull v,int d):x(v),e(d){}
    num operator+(num b) const
    {
        num a=*this;
        if(a.x==0)return b;
        if(b.x==0)return a;
        if(a.e<b.e)swap(a,b);
        int d=a.e-b.e;
        if(d>=64)return a;
        u128 z=(u128)a.x+(b.x>>d);
        if(z>>64)
        {
            a.x=(ull)(z>>1);
            a.e++;
        }
        else a.x=(ull)z;
        return a;
    }
    num operator*(num b) const
    {
        if(x==0||b.x==0)return 0;
        u128 z=(u128)x*b.x;
        if(z>>127)return num((ull)(z>>64),e+b.e+64);
        return num((ull)(z>>63),e+b.e+63);
    }
    num operator/(num b) const
    {
        if(x==0)return 0;
        if(x>=b.x)return num((ull)(((u128)x<<63)/b.x),e-b.e-63);
        return num((ull)(((u128)x<<64)/b.x),e-b.e-64);
    }
    num &operator+=(num b)
    {
        return *this=*this+b;
    }
    num &operator*=(num b)
    {
        return *this=*this*b;
    }
    explicit operator bool() const
    {
        return x!=0;
    }
    double val() const
    {
        return ldexp((double)x,e);
    }
};
const int N=405;
const int Q=405;
const int K=505;
const int D=12;
int n,deg[N],cnt[N][N],id[N],ord[N],c[Q][K],len[Q],tot;
num f[N][N],b[N],h[N][N],bak[D][N][N],bb[D][N];
bool use[N];
void del(int l,int r,int p)
{
    num v=0;
    for(int i=l;i<r;i++)v+=f[p][id[i]];
    v=num(1)/v;
    b[p]*=v;
    for(int i=l;i<r;i++)f[p][id[i]]*=v;
    for(int i=l;i<r;i++)
    {
        int u=id[i];
        v=f[u][p];
        if(!v)continue;
        b[u]+=v*b[p];
        for(int j=l;j<r;j++)
        {
            int w=id[j];
            f[u][w]+=v*f[p][w];
        }
        f[u][u]=0;
        f[u][p]=0;
    }
    ord[tot++]=p;
}
void elim(int l,int r,int x,int y)
{
    if(x==l)
    {
        for(int i=x;i<y;i++)del(i+1,r,id[i]);
    }
    else
    {
        for(int i=y-1;i>=x;i--)del(l,i,id[i]);
    }
}
void dfs(int d,int l,int r)
{
    if(r-l==1)
    {
        int t=id[l];
        for(int i=0;i<n;i++)h[i][t]=0;
        for(int i=tot-1;i>=0;i--)
        {
            int p=ord[i];
            h[p][t]=b[p];
            for(int j=0;j<n;j++)if(j!=p&&f[p][j])h[p][t]+=f[p][j]*h[j][t];
        }
        return;
    }
    int mid=(l+r)/2;
    int cur=tot;
    for(int i=l;i<r;i++)
    {
        int u=id[i];
        bb[d][u]=b[u];
        for(int j=l;j<r;j++)bak[d][u][id[j]]=f[u][id[j]];
    }
    elim(l,r,mid,r);
    dfs(d+1,l,mid);
    tot=cur;
    for(int i=l;i<r;i++)
    {
        int u=id[i];
        b[u]=bb[d][u];
        for(int j=l;j<r;j++)f[u][id[j]]=bak[d][u][id[j]];
    }
    elim(l,r,l,mid);
    dfs(d+1,mid,r);
    tot=cur;
}
int main()
{
    ios::sync_with_stdio(false);
    cin.tie(nullptr);
    int m,q;
    cin>>n>>m>>q;
    for(int i=1;i<=m;i++)
    {
        int u,v;
        cin>>u>>v;
        cnt[u][v]++;
        deg[u]++;
    }
    for(int i=1;i<=q;i++)
    {
        cin>>len[i];
        for(int j=1;j<=len[i];j++)
        {
            cin>>c[i][j];
            if(j>1)use[c[i][j]]=1;
        }
    }
    for(int i=0;i<n;i++)
    {
        b[i]=1;
        for(int j=0;j<n;j++)
        {
            if(i!=j)f[i][j]=num(cnt[i][j])/num(deg[i]);
        }
    }
    int s=0;
    for(int i=0;i<n;i++)if(!use[i])id[s++]=i;
    int l=s;
    for(int i=0;i<n;i++)if(use[i])id[s++]=i;
    elim(0,n,0,l);
    dfs(0,l,n);
    cout<<fixed<<setprecision(12);
    for(int i=1;i<=q;i++)
    {
        num ans=0;
        for(int j=2;j<=len[i];j++)ans+=h[c[i][j-1]][c[i][j]];
        cout<<ans.val()<<'\n';
    }
    return 0;
}