[ABC319G] Counting Shortest Paths

· · 题解

题意:

现在又一个 n 个点的无向完全图,删除了 m 条边之后,问你从 1 号点到达 n 号点的最短路径条数是多少?

思路:

赛后第一次切 G 题诶!

看到这里,大家可能会想到 P1144 最短路计数,直接暴力删边直接做的话,可以使用 BFS 来求,但是这样的时间复杂度也是 O(n^2) 的,所以要想想怎么优化。

其实有一个挺简单的做法,定义 f_i 为到达 i 号点的最短路径条数,dis_i 为到达 i 号点的最短路径,s_i 为当前层次不能到达 i 号点的条数,sum 为当前层次所有点的最短路径数量之和。

我们从一号点开始入队(初始肯定定义 dis_1=f_1=1),然后宽搜一下,设队头为 u,对于所有删除了 u 号点的边,找到删除的边的另外一条边,将 s_v \to s_v + f_u(因为这条边被删除了,所以这个点不能直接过去,所以累加在 s_v 中)。

如果当前队为空了,我们需要加入新的点了,遍历这 n 个点,寻找每一个没有找到最短路径的并且 s_i \ne sum 的点 i,然后将 dis_i \to dis_u+1,同时将这个这个点的最短路径的条数 f_i \to (sum-s_i),最后将 i 入队。

现在想想为什么要这样?

我们知道 sum 为当前所有点的最短路径数量之和,如果 s_i=sum 的话,这说明在当前层次没有一个点可以通向 i 这个点,所以不用访问。

还有如果 dis_i 已经有值了,说明前面已经有一个最近的路径了,所以也不用访问。

更新 i 号点的答案就是当前层次所有的条数减去不能到达当前点的条数(毕竟在同一层次的点,到达 i 号点的距离都是 1,毕竟之前是一个完全图)。

因为 sum 和 s_i 都是当前层次的,即每次入队之后就将 sum=s_i=0。

最后我们的答案就是 f_n。(注意如果 f_n 为 0 就输出 -1)

因为是按照层次分的,所以时间复杂度为:O(n+m)。

新增部分:

因为有模数,所以可能 sum 不等于 g_i 但是 sum 模 mod 同于 g_i,这样会导致计算错误,其实,我们只需要讲模式改为 998244353 的倍数即可(这样不好改变答案),最后再重新模上 998244353 即可。

我这里新设置的模式是 998244353^2,这样就可以避免 after_contest 卡 998244353 的数据。

完整代码:

#include<bits/stdc++.h>
using namespace std;
typedef long long ll;
const int N=200200;
const ll mod=(998244353ll)*(998244353ll);
inline ll read(){
    ll x=0,f=1;
    char c=getchar();
    while(c<'0'||c>'9'){
        if(c=='-')
          f=-1;
        c=getchar();
    }
    while(c>='0'&&c<='9'){
        x=(x<<1)+(x<<3)+(c^48);
        c=getchar();
    }
    return x*f;
}
inline void write(ll x){
    if(x<0){
        putchar('-');
        x=-x;
    }
    if(x>9)
      write(x/10);
    putchar(x%10+'0');
}
ll n,m,sum=0;
ll dis[N],f[N],s[N];
vector<ll> E[N];
queue<ll> q;
void add(ll u,ll v){ //建边 
    E[u].push_back(v);
    E[v].push_back(u);
} 
int main(){
    n=read(),m=read(); 
    for(int u,v,i=1;i<=m;i++){
        u=read(),v=read();
        add(u,v);
    }
    q.push(1);
    dis[1]=1;
    f[1]=1;
    while(!q.empty()){
        ll u=q.front();
        q.pop();
        for(auto v:E[u]) //删除的边 
          s[v]=(s[v]+f[u])%mod;
        sum=(sum+f[u])%mod; //累加条数 
        if(q.empty()){ //增点 
            for(ll i=1;i<=n;i++){
                if(!dis[i]&&s[i]!=sum){
                    dis[i]=dis[u]+1;
                    f[i]=(sum-s[i]+mod)%mod;
                    q.push(i);
                }
                s[i]=0;
            }
            sum=0;
        }
    }
    if(!f[n])
      puts("-1"); 
    else
      write(f[n]%998244353ll);
    return 0;
}