【模板】矩阵快速幂

GitHub跳转原题关系图返回列表

把普通快速幂中的乘法换成矩阵乘法,用指数二进制拆分求矩阵高次幂。

OJ: luogu

题目 ID: P3390

难度:普及/提高-

标签:数学快速幂矩阵模板题

日期: 2026-07-06 23:52

题意

给定一个 n x n 的矩阵 A 和一个非负整数 k,要求输出:

text
A^k

答案中的每个数都对 1000000007 取模。

特殊地,A^0 是单位矩阵。

思路

先看最直接的做法:从单位矩阵开始,连续乘 kA

这个做法完全按照定义计算,适合小数据验证:

cpp
// brute.cpp:小数据暴力解,直接把矩阵连乘 k 次。
#include <bits/stdc++.h>
using namespace std;

const int MAXN = 25;
const long long MOD = 1000000007LL;

int n;
long long k;

struct Matrix {
    long long a[MAXN][MAXN];
};

Matrix multiply_matrix(const Matrix &x, const Matrix &y) {
    Matrix z;
    for (int i = 1; i <= n; i++) {
        for (int j = 1; j <= n; j++) {
            z.a[i][j] = 0;
        }
    }

    for (int i = 1; i <= n; i++) {
        for (int t = 1; t <= n; t++) {
            for (int j = 1; j <= n; j++) {
                z.a[i][j] = (z.a[i][j] + x.a[i][t] * y.a[t][j]) % MOD;
            }
        }
    }
    return z;
}

Matrix identity_matrix() {
    Matrix res;
    for (int i = 1; i <= n; i++) {
        for (int j = 1; j <= n; j++) {
            res.a[i][j] = (i == j ? 1 : 0);
        }
    }
    return res;
}

int main() {
    ios::sync_with_stdio(false);
    cin.tie(nullptr);

    cin >> n >> k;
    Matrix base;
    for (int i = 1; i <= n; i++) {
        for (int j = 1; j <= n; j++) {
            long long x;
            cin >> x;
            x %= MOD;
            if (x < 0) x += MOD;
            base.a[i][j] = x;
        }
    }

    Matrix ans = identity_matrix();
    for (long long step = 1; step <= k; step++) {
        ans = multiply_matrix(ans, base);
    }

    for (int i = 1; i <= n; i++) {
        for (int j = 1; j <= n; j++) {
            if (j > 1) cout << ' ';
            cout << ans.a[i][j];
        }
        cout << "\n";
    }

    return 0;
}

k <= 10^12,不能真的乘 k 次。

普通快速幂能把 a^k 从连乘 k 次优化到 O(logk)O(log k) 次乘法。它依赖的是乘法结合律。

矩阵乘法也满足结合律:

text
(A * B) * C = A * (B * C)

所以可以把普通快速幂里的“乘法”替换成“矩阵乘法”。

例如:

指数拆分 对应矩阵
13 = 8 + 4 + 1 A^13 = A^8 * A^4 * A

我们维护两个矩阵:

  • res:当前答案,初始为单位矩阵;
  • base:当前二进制位对应的矩阵幂,初始为 A

每次看 k 的最低位:

  • 如果最低位是 1,就把 base 乘进 res
  • 然后令 base = base * base
  • 最后把 k 除以 2,继续处理下一位。

代码

cpp
#include <bits/stdc++.h>
using namespace std;

const int MAXN = 105;
const long long MOD = 1000000007LL;

int n;
long long k;

struct Matrix {
    long long a[MAXN][MAXN];
};

Matrix multiply_matrix(const Matrix &x, const Matrix &y) {
    Matrix z;
    for (int i = 1; i <= n; i++) {
        for (int j = 1; j <= n; j++) {
            z.a[i][j] = 0;
        }
    }

    for (int i = 1; i <= n; i++) {
        for (int t = 1; t <= n; t++) {
            if (x.a[i][t] == 0) continue;
            for (int j = 1; j <= n; j++) {
                z.a[i][j] = (z.a[i][j] + x.a[i][t] * y.a[t][j]) % MOD;
            }
        }
    }
    return z;
}

Matrix identity_matrix() {
    Matrix res;
    for (int i = 1; i <= n; i++) {
        for (int j = 1; j <= n; j++) {
            res.a[i][j] = (i == j ? 1 : 0);
        }
    }
    return res;
}

Matrix matrix_power(Matrix base, long long exp) {
    Matrix res = identity_matrix();

    // 和普通快速幂一样,按指数的二进制位决定是否乘当前底数。
    while (exp > 0) {
        if (exp % 2 == 1) {
            res = multiply_matrix(res, base);
        }
        base = multiply_matrix(base, base);
        exp /= 2;
    }

    return res;
}

int main() {
    ios::sync_with_stdio(false);
    cin.tie(nullptr);

    cin >> n >> k;
    Matrix base;
    for (int i = 1; i <= n; i++) {
        for (int j = 1; j <= n; j++) {
            long long x;
            cin >> x;
            x %= MOD;
            if (x < 0) x += MOD;
            base.a[i][j] = x;
        }
    }

    Matrix ans = matrix_power(base, k);

    for (int i = 1; i <= n; i++) {
        for (int j = 1; j <= n; j++) {
            if (j > 1) cout << ' ';
            cout << ans.a[i][j];
        }
        cout << "\n";
    }

    return 0;
}

复杂度

一次矩阵乘法需要 O(n3)O(n^3)

快速幂会处理 O(logk)O(log k) 个二进制位。

时间复杂度:O(n3logk)O(n^3 log k)

空间复杂度:O(n2)O(n^2)

总结

矩阵快速幂本质上不是一个全新的算法,而是普通快速幂的替换版本:

text
整数乘法 -> 矩阵乘法
整数 1   -> 单位矩阵

只要记住矩阵乘法满足结合律,就可以自然地用二进制拆分指数。