2025 ICPC武汉邀请赛 G 题思路分享(根号分治,dp)

题意

给定 \(n \times m\) 的网格,每个点有权值 \(a_{i,j}\),定义路径的价值为路径上不同 \(a_{i,j}\) 的数量,求 \((1,1)\)\((n,m)\) 所有路径的价值之和,模 \(998244353\).

\(1\le n\times m \le 10^5\).

思路

对不同权值单独处理,记下标集合为 \(S\),本质是计算经过 \(S\) 中至少一个点的路径数量.

  • \(|S|\gt \sqrt{nm}\) 时,直接暴力 \(dp\) 计算不经过 \(S\) 中任何一个点的路径数量,时间复杂度 \(\mathcal{O}(nm)\).

  • \(|S| \le \sqrt{nm}\) 时,令 \(dp_i\) 表示从 \((1,1)\)\(p_i\),且 \(p_i\) 是第一个权值为当前处理权值的点的路径数量.

    \(W(p_1,p_2)\) 为从 \(p_1\)\(p_2\) 的路径数量,初始钦定 \(dp_i = W((1,1),p_i)\).

    \(|S|\)\((x,y)\) 字典序升序排序,若 \(p_i\)\(p_j\) 可达,则 \(dp_j \leftarrow dp_j - dp_i \cdot W(p_i,p_j)\).

    总贡献为 \(\sum{dp_i \cdot W(p_i,(n,m))}\),时间复杂度 \(\mathcal{O}(|S|^2)\).

总体时间复杂度 \(\mathcal{O}(nm\sqrt{nm})\).

代码

//author:kzssCCC

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

template<int MOD>
struct modint {
    int val;
    
    modint() : val(0) {}
    modint(long long v) {
        val = v % MOD;
        if (val < 0) val += MOD;
    }
    
    modint& operator++() { val = (val + 1 == MOD ? 0 : val + 1); return *this; }
    modint& operator--() { val = (val == 0 ? MOD - 1 : val - 1); return *this; }
    modint operator++(int) { modint res = *this; ++*this; return res; }
    modint operator--(int) { modint res = *this; --*this; return res; }
    
    modint& operator+=(const modint& o) { val += o.val; if (val >= MOD) val -= MOD; return *this; }
    modint& operator-=(const modint& o) { val -= o.val; if (val < 0) val += MOD; return *this; }
    modint& operator*=(const modint& o) { val = 1LL * val * o.val % MOD; return *this; }
    modint& operator/=(const modint& o) { return *this *= o.inv(); }
    
    friend modint operator+(modint a, const modint& b) { return a += b; }
    friend modint operator-(modint a, const modint& b) { return a -= b; }
    friend modint operator*(modint a, const modint& b) { return a *= b; }
    friend modint operator/(modint a, const modint& b) { return a /= b; }

    friend bool operator==(const modint& a, const modint& b) { return a.val == b.val; }
    friend bool operator!=(const modint& a, const modint& b) { return a.val != b.val; }
    modint operator-() const { modint res = *this; res.val = (res.val == 0 ? 0 : MOD - res.val); return res;};
    modint operator+() const { return *this; };
    
    modint qpow(long long p) const {
        modint res = 1, a = *this;
        while (p > 0) {
            if (p & 1) res *= a;
            a *= a;
            p >>= 1;
        }
        return res;
    }

    modint inv() const {
        return qpow(MOD - 2);
    }
    
    friend std::ostream& operator<<(std::ostream& os, const modint& m) { return os << m.val; }
    friend std::istream& operator>>(std::istream& is, modint& m) { long long v; is >> v; m = modint(v); return is; }
};
using Z = modint<998244353>;

struct binom{
    int n;
    vector<Z> fac,inv;

    binom(int _n){
        n = _n;
        fac.assign(n+1,0);
        inv.assign(n+1,0);
        
        fac[0] = 1;
        for (int i=1;i<=n;i++){
            fac[i] = fac[i-1]*i;
        }       
        inv[n] = fac[n].inv();
        for (int i=n-1;i>=0;i--){
            inv[i] = inv[i+1]*(i+1);
        }
    }

    Z C(int a,int b){
        if (a<0 || a>n || b<0 || b>n || a<b) return 0;
        return fac[a]*inv[b]*inv[a-b];
    }

    Z A(int a,int b){
        if (a<0 || a>n || b<0 || b>n || a<b) return 0;
        return fac[a]*inv[a-b];
    }
};

binom bn(1e5);
Z cal(int x1,int y1,int x2,int y2){
    return x1<=x2&&y1<=y2?bn.C(x2-x1+y2-y1,x2-x1):0;
}

void solve(){
    int n,m;
    cin >> n >> m;
    int block = sqrt(n*m);

    vector<vector<int>> a(n+1,vector<int>(m+1));
    map<int,vector<pair<int,int>>> mp;
    for (int i=1;i<=n;i++){
        for (int j=1;j<=m;j++){
            cin >> a[i][j];
            mp[a[i][j]].emplace_back(i,j);
        }
    }

    Z res = 0;
    for (auto& [v,vec]:mp){
        if (vec.size()>block){
            vector<vector<bool>> valid(n+1,vector<bool>(m+1,true));
            for (auto& [x,y]:vec){
                valid[x][y] = false;
            }
            if (!valid[1][1] || !valid[n][m]){
                res += cal(1,1,n,m);
                continue;
            }
            vector<vector<Z>> dp(n+1,vector<Z>(m+1));
            dp[1][1] = 1;
            for (int i=1;i<=n;i++){
                for (int j=1;j<=m;j++){
                    if (i==1 && j==1 || !valid[i][j]) continue;
                    dp[i][j] = dp[i-1][j]+dp[i][j-1];
                }
            }
            res += cal(1,1,n,m)-dp[n][m];
        }
        else{
            sort(vec.begin(),vec.end());
            int N = vec.size();
            vector<Z> dp(N);
            for (int i=0;i<N;i++){
                dp[i] = cal(1,1,vec[i].first,vec[i].second);
            }

            for (int i=0;i<N;i++){
                for (int j=i+1;j<N;j++){
                    dp[j] -= dp[i]*cal(vec[i].first,vec[i].second,vec[j].first,vec[j].second);
                }
            }

            Z cur = 0;
            for (int i=0;i<N;i++){
                cur += dp[i]*cal(vec[i].first,vec[i].second,n,m);
            }
            res += cur;
        }
    }

    cout << res << '\n';
}

int main(){
    ios::sync_with_stdio(false);
    cin.tie(0);
    
    int t = 1;
    cin >> t;
    while (t--) solve();

    return 0;
}
posted @ 2026-07-21 19:11  kzssCCC  阅读(1)  评论(0)    收藏  举报