矩阵快速幂模板P3390
P3390【模板】矩阵快速幂(详解 + 代码)
题目背景
一个 \(m \times n\) 的矩阵是一个整数构成的矩形阵列,形如:
\[A =
\begin{pmatrix}
a_{11} & a_{12} & \cdots & a_{1n} \\
a_{21} & a_{22} & \cdots & a_{2n} \\
\vdots & \vdots & & \vdots \\
a_{m1} & a_{m2} & \cdots & a_{mn}
\end{pmatrix}
\]
矩阵乘法
若矩阵 \(A\) 的大小为 \(m \times n\),\(B\) 的大小为 \(n \times p\),则乘积 \(C = AB\) 是一个 \(m \times p\) 的矩阵:
\[c_{ij} = \sum_{k=1}^{n} a_{ik} b_{kj}
\quad (1 \le i \le m,\; 1 \le j \le p)
\]
矩阵乘法满足结合律:
\[(AB)C = A(BC)
\]
矩阵快速幂
对于 \(n \times n\) 的方阵 \(A\),定义其幂次为:
- \(A^1 = A\)
- \(A^k = A \times A^{k-1}\)
- \(A^0 = I\)(单位矩阵)
单位矩阵 \(I\) 形如:
\[I =
\begin{pmatrix}
1 & 0 & \cdots & 0 \\
0 & 1 & \cdots & 0 \\
\vdots & \vdots & \ddots & \vdots \\
0 & 0 & \cdots & 1
\end{pmatrix}
\]
题目描述
给定一个 \(n \times n\) 的矩阵 \(A\),求 \(A^k\)。
输入格式
第一行两个整数 \(n, k\)。
接下来 \(n\) 行,每行 \(n\) 个整数,表示矩阵 \(A\)。
输出格式
输出 \(A^k\)。
共 \(n\) 行,每行 \(n\) 个整数,所有结果对 \(10^9 + 7\) 取模。
输入输出样例
样例 1
输入
2 1
1 1
1 1
输出
1 1
1 1
样例 2
输入:
3 5
1 2 3
4 5 6
7 8 9
输出
121824 149688 177552
275886 338985 402084
429948 528282 626616
数据范围
- \(1 \le n \le 100\)
- \(0 \le k \le 10^{12}\)
- \(|A_{ij}| \le 1000\)
解题思路
由于 \(k\) 最大可达 \(10^{12}\),直接循环乘 \(k\) 次显然会超时。
这里需要使用 快速幂(二进制拆分) 的思想:
- 如果 \(k\) 是奇数:\(A^k = A \times A^{k-1}\)
- 如果 \(k\) 是偶数:\(A^k = (A^{k/2})^2\)
结合矩阵乘法,时间复杂度优化为 \(O(n^3 \log k)\)。
实现代码:
点击查看代码
#include <bits/stdc++.h>
using namespace std;
#define int long long
const int MOD=1e9+7;
struct matrix{
int n;
int a[105][105];
};
matrix mul(matrix A,matrix B)//矩阵乘法
{
matrix C;
C.n=A.n;
for(int i=1;i<=C.n;i++)
{
for(int j=1;j<=C.n;j++)
{
C.a[i][j]=0;
}
}//初始化全为0
for(int i=1;i<=C.n;i++)
{
for(int k=1;k<=C.n;k++)
{
for(int j=1;j<=C.n;j++)
{
C.a[i][j]=(C.a[i][j]+A.a[i][k]*B.a[k][j])%MOD;
}
}
}
return C;
}
matrix fast_pow(matrix A,int n)
{
matrix R;
R.n=A.n;
for(int i=1;i<=R.n;i++)
{
for(int j=1;j<=R.n;j++)
{
R.a[i][j]=(i==j)?1:0;
}
}
while(n>0)
{
if(n&1)
{R=mul(R,A);}
A=mul(A,A);
n>>=1;
}
return R;
}
signed main()
{
ios::sync_with_stdio(0);
cin.tie(0);
cout.tie(0);
int n,k;
cin>>n>>k;
matrix A;
A.n=n;
for(int i=1;i<=n;i++)
{
for(int j=1;j<=n;j++)
{
cin>>A.a[i][j];
A.a[i][j]%=MOD;
if(A.a[i][j]<0)
A.a[i][j]+=MOD;
}
}
matrix R;
R=fast_pow(A,k);
for(int i=1;i<=n;i++)
{
for(int j=1;j<=n;j++)
{
if(j!=1) cout<<' ';
cout<<R.a[i][j];
}
cout<<endl;
}
system("pause");
return 0;
}

浙公网安备 33010602011771号