利用矩阵乘法结合律先计算 K^T V,再计算按行缩放后的 Q 与该小矩阵的乘积。
OJ: shumeng
题目 ID: CSP202305B
难度:普及+/提高-
标签:线性代数矩阵乘法结合律
日期: 2026-07-31 16:21
形式化题目
给定
其中
思路
先看直接按照题目公式展开的朴素做法,它把
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;
}直接计算
利用矩阵乘法结合律
矩阵乘法满足结合律,可以调整括号顺序:
观察维数:
计算步骤
- 先算
,其中 - 对答案第
行,把 先乘上 ,再与 相乘:
实现时统一使用 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;
}复杂度
计算
总结
矩阵乘法不满足交换律,但满足结合律。面对形如