AVX2水平求和汇编实现问题:自定义内联汇编函数失效排查
AVX2水平求和汇编内联函数错误排查
先看你原有的C实现逻辑:
static inline float hsum256_ps(__m256 v) { __m128 lo = _mm256_castps256_ps128(v); __m128 hi = _mm256_extractf128_ps(v, 1); __m128 sum128 = _mm_add_ps(lo, hi); __m128 shuf = _mm_movehdup_ps(sum128); __m128 sums = _mm_add_ps(sum128, shuf); shuf = _mm_movehl_ps(shuf, sums); sums = _mm_add_ss(sums, shuf); return _mm_cvtss_f32(sums); }
你写的汇编版本有几个关键错误,导致逻辑完全偏离:
- 操作数顺序搞反:x86(包括AVX)汇编的指令格式是
指令 源操作数1, 源操作数2, 目标操作数,你写的vaddps %%xmm0, %%xmm0, %%xmm1是把xmm0加xmm0存到xmm1,但原逻辑是lo(xmm0)加hi(xmm1)得到sum128,这一步就错了。 - 缺少AVX前缀:
movshdup和movhlps应该用vmovshdup和vmovhlps,非VEX前缀的指令会切换到SSE模式,破坏AVX寄存器状态,也容易出逻辑问题。 - 寄存器使用逻辑混乱:后续的加法和shuffle操作完全没对应原C的步骤,比如你反复把xmm0加到自己存到xmm1,完全和水平求和的逻辑不沾边。
修正后的汇编代码如下,每一步都对应原C的逻辑:
static inline float hsum256_ps(__m256 v){ float r; __asm__ __volatile__ ( // 提取低128位到xmm0,对应lo = _mm256_castps256_ps128(v) "vextractf128 $0, %1, %%xmm0 \n\t" // 提取高128位到xmm1,对应hi = _mm256_extractf128_ps(v, 1) "vextractf128 $1, %1, %%xmm1 \n\t" // sum128 = lo + hi,结果存在xmm0 "vaddps %%xmm1, %%xmm0, %%xmm0 \n\t" // shuf = _mm_movehdup_ps(sum128),把xmm0的高半部分复制到低半部分,存在xmm1 "vmovshdup %%xmm0, %%xmm1 \n\t" // sums = sum128 + shuf,结果存在xmm0 "vaddps %%xmm1, %%xmm0, %%xmm0 \n\t" // shuf = _mm_movehl_ps(shuf, sums),把xmm0的高64位移到xmm1的低64位 "vmovhlps %%xmm0, %%xmm1, %%xmm1 \n\t" // sums = sums + shuf(单精度加法),结果存在xmm0 "vaddss %%xmm1, %%xmm0, %%xmm0 \n\t" // 把xmm0的低32位单精度值存到r "vmovss %%xmm0, %0 \n\t" : "=m"(r) : "x"(v) : "xmm0", "xmm1" ); return r; }
另外补充个小优化:其实可以不用vextractf128 $0,直接把YMM寄存器的低128位当成XMM寄存器使用,但上面的代码严格对应原C逻辑,更容易理解。
内容的提问来源于stack exchange,提问作者Googlebot
相关产品推荐
相关产品推荐

