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

基于AVX的矩阵乘法实现问题:无法读写myC向量数据

用AVX指令集实现矩阵乘法的问题排查与修复

问题背景

正在基于AVX指令集编写矩阵乘法代码,结合MPI与SUMMA算法做分块处理。其中alignedBuffA和alignedBuffB是600×600的float*对齐数组,myC是存储结果本地块的vector<float>(结构为给定)。目前遇到两个核心问题:无法加载myC中的历史数据,无法将计算结果正确存储回vector(因加载报错暂未验证存储逻辑)。

核心代码片段

for(int i=0; i<blockRowsA; i++){
    for(int j=0; j<blockColsB; j+=8){
        // Load 8 floats from myC into X
        __m256 X = _mm256_load_ps(myC.data() + i*blockColsB+j);

        for(int l=0; l<blockRowsB; l++){
            // Calculate the result
            __m256 A256 = _mm256_set1_ps(alignedBuffA[i*blockRowsA + l]);
            __m256 B256 = _mm256_load_ps(&alignedBuffB[l*blockRowsB + j]);

            X = _mm256_fmadd_ps(A256, B256, X);
        }
        // Store X back into myC
        alignas(32) float tempArray[8];
        _mm256_storeu_ps(tempArray, X);
        myC.assign(tempArray, tempArray+8);
    }
}

预期算法逻辑

  • 从B中取出8个元素
  • 将每个元素与A第一行的第一个元素相乘
  • 累加B第二行8个元素与A第一行第二个元素相乘的结果
  • 遍历完B的所有行
  • 对B的下一组列重复上述步骤

简化原型代码(无MPI)

#include <cmath>
#include <fstream>
#include <iostream>
#include <iomanip>
#include <vector>

#include <immintrin.h>
#include <numeric>

void init_data(std::vector<float>& data, int rows, int cols) {
    for(int i=0; i<rows; i++)
        for(int j=0; j<cols; j++)
            data[i*cols+j] = (rows-i+j) % 4;
}

int main (int argc, char *argv[]){
    int blockRowsA = 600;
    int blockRowsB = 600;
    int blockColsB = 600;

    std::vector<float> myA(blockRowsA*blockRowsB);
    std::vector<float> myB(blockRowsB*blockColsB);
    std::vector<float> myC(blockRowsA*blockColsB);

    init_data(myA, blockRowsA, blockRowsB);
    init_data(myA, blockRowsB, blockColsB);

    // I don't know how I could read from a vector into __m256 so I used this
    float* alignedBuffA = static_cast<float*>(_mm_malloc(blockRowsA*blockRowsB * sizeof(float),32));
    float* alignedBuffB = static_cast<float*>(_mm_malloc(blockRowsB*blockColsB * sizeof(float),32));
    std::copy(myA.begin(), myA.end(), alignedBuffA);
    std::copy(myB.begin(), myB.end(), alignedBuffB);

    for(int i=0; i<blockRowsA; i++){
        for(int j=0; j<blockColsB; j+=8){
            __m256 X = _mm256_load_ps(myC.data() + i*blockColsB+j);

            for(int l=0; l<blockRowsB; l++){
                __m256 A256 = _mm256_set1_ps(alignedBuffA[i*blockRowsA + l]);
                __m256 B256 = _mm256_load_ps(&alignedBuffB[l*blockRowsB + j]);

                X = _mm256_fmadd_ps(A256, B256, X);
            }
            alignas(32) float tempArray[8];
            _mm256_storeu_ps(tempArray, X);
            std::cout << tempArray[0] << std::endl;
            myC.assign(tempArray, tempArray+8);
        }
    }
}

问题分析与修复方案

1. 加载myC数据失败的原因与修复

_mm256_load_ps要求内存必须32字节对齐,但std::vector默认分配的内存不保证对齐,直接使用会触发未定义行为(如崩溃)。

  • 修复:用_mm256_loadu_ps替代,该指令支持非对齐内存加载;或者给vector自定义对齐分配器(复杂度较高,优先选前者)。

2. 存储回vector的错误与修复

