矩阵乘法:C代码转VFX优化后结果不符,求排查
AVX优化矩阵乘法代码错误排查与修正
问题分析
原代码是基于普通循环的方阵乘法,每次取mat1的两个元素与mat2的对应行做乘积累加;转换为AVX优化版本后结果不符合预期,主要存在以下几个关键错误:
1. 乘法操作的变量误用
原逻辑中,mat11(对应AVX变量A2)需要和mat2[(k*2)+1][j](对应B2)相乘,但优化代码中错误使用了A1(对应mat01)参与B2的乘法,导致计算逻辑完全错误:
// 错误代码 C2 = _mm256_mullo_epi32(A1, B2); // 正确逻辑应该是 C2 = _mm256_mullo_epi32(A2, B2);
2. 循环步长与AVX寄存器宽度不匹配
__m256i寄存器是256位,每个epi32元素占4字节,因此一个寄存器最多存储8个int类型数据。优化代码中j += 16的步长过大,单次循环仅处理8个元素,剩余8个元素完全未被处理,导致结果缺失。
3. 存储操作的类型转换错误
_mm256_storeu_epi32要求目标指针为__m256i*类型,但代码中错误转换为(__m128*),这会导致存储时的内存操作长度不匹配,破坏数据结构。
4. 结果矩阵未初始化(潜在问题)
原代码使用result[i][j] += ...进行累加,如果result矩阵未预先初始化为0,无论普通版本还是优化版本都会得到错误结果,优化版本的乱序内存操作会让错误更明显。
修正后的AVX优化代码
#include <immintrin.h> // 假设N是8的倍数,若不是需要处理剩余元素的边界逻辑 void mat_mul_avx(int N, int mat1[N][N], int mat2[N][N], int result[N][N]) { __m256i A1, A2, B1, B2, D, C1, C2; int i, j, k; // 先初始化结果矩阵为0 for (i = 0; i < N; ++i) { for (j = 0; j < N; ++j) { result[i][j] = 0; } } for (i = 0; i < N; ++i) { for (k = 0; k < N/2; ++k) { A1 = _mm256_set1_epi32(mat1[i][k*2]); A2 = _mm256_set1_epi32(mat1[i][(k * 2) + 1]); // 步长改为8,匹配__m256i的8个int容量 for (j = 0; j < N; j += 8) { B1 = _mm256_loadu_epi32((__m256i*) &mat2[k*2][j]); B2 = _mm256_loadu_epi32((__m256i*) &mat2[(k * 2) + 1][j]); D = _mm256_loadu_epi32((__m256i*) &result[i][j]); // 修正乘法变量,使用A2对应mat11的乘法 C1 = _mm256_mullo_epi32(A1, B1); C2 = _mm256_mullo_epi32(A2, B2); D = _mm256_add_epi32(D, C1); D = _mm256_add_epi32(D, C2); // 修正存储的类型转换 _mm256_storeu_epi32((__m256i*) &result[i][j], D); } } } // 若N不是8的倍数,处理剩余的j元素(边界逻辑) if (N % 8 != 0) { for (i = 0; i < N; ++i) { for (k = 0; k < N/2; ++k) { int mat01 = mat1[i][k*2]; int mat11 = mat1[i][(k*2)+1]; for (j = N - (N%8); j < N; ++j) { int n = mat01 * mat2[k*2][j]; int m = mat11 * mat2[(k*2)+1][j]; result[i][j] += n + m; } } } } }
额外说明
- 如果
N不是2的倍数,原代码的k < N/2会忽略最后一行mat1的元素,这个问题在普通版本和优化版本中都存在,需要根据需求补充边界处理。 - 使用
_mm256_loadu_epi32而非_mm256_load_epi32是因为矩阵内存不一定满足256位对齐要求,loadu支持非对齐加载,兼容性更好。
内容的提问来源于stack exchange,提问作者Yazeed Zaid
相关产品推荐
相关产品推荐

