如何用SSE向量化实现向量-矩阵乘积?代码结果不符排查
问题:修正SSE向量化实现的结果不符问题
我有如下串行C++函数:
void routine2(float alpha, float beta) { unsigned int i, j; for (i = 0; i < N; i++) for (j = 0; j < N; j++) w[i] = w[i] - beta + alpha * A[i][j] * x[j]; }
以下是我编写的SSE向量化版本,但运行后结果与串行版本不一致:
void routine2_vec(float alpha, float beta) { __m128 alpha_2 = _mm_set1_ps(alpha); __m128 beta_2 = _mm_set1_ps(beta); __m128 w2,num1,A2,X2; __m128 multiple1, multiple2,sub; unsigned int i, j; for ( i = 0; i < N; ++i) { w2 = _mm_setzero_ps(); for ( j = 0; j < (N/4)*4; j += 4) { num1 = _mm_loadu_ps(&w[i]); A2 = _mm_loadu_ps(&A[i][j]); X2 = _mm_loadu_ps(&x[j]); multiple1 = _mm_mul_ps(A2, X2); multiple2 = _mm_mul_ps(alpha_2, multiple1); sub = _mm_sub_ps(num1, beta_2); w2 = _mm_add_ps(sub, multiple2); _mm_store_ss(&w[i], w2); } } }
为验证结果一致性,我编写了测试对比函数:
int routine2_test(float alpha, float beta) { unsigned int i, j; for (i = 0; i < N; i++) { for (j = 0; j < N; j++) { test2[i] = test2[i] - beta + alpha * A[i][j] * x[j]; } } // 对比结果 for (j = 0; j < N; j++) { if (equal(w[j], test2[j]) == 1) { printf("\n The result of w[%d] is not equal to test2[%d] \n", j, j); return 1; } } return 0; } unsigned short int equal(float a, float b) { float temp = a - b; //printf("\n %f %f", a, b); if ((fabs(temp) / fabs(b)) < EPSILON) return 0; // 结果一致 else return 1; // 结果不一致 }
测试逻辑:串行代码将结果存入test2数组,SSE版本将结果存入w数组,通过equal函数对比两者精度范围内的一致性。
完整补充代码(包含主函数与初始化)
#include <stdio.h> #include <time.h> #include <pmmintrin.h> #include <process.h> #include <chrono> #include <iostream> #include <immintrin.h> #include <omp.h> #define M 256*128 #define ARITHMETIC_OPERATIONS1 3*M #define TIMES1 1 #define N 2048 #define ARITHMETIC_OPERATIONS2 4*N*N #define TIMES2 1 constexpr auto EPSILON =0.0001; // 函数声明 void initialize(); void routine2(float alpha, float beta); unsigned short int equal(float a, float b); int routine2_test(float, float); void routine2_vec(float alpha, float beta); __declspec(align(64)) float test2[N]; __declspec(align(64)) float A[N][N], x[N], w[N]; int main() { float alpha = 0.023f, beta = 0.045f; double run_time, start_time; unsigned int t; initialize(); printf("\nRoutine2:"); start_time = omp_get_wtime(); // 启动计时器 for (t = 0; t < TIMES2; t++){ // routine2(alpha, beta); routine2_vec(alpha, beta); routine2_test(alpha, beta); } run_time = omp_get_wtime() - start_time; // 结束计时器 printf("\n Time elapsed is %f secs \n %e FLOPs achieved\n", run_time, (double)(ARITHMETIC_OPERATIONS2) / ((double)run_time / TIMES2)); return 0; } void initialize() { unsigned int i, j; // 初始化routine2数组 for (i = 0; i < N; i++) for (j = 0; j < N; j++) { A[i][j] = (i % 99) + (j % 14) + 0.013f; } // 初始化routine2数组 for (i = 0; i < N; i++) { x[i] = (i % 19) - 0.01f; w[i] = (i % 5) - 0.002f; test2[i]= (i % 5) - 0.002f; } }
问题分析与修正方案
你的SSE版本存在多个逻辑错误,导致结果与串行版本不一致:
1. 核心计算逻辑错误
串行代码中,w[i]在循环中是累计更新的:每一轮j循环,w[i]都会减去一次beta,再加上alpha*A[i][j]*x[j]。而你的SSE代码中,每次j循环都直接从原始w[i]取值计算,然后覆盖存储,相当于只保留了最后一组4个j值的计算结果,完全丢失了前面的累计过程。
2. Beta的处理错误
串行代码中,每遍历一个j,w[i]就减去一次beta,总共会减去N次beta。你的SSE代码中没有处理这个累计减beta的逻辑,反而每次都只减一次beta,这完全不符合原逻辑。
3. 未处理剩余元素(N不是4的倍数时)
当N无法被4整除时,(N/4)*4之后的剩余元素没有处理,会导致结果缺失。
修正后的SSE代码
void routine2_vec(float alpha, float beta) { __m128 alpha_vec = _mm_set1_ps(alpha); // 每4个元素对应的beta总和:4*beta __m128 beta_vec_4 = _mm_set1_ps(4.0f * beta); // 单个beta,处理剩余元素时用 const float beta_scalar = beta; for (unsigned int i = 0; i < N; ++i) { // 加载初始w[i]值,广播到4个元素位置,用于累计计算 __m128 w_accum = _mm_set1_ps(w[i]); // 累计减去的beta总数:初始为0 float total_beta = 0.0f; unsigned int j; // 处理4个元素为一组的循环 for (j = 0; j < (N / 4) * 4; j += 4) { // 加载A的一行中连续4个元素(数组已对齐,用_mm_load_ps更高效) __m128 A_vec = _mm_load_ps(&A[i][j]); // 加载x中连续4个元素 __m128 x_vec = _mm_load_ps(&x[j]); // 计算 alpha*A[i][j]*x[j] 的4个元素 __m128 prod = _mm_mul_ps(A_vec, x_vec); prod = _mm_mul_ps(prod, alpha_vec); // 累计加法:w_accum += prod w_accum = _mm_add_ps(w_accum, prod); // 累计beta:这一组4个j,对应减4次beta total_beta += 4.0f * beta; } // 处理剩余的不足4个的元素 for (; j < N; ++j) { w_accum[0] += alpha * A[i][j] * x[j]; total_beta += beta_scalar; } // 最后一次性减去所有beta总和 w[i] = w_accum[0] - total_beta; } }
额外优化建议
- 因为数组已经按64字节对齐,使用
_mm_load_ps代替_mm_loadu_ps,提升加载效率。 - 将beta的累计从循环内的多次操作改为最后一次性计算,减少浮点操作次数,同时避免循环内的频繁内存读写。
- 原串行代码的计算顺序调整为
w[i] = w[i] + alpha*A[i][j]*x[j],最后再减去N*beta,利用加法交换律保持结果一致的同时,提升计算效率。
内容的提问来源于stack exchange,提问作者Mojtaba Sayari
相关产品推荐
相关产品推荐

