如何使用x86-64 Intrinsics将指定C代码转换为汇编级优化实现?
用AVX2 Intrinsics优化你的矩阵运算代码
以下是针对你提供的标量代码,基于AVX2指令集的intrinsics优化实现,同时附带关键逻辑的解释:
#include <immintrin.h> #include <stdint.h> // 替换为实际的N值,或改为运行时参数 #define N 1024 // 数组对齐到32字节(AVX2要求的最优内存对齐) float A[N][N] __attribute__((aligned(32))); float u1[N] __attribute__((aligned(32))), u2[N] __attribute__((aligned(32))); float v1[N] __attribute__((aligned(32))), v2[N] __attribute__((aligned(32))); float x[N] __attribute__((aligned(32))), y[N] __attribute__((aligned(32))); float z[N] __attribute__((aligned(32))), w[N] __attribute__((aligned(32))); void optimized_routine(float alpha, float beta) { uint32_t i, j; const uint32_t vec_width = 8; // AVX2一次处理8个单精度浮点数 uint32_t full_vecs = N / vec_width; uint32_t remainder = N % vec_width; // 第一部分:A[i][j] += u1[i]*v1[j] + u2[i]*v2[j] for (i = 0; i < N; i++) { // 将u1[i]、u2[i]广播为8元素向量,用于批量乘法 __m256 u1_broadcast = _mm256_set1_ps(u1[i]); __m256 u2_broadcast = _mm256_set1_ps(u2[i]); // 处理整向量块 for (j = 0; j < full_vecs * vec_width; j += vec_width) { __m256 v1_vec = _mm256_load_ps(&v1[j]); __m256 v2_vec = _mm256_load_ps(&v2[j]); __m256 A_row_vec = _mm256_load_ps(&A[i][j]); // 计算两个乘积的和,再累加到A的对应行 __m256 prod1 = _mm256_mul_ps(u1_broadcast, v1_vec); __m256 prod2 = _mm256_mul_ps(u2_broadcast, v2_vec); __m256 sum_prod = _mm256_add_ps(prod1, prod2); A_row_vec = _mm256_add_ps(A_row_vec, sum_prod); _mm256_store_ps(&A[i][j], A_row_vec); } // 处理剩余不足8个的元素 for (; j < N; j++) { A[i][j] += u1[i] * v1[j] + u2[i] * v2[j]; } } // 第二部分:x[i] += beta * A[j][i] * y[j](列向量点积) __m256 beta_broadcast = _mm256_set1_ps(beta); for (i = 0; i < N; i++) { float col_sum = 0.0f; // 批量计算8个元素的乘积和 for (j = 0; j < full_vecs * vec_width; j += vec_width) { // A的列可能不对齐,用无对齐加载 __m256 A_col_vec = _mm256_loadu_ps(&A[j][i]); __m256 y_vec = _mm256_load_ps(&y[j]); __m256 prod = _mm256_mul_ps(A_col_vec, y_vec); prod = _mm256_mul_ps(prod, beta_broadcast); // 将向量元素累加为标量 float temp[8]; _mm256_storeu_ps(temp, prod); col_sum += temp[0] + temp[1] + temp[2] + temp[3] + temp[4] + temp[5] + temp[6] + temp[7]; } // 处理剩余元素 for (; j < N; j++) { col_sum += beta * A[j][i] * y[j]; } x[i] += col_sum; } // 第三部分:x[i] += z[i](简单向量加法) for (i = 0; i < full_vecs * vec_width; i += vec_width) { __m256 x_vec = _mm256_load_ps(&x[i]); __m256 z_vec = _mm256_load_ps(&z[i]); x_vec = _mm256_add_ps(x_vec, z_vec); _mm256_store_ps(&x[i], x_vec); } for (; i < N; i++) { x[i] += z[i]; } // 第四部分:w[i] += alpha * A[i][j] * x[j](行向量点积) __m256 alpha_broadcast = _mm256_set1_ps(alpha); for (i = 0; i < N; i++) { __m256 row_accum = _mm256_setzero_ps(); // 批量计算行内元素的乘积和 for (j = 0; j < full_vecs * vec_width; j += vec_width) { __m256 A_row_vec = _mm256_load_ps(&A[i][j]); __m256 x_vec = _mm256_load_ps(&x[j]); __m256 prod = _mm256_mul_ps(A_row_vec, x_vec); prod = _mm256_mul_ps(prod, alpha_broadcast); row_accum = _mm256_add_ps(row_accum, prod); } // 将累加向量转为标量 float temp[8]; _mm256_storeu_ps(temp, row_accum); float row_sum = temp[0] + temp[1] + temp[2] + temp[3] + temp[4] + temp[5] + temp[6] + temp[7]; // 处理剩余元素 for (; j < N; j++) { row_sum += alpha * A[i][j] * x[j]; } w[i] += row_sum; } }
核心优化逻辑说明
内存对齐:所有数组通过
__attribute__((aligned(32)))对齐到32字节,AVX2的对齐内存访问(_mm256_load_ps/_mm256_store_ps)比无对齐访问性能更高;若无法保证对齐,改用_mm256_loadu_ps/_mm256_storeu_ps。向量广播:
_mm256_set1_ps将单个标量扩展为8元素向量,实现标量与整个向量的批量乘法,避免重复加载标量。批量运算:
_mm256_mul_ps和_mm256_add_ps一次完成8个浮点数的乘/加操作,相比标量循环理论性能提升8倍。余数处理:当N不是8的整数倍时,用标量循环处理剩余元素,确保覆盖所有数据且不越界。
点积优化:第二、第四部分的点积计算,先通过向量运算批量计算乘积,再将向量元素累加为标量,减少标量循环的执行次数。
编译与使用
- GCC/Clang编译时需添加
-mavx2参数启用AVX2支持;MSVC使用/arch:AVX2参数。 - 若N为运行时参数,需调整
full_vecs和remainder的计算逻辑,保持一致即可。
内容的提问来源于stack exchange,提问作者jebateapie123
相关产品推荐
相关产品推荐

