矩阵快速幂算法
极难2矩阵快速幂:让计算飞起来的魔法
你学会了普通数的快速幂(比如快速计算 2 的 100 次方),也学会了矩阵乘法(两个矩阵相乘得到新矩阵)。现在把它们俩结合起来,就能得到 矩阵快速幂 —— 一种高效计算矩阵的 n 次方(即 n 个相同的矩阵相乘)的方法。这在很多递推问题(比如求斐波那契数列第 n 项)、图论问题(比如计算图中长度为 k 的路径条数)中超级有用。简单说,就是给矩阵装上了“加速器”,让原本需要 n 次乘法的事情,变成 log n 次。
为什么要学矩阵快速幂?—— 生活中的两个类比
类比1:复利计算中的“加速”
假设你存了一笔钱,年利率 5%,每年末把利息加入本金(复利)。你想知道 100 年后的总资产。如果一年一年算,要算 100 次乘法。但是,如果你用一个矩阵来表示“一年后的状态变化”:
[ 本金+利息 ] = [ 1.05 0 ] * [ 本金 ]
[ 利息累计 ] [ 0.05 1 ] [ 利息? ] (简化模型)
那么 100 年后的状态,就等于这个矩阵的 100 次方乘以初始状态。矩阵快速幂让你只用大约 log₂(100) ≈ 7 次矩阵乘法就能得到结果,而不是 100 次。
类比2:同学之间的消息传递
你在班级里,想让一个消息传遍全班。每个人可以同时告诉多个同学。如果用一个矩阵表示“谁可以告诉谁”(1 表示可以,0 表示不可以),那么矩阵的 n 次方就表示“通过 n 步传递,谁能传到谁”。比如矩阵的 2 次方就表示“经过 2 次传递后”的路径数。现实中的社交网络、朋友推荐、网页排名都会用到这种思想。
数学原理:把快速幂思想搬到矩阵上
矩阵的幂运算和数的幂运算非常像,只不过底数变成了一个 方阵(只有行数等于列数的矩阵才能不断自乘)。求矩阵 M 的 n 次方(记为 Mⁿ),就是把 M 自己乘自己 n 次。
不过注意:矩阵乘法没有交换律(A×B 一般不等于 B×A),但 结合律 依然成立,比如 (A×B)×C = A×(B×C)。而快速幂的二进制分解思想只依赖结合律,所以完全可以用!我们只需要把普通快速幂中的:
- 底数 → 矩阵
- 乘法 → 矩阵乘法
- 初始结果 (1) → 单位矩阵 I
单位矩阵 I 就像数字 1 一样:任何矩阵乘以 I(或 I 乘以任何矩阵)都等于原矩阵。单位矩阵是一个对角线全是 1、其他位置全是 0 的方阵。例如 2×2 的单位矩阵:
[ 1 0 ]
[ 0 1 ]
为什么初始化结果为单位矩阵?因为最终结果 Mⁿ 是通过不断乘以 M 的若干次幂累积得到的,初始结果必须是“乘法单位元”,这样第一次乘 M 时得到 M,第二次乘 M² 时得到 M³…… 最后才能得到 Mⁿ。
算法步骤详解(和普通快速幂对比)
普通快速幂(求 xⁿ)的步骤:
- 初始化 res = 1
- base = x
- while n > 0:
- 如果 n 是奇数(n&1),则 res = res * base
- base = base * base
- n >>= 1(右移一位,相当于整除 2)
- 返回 res
矩阵快速幂完全一样,只是把:
- 数字
1换成单位矩阵 I - 数字乘法换成
矩阵乘法
举个例子:求 M = [[1,1],[1,0]] 的 10 次方(斐波那契矩阵)
这个矩阵的特点是:Mⁿ 的结果会包含斐波那契数。例如 M¹⁰ 的左上角就是斐波那契第 11 项(F11=89),右上角是 F10=55,左下角是 F10=55,右下角是 F9=34。
用二进制分解:10 的二进制是 1010(即 8+2)。我们只需要把 M² 乘一次,M⁸ 乘一次,而 M⁴ 不需要乘,因为二进制第 2 位是 0。整个过程:
- 初始化 res = I
- 第一轮:n=10(偶数),不乘;base 平方为 M²;n>>=1 → 5
- 第二轮:n=5(奇数),res *= base(此时 base = M²),得到 res = M²;base 平方为 M⁴;n>>=1 → 2
- 第三轮:n=2(偶数),不乘;base 平方为 M⁸;n>>=1 → 1
- 第四轮:n=1(奇数),res *= base(此时 base = M⁸),得到 res = M² × M⁸ = M¹⁰;n>>=1 → 0
- 结束,返回 res
只用了 4 次循环,而直接乘需要 9 次矩阵乘法。
代码实现
为了方便,我们假设矩阵是方阵,用二维数组(C++ 用 vector<vector<long long>>,Python 用 list of list)表示。下面给出完整实现,并加上详细注释。
C++ 代码(含完整测试)
#include <iostream>
#include <vector>
using namespace std;
// 矩阵乘法(两个 n×n 方阵相乘)
vector<vector<long long>> matMul(const vector<vector<long long>>& A,
const vector<vector<long long>>& B) {
int n = A.size(); // 假设A和B都是n×n
vector<vector<long long>> C(n, vector<long long>(n, 0));
for (int i = 0; i < n; i++) {
for (int j = 0; j < n; j++) {
for (int k = 0; k < n; k++) {
C[i][j] += A[i][k] * B[k][j];
}
}
}
return C;
}
// 矩阵快速幂:计算 M^n (M是方阵)
vector<vector<long long>> matPow(vector<vector<long long>> M, long long n) {
int size = M.size(); // 矩阵大小
// 初始化结果矩阵为单位矩阵 I
vector<vector<long long>> res(size, vector<long long>(size, 0));
for (int i = 0; i < size; i++) res[i][i] = 1;
while (n > 0) {
if (n & 1LL) { // 如果 n 的当前二进制位是1
res = matMul(res, M); // 结果乘当前底数
}
M = matMul(M, M); // 底数平方
n >>= 1; // n 右移一位
}
return res;
}
// 打印矩阵
void printMat(const vector<vector<long long>>& M) {
for (auto& row : M) {
for (auto val : row) cout << val << " ";
cout << endl;
}
}
int main() {
// 使用斐波那契矩阵 [[1,1],[1,0]]
vector<vector<long long>> M = {{1, 1}, {1, 0}};
long long n = 10;
auto ans = matPow(M, n);
cout << "M^10 = " << endl;
printMat(ans);
// 输出结果:
// 89 55
// 55 34
// 验证:F11=89, F10=55, 正确
return 0;
}
Python 代码(含完整测试)
def mat_mul(A: list, B: list) -> list:
"""方阵乘法,返回 A×B"""
n = len(A)
C = [[0] * n for _ in range(n)]
for i in range(n):
for j in range(n):
total = 0
for k in range(n):
total += A[i][k] * B[k][j]
C[i][j] = total
return C
def mat_pow(M: list, n: int) -> list:
"""矩阵快速幂,返回 M^n"""
size = len(M)
# 生成单位矩阵:对角线为1,其余为0
res = [[1 if i == j else 0 for j in range(size)] for i in range(size)]
base = M # 底数
while n > 0:
if n & 1: # 如果当前位是1
res = mat_mul(res, base)
base = mat_mul(base, base) # 平方
n >>= 1 # 右移一位
return res
# 测试
if __name__ == "__main__":
M = [[1, 1], [1, 0]]
n = 10
ans = mat_pow(M, n)
for row in ans:
print(row) # 输出 [89, 55] 和 [55, 34]
新手最容易犯的 3 个错误
错误1:矩阵乘法写错了维度或顺序
- 普通数的乘法交换律 让我们觉得
res = res * base和base * res一样,但矩阵乘法不满足交换律!一定要保持顺序:res = base * res还是res = res * base取决于你的矩阵定义。通常我们按照“左边乘右边”的顺序:res = res * base表示结果不断左乘底数,这样最终得到 Mⁿ。搞反了可能会得到不同结果。 - 特别提醒:在快速幂循环中,
res和base的乘法顺序不能颠倒。建议固定为res = res * base。
错误2:忘记把结果初始化为单位矩阵
- 如果初始化成零矩阵,那么
res * base永远是零矩阵,最后得到零矩阵,答案全错。单位矩阵 I 是矩阵乘法中的“1”,绝对不能省。
错误3:没考虑 n=0 的情况
- 任何矩阵的 0 次方定义为单位矩阵。我们的算法中,当 n=0 时 while 循环不会进入,直接返回单位矩阵,这正好是正确的。如果你自己手动写递归版本,要特别注意处理 n=0 的边界。
一个隐藏错误:乘法溢出
- 矩阵元素可能变得很大(比如斐波那契数增长很快),如果题目要求取模(比如模 1e9+7),要在矩阵乘法每一步中取模,否则 long long 也可能溢出。
完整可运行的示例:快速求斐波那契数列第 n 项
我们知道,斐波那契数列可以这样用矩阵表示:
[ F(n) ] = [ 1 1 ] ^ (n-1) * [ F(1) ] = [ 1 1 ] ^ (n-1) * [ 1 ]
[ F(n-1) ] [ 1 0 ] [ F(0) ] [ 1 0 ] [ 0 ]
所以求第 n 项,只需计算矩阵的 (n-1) 次方,然后取左上角元素。下面给一个完整 Python 程序,包含取模运算:
MOD = 10**9 + 7
def mat_mul(A, B):
"""矩阵乘法,带取模"""
n = len(A)
C = [[0]*n for _ in range(n)]
for i in range(n):
for j in range(n):
total = 0
for k in range(n):
total = (total + A[i][k] * B[k][j]) % MOD
C[i][j] = total
return C
def mat_pow(M, n):
"""矩阵快速幂,带取模"""
size = len(M)
res = [[1 if i==j else 0 for j in range(size)] for i in range(size)]
base = M
while n:
if n & 1:
res = mat_mul(res, base)
base = mat_mul(base, base)
n >>= 1
return res
def fib(n):
"""返回斐波那契数列第 n 项(F1=1, F2=1),n>=1"""
if n <= 2:
return 1
M = [[1,1],[1,0]]
# 需要 M^(n-1)
power = mat_pow(M, n-1)
return power[0][0] # 左上角就是 F(n)
# 测试
print(fib(10)) # 55
print(fib(100)) # 354224848179261915075 但这里取模了
# 实际模 MOD 后:687995182
(注意:上面没有取模,如果想看取模结果,可以在 fib 函数中对结果取模,或者直接在 mat_mul 中取模)
练习
- 计算矩阵
M = [[2,0],[0,3]]的 5 次方,并验证结果(对角矩阵的幂就是对角元素的幂)。 - 尝试用矩阵快速幂求递推数列
a(n) = 2*a(n-1) + 3*a(n-2)的第 n 项(提示:构造 2×2 矩阵,类似于斐波那契)。 - 如果矩阵大小是 100×100,n 是 10^9,估算一下需要多少次矩阵乘法?(100^3 × log2(10^9) ≈ 1e6×30 = 3e7,可能很大,实际中通常用稀疏矩阵优化。)
总结
矩阵快速幂把普通快速幂的“乘法和平方”替换成“矩阵乘法和矩阵平方”,再配合单位矩阵,就能高效计算矩阵的任意次幂。它是处理线性递推问题(斐波那契、线性递推数列)、图论问题(路径计数、可达性)的超级工具。学好了它,你就掌握了一个在竞赛和实际编程中非常实用的“加速器”。
相关知识点指引
- 快速幂算法:理解二进制分解的思想(一切的起点)。
- 矩阵乘法:掌握矩阵相乘的方法和性质(结合律、不交换)。
- 线性递推:如何把递推式转化成矩阵乘法。
- 取模运算:大数运算中防止溢出的标准做法。
- 图论中的邻接矩阵:矩阵的幂表示路径数量。
现在,你可以尝试用矩阵快速幂去解决一些有趣的问题了,比如计算斐波那契数列的第 10^18 项(取模),或者计算一个社交网络中 K 步内能认识多少人。
例题精讲
使用矩阵快速幂计算一个大小为 n×n 的矩阵的 k 次幂(k 较大),假设矩阵乘法采用直接定义 O(n³) 的算法,则总的时间复杂度为:
矩阵快速幂算法要求参与运算的矩阵必须是方阵。
下列哪个问题不能直接通过矩阵快速幂高效求解?
若矩阵 A 不可逆(奇异),则使用矩阵快速幂计算 A^n 时,结果矩阵在 n 足够大时一定会变为零矩阵。
下面是矩阵快速幂的迭代实现(伪代码),请补全缺失的语句,使指数 n 正确递减。
function matPow(mat, n):
// mat 为方阵,返回 mat^n
size = len(mat)
res = identityMatrix(size) // 单位矩阵
while n > 0:
if n % 2 == 1:
res = mul(res, mat)
mat = mul(mat, mat)
n = ___ // 填空
return res