2025 ICPC Wuhan Invitational Contest G.Path Summing Problem (根号分治+容斥+DP)

题干:有一个 \(n\) 行 \(m\) 列的网格。网格里的每个格子都写着一个整数,其中第 \(i\) 行第 \(j\) 列的格子里写着整数
\(a_{i,j}\) 。
令 \((i, j)\) 表示位于第 \(i\) 行第 \(j\) 列的格子。您现在需要从 \((1, 1)\) 出发并前往 \((n, m)\)。当您位于格子 \((i, j)\)
时,您可以选择走到右方的格子 \((i, j + 1)\)(若 \(j\) \(<\) \(m\) ),也可以选择走到下方的格子 \((i + 1, j)\)(若
\(i < n\))。
令 \(S\) 表示路径上每个格子里的整数形成的集合,包括 \(a_{1,1}\) 和 \(a_{n,m}\)。路径的价值定义为 \(S\) 中元素的数量
(请回忆:集合中不包含重复元素)。对于所有可能的路径,求它们的价值之和。
思路:考虑每个点数值为 \(x\) 对答案的贡献,当有一条路径第一次经过 \(x\) 时会对答案做出 \(1\) 的贡献。
如何统计? :我们有朴素的dp。
DP: \(dp_{i,j}\) 表示对 \(x\) 这个数值,由 \((1,1)\) 到 \((i,j)\) 的经过 \(x\) 的路径总数。
前置芝士:\((1,1)\) 到 \((n,m)\) 的路径总数为 \(\begin{pmatrix}n+m-2\\n-1\end{pmatrix}\) 。
有:
复杂度是 \(O(knm)\) , \(k\) 为不同 \(x\) 的个数。
显然复杂度是超的,我们只好寻求别的思路。
容斥:有显然的容斥:对于两个数值为 \(x\) 的点,若分别位于左上右下,合法的先经过右下的路径数= \((1,1)\) 到右下点的路径数 - 通过左上的合法的路径数。我们将这个关系推导至 \(i\) 个点,记 \(f_i\) 为通过第 \(i\) 个点的合法的路径数,有
注意这里的合法的 \(j\) 满足 \(x_j\leq x_i\) 且 \(y_j\leq y_i\)。
那么每一个 \(x\) 对答案的贡献为 \(\sum_{i=1}^{w} f_i*\begin{pmatrix}n+m-x_i-x_j\\n-x_i\end{pmatrix}~~~~~~~~(G_{i,j}=x)\)
复杂度是 \(O(kw^2)\) , \(w\) 为 \(x\) 的个数。
对比这两个算法的复杂度,我们会发现无论哪种做法都会超时,这时候不妨考虑根号分治,能让我们将复杂度卡到\(O(nm\sqrt{nm})\)。
具体就是对每一个不同的 \(w\) 选择合适的方法。
代码如下:
#include<bits/stdc++.h>
using namespace std;
using ll = long long;
using ull = unsigned long long;
using db = double;
using ldb = long double;
using pii = pair<ll,ll>;
#define CI const int
#define mp make_pair
#define int ll
CI maxn=1e5+5;
map<int,int>G[maxn];
int T,n,m,ans=0;
int v[maxn];
CI mod=998244353;
struct node{
int x,y,f;
friend bool operator<(node a,node b){
return a.x==b.x?a.y<b.y:a.x<b.x;
}
}c[maxn];
vector<node>col[maxn];
vector<int>id;
int ad(int a,int b){return (a%mod+b%mod)%mod;}
int ch(int a,int b){return (a%mod*b%mod)%mod;}
int ji(int a,int b){return (a%mod-b%mod+mod)%mod;}
void cl()
{
ans=0;id.clear();
for(int i=1;i<=n*m;i++)G[i].clear(),v[i]=0,col[i].clear();
}
ll qpow(int x,int y)
{
ll r=1;
for(;y;y>>=1)
{
if(y&1)r=ch(r,x);
x=ch(x,x);
}
return r;
}
int inv(int x){return qpow(x,mod-2);}
int fac[maxn],inf[maxn];
void init()
{
fac[0]=1;
for(int i=1;i<=1e5;i++)fac[i]=ch(fac[i-1],i);
inf[(int)1e5]=inv(fac[(int)1e5]);
for(int i=1e5-1;i>=0;i--)inf[i]=ch(inf[i+1],i+1);
}
int C(int y,int x)
{
if(y==0)return 1;
return ch(fac[x],ch(inf[y],inf[x-y]));
}
void solve1(int k)
{
sort(col[k].begin(),col[k].end());
for(auto &p:col[k])
{
int xi=p.x,yi=p.y;
for(auto &p2:col[k])
{
int xj=p2.x,yj=p2.y,sum=p2.f;
if(xi<xj||yi<yj)continue;
if(xi==xj&&yi==yj)continue;
p.f=ji(p.f,ch(sum,C(xi-xj,xi+yi-xj-yj)));
}
}
for(auto p:col[k])
{
ans=ad(ans,ch(p.f,C(n-p.x,n+m-p.x-p.y)));
}
}
void solve2(int k)
{
vector<vector<int>>dp(n+1,vector<int>(m+1,0));
for(int i=1;i<=n;i++)
{
for(int j=1;j<=m;j++)
{
if(G[i][j]==k)dp[i][j]=C(i-1,i+j-2);
else dp[i][j]=ad(dp[i-1][j],dp[i][j-1]);
}
}
ans=ad(ans,dp[n][m]);
}
signed main()
{
ios::sync_with_stdio(false);
cin.tie(nullptr);
cin>>T;
init();
while(T--)
{
cin>>n>>m;
for(int i=1;i<=n;i++)
for(int j=1;j<=m;j++)
{
cin>>G[i][j];
if(!v[G[i][j]])id.push_back(G[i][j]);
v[G[i][j]]++;
col[G[i][j]].push_back(node{i,j,C(i-1,i+j-2)});
}
int p=sqrt(n*m);
for(auto k:id)
{
if(v[k]<=p)solve1(k);
else solve2(k);
}
cout<<ans<<'\n';
cl();
}
return 0;
}

浙公网安备 33010602011771号