矩阵乘法分块优化无性能提升的原因探究
矩阵乘法优化性能未达预期求助
我研究了Creel的《矩阵乘法优化》视频,但无法复现视频中的性能提升,特此询问原因。
测试环境与现象
- 测试程序包含三个核心函数:朴素矩阵乘法、B矩阵原地转置后的乘法、B矩阵原地转置加分块的乘法
- 测试配置:矩阵大小
n=4000,块大小依次尝试1、10、20、50、100、200 - 硬件缓存:32KB L1D、256KB L2、4MB共享L3。按计算,块大小10时内存占用约6.4KB,完全可以放入L1缓存
- 实际结果:所有块大小的测试耗时均为50秒,和仅转置的乘法版本耗时完全一致
- 编译命令:
gcc -O3 -mavx2
测试代码
#include <stdlib.h> #include <stdio.h> #include <time.h> void matmul(size_t n, double A[n][n], double B[n][n], double result[n][n]) { for (size_t i = 0; i < n; i++) { for (size_t j = 0; j < n; j++) { double acc = 0; for (size_t k = 0; k < n; k++) { acc += A[i][k] * B[k][j]; } result[i][j] = acc; } } } void transpose(size_t n, double matrix[n][n]) { for (size_t i = 0; i < n; i++) { for (size_t j = 0; j < i; j++) { double temp = matrix[i][j]; matrix[i][j] = matrix[j][i]; matrix[j][i] = temp; } } } void matmulTrans(size_t n, double A[n][n], double B[n][n], double result[n][n]) { transpose(n, B); for (size_t i = 0; i < n; i++) { for (size_t j = 0; j < n; j++) { double acc = 0; for (size_t k = 0; k < n; k++) { acc += A[i][k] * B[j][k]; } result[i][j] = acc; } } } void matmulBlock(size_t n, double A[n][n], double B[n][n], double result[n][n], size_t blockSize) { transpose(n, B); for (size_t i = 0; i < n; i += blockSize) { for (size_t j = 0; j < n; j += blockSize) { for (size_t iBlock = i; iBlock < i + blockSize; iBlock++) { for (size_t jBlock = j; jBlock < j + blockSize; jBlock++) { double acc = 0; for (size_t k = 0; k < n; k++) { acc += A[iBlock][k] * B[jBlock][k]; } result[iBlock][jBlock] = acc; } } } } } int main(int argc, char **argv) { if (argc != 3) { printf("Provide two arguments!\n"); return 1; } int n = atoi(argv[1]); int blockSize = atoi(argv[2]); double (*A)[n] = malloc(n * n * sizeof(double)); double (*B)[n] = malloc(n * n * sizeof(double)); double (*result)[n] = malloc(n * n * sizeof(double)); clock_t time1 = clock(); matmulBlock(n, A, B, result, blockSize); clock_t time2 = clock(); // matmul(n, A, B, result); clock_t time3 = clock(); matmulTrans(n, A, B, result); clock_t time4 = clock(); printf("Blocked version: %lfs.\nNaive version: %lfs.\n" "Transposed version: %lfs.\n", (double) (time2 - time1) / CLOCKS_PER_SEC, (double) (time3 - time2) / CLOCKS_PER_SEC, (double) (time4 - time3) / CLOCKS_PER_SEC); free(A); free(B); free(result); return 0; }
内容的提问来源于stack exchange,提问作者asdfldsfdfjjf
相关产品推荐
相关产品推荐

