ABC323G Inversion of Tree 题解
题目描述
给定长为 \(n\) 的排列 \(p\) ,对 \(\forall 0\le k\lt n\) ,求满足 \(u\lt v\and p_u\gt p_v\) 的树边 \((u,v)\) 数量恰好为 \(k\) 的树的个数,对 \(998244353\) 取模。
数据范围
- \(2\le n\le 500\) 。
时间限制 \(\texttt{4s}\) ,空间限制 \(\texttt{1GB}\) 。
分析
矩阵树定理,对每对 \((i,j)\) ,若 \(p_i\gt p_j\) ,令系数为 \(x\) ,否则令系数为 \(1\) ,得到拉普拉斯矩阵 \(L\) ,则 \(k\) 的答案为 \(\det\big(L\) 的 \(n-1\) 阶主子式 \(\big)\) 的 \(x^k\) 项系数。
注意到 \(L\) 中各元素均为一次多项式,问题转化为求 \(f(x)=\det(A+xB)\) ,其中 \(A,B\) 为常数矩阵。
我们希望将 \(B\) 化成 \(I\) ,这样只需求 \(-A\) 的特征多项式,但可惜 \(B\) 不一定有逆元。
对 \((A;B)\) 高斯消元,如果某个时刻 \(B\) 没有主元:
- 如果 \(A\) 也没有主元,行列式为零。
- 否则将 \(A,B\) 的第 \(j\) 列乘以 \(-b_{j,i}\) 加到第 \(i\) 列,这样 \(B\) 的第 \(i\) 列全空。然后交换 \(A,B\) 的第 \(i\) 列,这会给行列式乘上 \(x\) 。记录交换次数 \(s\) ,最后将计算结果向低平移 \(s\) 次幂即可。
接下来考虑 \(A\) 的特征多项式 \(\det(xI-A)\) 怎么求。
后面内容同时作为 P7776 【模板】特征多项式 的题解,对应代码第 \(56\sim 92\) 行。
定义:若 \(i\gt j+1\) 时 \(h_{i,j}=0\) ,则矩阵 \(H\) 被称为上 Hessenberg 矩阵。
定理:相似矩阵的特征多项式相同。
证明:设 \(B=P^{-1}AP\) ,则 \(|xI-A|=|P(xI-B)P^{-1}|=|P|\cdot |xI-B|\cdot |P^{-1}|=|xI-B|\) 。
定理:任意矩阵 \(A\) 相似于某个上 Hessenberg 矩阵 \(H\) 。
下面给出一个构造性证明。
用初等矩阵消元,记 \(P(j,i(k))=E+k\cdot E_{j,i}\) (主对角为 \(1\) ,第 \((j,i)\) 个元素为 \(k\) ,其余为 \(0\) ),则 \(P^{-1}(j,i(k))=E-k\cdot E_{j,i}=P(j,i(-k))\) 。
左乘 \(P(j,i(k))\) ,将第 \(i\) 行的 \(k\) 倍加到第 \(j\) 行。
右乘 \(P(j,i(-k))\) ,将第 \(j\) 列的 \(-k\) 倍加到第 \(i\) 列。
对 \(i=1\sim n-1\) ,先给第 \(i+1\) 行乘以 \(a_{i+1,i}^{-1}\) ,再给第 \(i+1\) 列乘以 \(a_{i+1,i}\) ,从而将 \(a_{i+1,i}\) 消成 \(1\) 。
再对 \(j\gt i+1\) ,左乘 \(P(j,(i+1)(-a_{j,i}))\) 消去 \(a_{j,i}\) ,右乘 \(P(j,(i+1)(a_{j,i}))\) 只会影响第 \(i+1\) 列,不会破坏已经消元的前 \(i\) 列结构。
这样我们将 \(A\) 转化为了上 Hessenberg 矩阵:
记 \(H_i\) 为上述矩阵的 \(i\) 阶主子式, \(f_i(x)\) 为 \(H_i\) 的特征多项式。
按第 \(n\) 行展开:
按第 \(n-1\) 列展开,第 \(j\) 行贡献 \((-1)^{j+(n-1)}h_{j,n}\cdot f_{j-1}(x)\prod_{k=j+1}^{n-2}(-\beta_k)\) 。
整理得:
据此递推即可,时间复杂度 \(\mathcal O(n^3)\) 。
#include<bits/stdc++.h>
using namespace std;
const int maxn=505,mod=998244353;
int n,r=1,s,p[maxn];
int a[maxn][maxn],b[maxn][maxn],f[maxn][maxn];
int qpow(int a,int k)
{
int res=1;
for(;k;k>>=1,a=1ll*a*a%mod) if(k&1) res=1ll*res*a%mod;
return res;
}
int main()
{
scanf("%d",&n);
for(int i=1;i<=n;i++) scanf("%d",&p[i]);
for(int i=1;i<=n;i++)
for(int j=i+1;j<=n;j++)
{
auto s=p[i]>p[j]?b:a; // det(A+xB),若 p[i]>p[j] 贡献 x, 否则贡献 1
s[i][i]++,s[j][j]++,s[i][j]--,s[j][i]--;
}
n--; // 取 n-1 阶主子式
// 将 det(A+xB) 消成 det(A+xI) 的形式
for(int i=1;i<=n;i++)
{
if(!b[i][i]) for(int j=i+1;j<=n;j++) if(b[j][i]) {r*=-1,swap(a[i],a[j]),swap(b[i],b[j]);break;}
if(!b[i][i])
{
if(!a[i][i]) for(int j=i+1;j<=n;j++) if(a[j][i]) {r*=-1,swap(a[i],a[j]),swap(b[i],b[j]);break;}
if(!a[i][i]) // A,B 都找不到主元,行列式为零
{
for(int j=0;j<=n;j++) printf("0%c"," \n"[j==n]);
return 0;
}
// 将 A,B 的第 j 列乘以 -b[j][i] 加到第 i 列,将 B 的第 i 列消空
for(int j=1;j<i;j++)
for(int k=1,t=-b[j][i];k<=n;k++)
{
a[k][i]=(a[k][i]+1ll*a[k][j]*t)%mod;
b[k][i]=(b[k][i]+1ll*b[k][j]*t)%mod;
}
s++; // 交换 A,B 的第 i 列, 行列式乘以 x
for(int k=1;k<=n;k++) swap(a[k][i],b[k][i]);
}
r=1ll*r*b[i][i]%mod;
for(int k=1,v=qpow(b[i][i],mod-2);k<=n;k++) a[i][k]=1ll*a[i][k]*v%mod,b[i][k]=1ll*b[i][k]*v%mod;
for(int j=1;j<=n;j++) if(j!=i)
for(int k=1,t=-b[j][i];k<=n;k++)
{
a[j][k]=(a[j][k]+1ll*a[i][k]*t)%mod;
b[j][k]=(b[j][k]+1ll*b[i][k]*t)%mod;
}
}
// 计算 -A 特征多项式, 先取反得到 -A
for(int i=1;i<=n;i++) for(int j=1;j<=n;j++) a[i][j]*=-1;
// 核心思想:用 a[i+1][i] 消 a[j][i] ,得到上 Hessenberg 矩阵
for(int i=1;i<n;i++)
{
if(!a[i+1][i])
for(int j=i+2;j<=n;j++) if(a[j][i])
{
// 交换第 i+1 行和第 j 行,交换第 i+1 列和第 j 列
for(int k=1;k<=n;k++) swap(a[i+1][k],a[j][k]);
for(int k=1;k<=n;k++) swap(a[k][i+1],a[k][j]);
break;
}
if(!a[i+1][i]) continue;
// 第 i+1 行乘以 a[i+1][i]^{-1},第 i+1 列乘以 a[i+1][i]
int u=qpow(a[i+1][i],mod-2),v=a[i+1][i];
for(int k=1;k<=n;k++) a[i+1][k]=1ll*a[i+1][k]*u%mod;
for(int k=1;k<=n;k++) a[k][i+1]=1ll*a[k][i+1]*v%mod;
assert(a[i+1][i]==1);
for(int j=i+2;j<=n;j++)
{
// 第 i+1 行的 -a[j][i] 倍加到第 j 行,第 j 列的 a[j][i] 倍加到第 i+1 列
int v=a[j][i];
for(int k=1;k<=n;k++) a[j][k]=(a[j][k]-1ll*a[i+1][k]*v)%mod;
for(int k=1;k<=n;k++) a[k][i+1]=(a[k][i+1]+1ll*a[k][j]*v)%mod;
}
}
// 递推计算上 Hessenberg 矩阵的特征多项式
// f_i(x)=x*f_{i-1}(x)-\sum_{j=1}^ia_{j,i}f_{j-1}(x)\prod_{k=j+1}^ia_{k,k-1}
f[0][0]=1;
for(int i=1;i<=n;i++)
{
for(int k=1;k<=i;k++) f[i][k]=f[i-1][k-1];
for(int j=i,cur=1;j>=1;j--)
{
for(int k=0,v=1ll*a[j][i]*cur%mod;k<=j;k++) f[i][k]=(f[i][k]-1ll*f[j-1][k]*v)%mod;
cur=1ll*cur*a[j][j-1]%mod;
}
}
// 最后乘上 r*x^{-s} ,得到 det(A+xB)
for(int i=0;i<=n;i++) printf("%lld%c",i<=n-s?(1ll*r*f[n][i+s]%mod+mod)%mod:0," \n"[i==n]);
return 0;
}
本文来自博客园,作者:peiwenjun,转载请注明原文链接:https://www.cnblogs.com/peiwenjun/p/20014889
浙公网安备 33010602011771号