矩阵运算

利用矩阵乘法结合律先计算 K^T V,再计算按行缩放后的 Q 与该小矩阵的乘积。

OJ: shumeng

题目 ID: CSP202305B

难度:普及+/提高-

标签:线性代数矩阵乘法结合律

日期: 2026-07-31 16:21

形式化题目

给定 n×dn\times d 的矩阵 Q,K,VQ,K,V 和长度为 nn 的向量 WW,计算

(W(QKT))V,\left(W\cdot(QK^T)\right)V,

其中 WW 按行缩放 QKTQK^T(第 ii 行的每个元素都乘以 WiW_i),输出一个 n×dn\times d 矩阵。数据范围 n104, d20n\le 10^4,\ d\le 20,矩阵元素为整数。

思路

先看直接按照题目公式展开的朴素做法,它把 QKTQK^T 的每个元素显式算出来:

cpp
/**
 * Author by Rainboy blog: https://rainboylv.com github: https://github.com/rainboylvx
 * rbook: -> https://rbook.roj.ac.cn  https://rbook2.roj.ac.cn
 * rainboy的学习导航网站: https://idx.roj.ac.cn
 * create_at: 2026-07-31 16:21
 * update_at: 2026-08-17 22:40
 */
// brute.cpp:小数据暴力解,直接按公式 (W·(Q*K^T))*V 展开计算。
#include <bits/stdc++.h>
using namespace std;

int n, d;
long long q[10005][25];
long long k[10005][25];
long long v[10005][25];
long long w[10005];

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

    cin >> n >> d;
    for (int i = 0; i < n; i++)
        for (int j = 0; j < d; j++) cin >> q[i][j];
    for (int i = 0; i < n; i++)
        for (int j = 0; j < d; j++) cin >> k[i][j];
    for (int i = 0; i < n; i++)
        for (int j = 0; j < d; j++) cin >> v[i][j];
    for (int i = 0; i < n; i++) cin >> w[i];

    // 按原始公式:答案[row][col] = sum_{other} w[row] * (Q·K^T)[row][other] * V[other][col]
    for (int row = 0; row < n; row++) {
        for (int col = 0; col < d; col++) {
            long long answer = 0;
            for (int other = 0; other < n; other++) {
                long long dot = 0; // Q 第 row 行与 K 第 other 行的内积
                for (int x = 0; x < d; x++) {
                    dot += q[row][x] * k[other][x];
                }
                answer += w[row] * dot * v[other][col];
            }
            cout << answer << (col + 1 == d ? '\n' : ' ');
        }
    }

    return 0;
}

直接计算 QKTQK^T 需要 O(n2d)O(n^2d),当 n=104n=10^4 时约为 2×1092\times 10^9 次乘法,不可行。

利用矩阵乘法结合律

矩阵乘法满足结合律,可以调整括号顺序:

(W(QKT))V=(WQ)(KTV)\left(W\cdot(QK^T)\right)V=(W\cdot Q)(K^TV)。

观察维数:KTVK^TVd×dd\times d 的小矩阵,而 (WQ)(W\cdot Q) 仍是 n×dn\times d。中间矩阵从 n×nn\times n 缩小为 d×dd\times d,运算量大大下降。

计算步骤

  1. 先算 B=KTVB=K^TV,其中
    Bi,j=r=1nKr,iVr,jB_{i,j}=\sum_{r=1}^{n}K_{r,i}V_{r,j}。
  2. 对答案第 rr 行,把 QrQ_r 先乘上 WrW_r,再与 BB 相乘:
    Ar,j=i=1dWrQr,iBi,jA_{r,j}=\sum_{i=1}^{d}W_rQ_{r,i}B_{i,j}。

实现时统一使用 long long:中间与最终结果的数量级都可能远超 int 范围。

代码

cpp
/**
 * Author by Rainboy blog: https://rainboylv.com github: https://github.com/rainboylvx
 * rbook: -> https://rbook.roj.ac.cn  https://rbook2.roj.ac.cn
 * rainboy的学习导航网站: https://idx.roj.ac.cn
 * create_at: 2026-07-31 16:21
 * update_at: 2026-08-17 22:40
 */
#include <bits/stdc++.h>
using namespace std;

int n, d;
long long q[10005][25]; // Q 矩阵
long long k[10005][25]; // K 矩阵
long long v[10005][25]; // V 矩阵
long long w[10005];     // 权重向量 W
long long middle[25][25]; // K^T * V 的结果矩阵

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

    cin >> n >> d;
    for (int i = 0; i < n; i++)
        for (int j = 0; j < d; j++) cin >> q[i][j];
    for (int i = 0; i < n; i++)
        for (int j = 0; j < d; j++) cin >> k[i][j];
    for (int i = 0; i < n; i++)
        for (int j = 0; j < d; j++) cin >> v[i][j];
    for (int i = 0; i < n; i++) cin >> w[i];

    // 利用结合律 (W·Q)*K^T*V = (W·Q)*(K^T*V)
    // 先计算 K^T * V,结果只有 d*d 个元素,避免 n*n 的中间矩阵
    for (int i = 0; i < d; i++) {
        for (int j = 0; j < d; j++) {
            long long sum = 0;
            for (int row = 0; row < n; row++) {
                sum += k[row][i] * v[row][j];
            }
            middle[i][j] = sum;
        }
    }

    // W 只按行缩放 Q,再与 d*d 矩阵相乘得到最终结果
    for (int row = 0; row < n; row++) {
        for (int j = 0; j < d; j++) {
            long long answer = 0;
            for (int i = 0; i < d; i++) {
                answer += w[row] * q[row][i] * middle[i][j];
            }
            cout << answer << (j + 1 == d ? '\n' : ' ');
        }
    }

    return 0;
}

复杂度

计算 KTVK^TV 与最终结果都需要 O(nd2)O(nd^2) 时间;存储三个输入矩阵和 d×dd\times d 中间矩阵需要 O(nd+d2)O(nd+d^2) 空间。

总结

矩阵乘法不满足交换律,但满足结合律。面对形如 ABCABC 的乘积,应优先选择中间结果规模更小的括号化顺序;本题通过先算 KTVK^TVn×nn\times n 的中间矩阵变成了 d×dd\times d