You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

矩阵乘法: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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.16 08:01:00