原代码中myC.assign(tempArray, tempArray+8)会把整个vector替换成8个元素,完全不符合“写入对应位置”的需求。

  • 修复:直接用_mm256_storeu_ps将计算结果写入myC的目标地址,无需临时数组:
    _mm256_storeu_ps(myC.data() + i*blockColsB + j, X);
    

3. 其他潜在错误修复

  • 初始化错误:原代码中init_data(myA, blockRowsB, blockColsB)应该初始化myB,改为init_data(myB, blockRowsB, blockColsB);
  • 数组索引错误:
    • A数组索引:alignedBuffA[i*blockRowsA + l] → alignedBuffA[i*blockRowsB + l](myA是blockRowsA行×blockRowsB列,行优先存储);
    • B数组索引:alignedBuffB[l*blockRowsB + j] → alignedBuffB[l*blockColsB + j](myB是blockRowsB行×blockColsB列,行优先存储);
  • 未初始化myC:SUMMA算法需要累加历史结果,需将myC初始化为0:std::vector<float> myC(blockRowsA*blockColsB, 0.0f)。

修复后的核心代码

for(int i=0; i<blockRowsA; i++){
    for(int j=0; j<blockColsB; j+=8){
        // 非对齐加载myC中的历史结果
        __m256 X = _mm256_loadu_ps(myC.data() + i*blockColsB + j);

        for(int l=0; l<blockRowsB; l++){
            // 修正A数组索引
            __m256 A256 = _mm256_set1_ps(alignedBuffA[i*blockRowsB + l]);
            // 修正B数组索引,因alignedBuffB是对齐内存,用load_ps更高效
            __m256 B256 = _mm256_load_ps(&alignedBuffB[l*blockColsB + j]);

            X = _mm256_fmadd_ps(A256, B256, X);
        }
        // 直接将结果存储回myC对应位置
        _mm256_storeu_ps(myC.data() + i*blockColsB + j, X);
    }
}

修复后的完整原型代码

#include <cmath>
#include <fstream>
#include <iostream>
#include <iomanip>
#include <vector>

#include <immintrin.h>
#include <numeric>

void init_data(std::vector<float>& data, int rows, int cols) {
    for(int i=0; i<rows; i++)
        for(int j=0; j<cols; j++)
            data[i*cols+j] = (rows-i+j) % 4;
}

int main (int argc, char *argv[]){
    int blockRowsA = 600;
    int blockRowsB = 600;
    int blockColsB = 600;

    std::vector<float> myA(blockRowsA*blockRowsB);
    std::vector<float> myB(blockRowsB*blockColsB);
    std::vector<float> myC(blockRowsA*blockColsB, 0.0f); // 初始化myC为0

    init_data(myA, blockRowsA, blockRowsB);
    init_data(myB, blockRowsB, blockColsB); // 修正初始化对象

    // 分配对齐内存并拷贝数据
    float* alignedBuffA = static_cast<float*>(_mm_malloc(blockRowsA*blockRowsB * sizeof(float),32));
    float* alignedBuffB = static_cast<float*>(_mm_malloc(blockRowsB*blockColsB * sizeof(float),32));
    std::copy(myA.begin(), myA.end(), alignedBuffA);
    std::copy(myB.begin(), myB.end(), alignedBuffB);

    for(int i=0; i<blockRowsA; i++){
        for(int j=0; j<blockColsB; j+=8){
            // 非对齐加载myC数据
            __m256 X = _mm256_loadu_ps(myC.data() + i*blockColsB + j);

            for(int l=0; l<blockRowsB; l++){
                // 修正A的索引
                __m256 A256 = _mm256_set1_ps(alignedBuffA[i*blockRowsB + l]);
                // 修正B的索引,对齐内存用load_ps
                __m256 B256 = _mm256_load_ps(&alignedBuffB[l*blockColsB + j]);

                X = _mm256_fmadd_ps(A256, B256, X);
            }
            // 直接存储回myC对应位置
            _mm256_storeu_ps(myC.data() + i*blockColsB + j, X);
        }
    }

    // 释放对齐内存
    _mm_free(alignedBuffA);
    _mm_free(alignedBuffB);

    // 输出部分结果验证
    std::cout << "myC[0][0] = " << myC[0] << std::endl;
    return 0;
}

内容的提问来源于stack exchange,提问作者Mikecraft1224

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.19 18:34:59