为何_mm256_fmadd_ps指令会引发浮点精度损失?
AVX FMA指令
_mm256_fmadd_ps的精度差异问题 使用_mm256_fmadd_ps指令执行c = a * b + c的乘积累加操作时,发现在特定场景下(尤其是c已存在非零值时),该指令相比先执行乘法(mul)再执行加法(add)的操作会产生精度损失。
测试代码
#include <stdio.h> #include <stdlib.h> #include <string.h> #include <math.h> #include <time.h> #include <intrin.h> static inline void multiply_scalar_and_accumulate_generic(float *out, const float *in, const float scalar, unsigned int cnt) { const float *aPtr = (float *)in; float *cPtr = (float *)out; while (cnt >= 8) { *cPtr = (*aPtr++) * scalar + (*cPtr); cPtr++; *cPtr = (*aPtr++) * scalar + (*cPtr); cPtr++; *cPtr = (*aPtr++) * scalar + (*cPtr); cPtr++; *cPtr = (*aPtr++) * scalar + (*cPtr); cPtr++; *cPtr = (*aPtr++) * scalar + (*cPtr); cPtr++; *cPtr = (*aPtr++) * scalar + (*cPtr); cPtr++; *cPtr = (*aPtr++) * scalar + (*cPtr); cPtr++; *cPtr = (*aPtr++) * scalar + (*cPtr); cPtr++; cnt -= 8; } while (cnt-- > 0) { *cPtr = (*aPtr++) * scalar + (*cPtr); cPtr++; } return; } static inline void multiply_scalar_and_accumulate_avx(float *out, const float *in, const float scalar, unsigned int cnt) { unsigned int idx = 0; const float *aPtr = (float *)in; float *cPtr = (float *)out; __m256 aVal; __m256 cVal; const __m256 bVal = _mm256_set1_ps(scalar); for (; idx < cnt; idx += 8) { aVal = _mm256_loadu_ps(aPtr); cVal = _mm256_loadu_ps(cPtr); cVal = _mm256_add_ps(cVal, _mm256_mul_ps(aVal, bVal)); _mm256_storeu_ps(cPtr, cVal); aPtr += 8; cPtr += 8; } for (; idx < cnt; idx++) { *cPtr = (*aPtr++) * scalar + (*cPtr); cPtr++; } return; } static inline void multiply_scalar_and_accumulate_avx_fma(float *out, const float *in, const float scalar, unsigned int cnt) { unsigned int idx = 0; const float *aPtr = (float *)in; float *cPtr = (float *)out; __m256 aVal; __m256 cVal; const __m256 bVal = _mm256_set1_ps(scalar); for (; idx < cnt; idx += 8) { aVal = _mm256_loadu_ps(aPtr); cVal = _mm256_loadu_ps(cPtr); cVal = _mm256_fmadd_ps(aVal, bVal, cVal); _mm256_storeu_ps(cPtr, cVal); aPtr += 8; cPtr += 8; } for (; idx < cnt; idx++) { *cPtr = (*aPtr++) * scalar + (*cPtr); cPtr++; } return; } int main(void) { #define TEST_COUNT (0x4000) float *in = NULL; float *ref = NULL; float *out_avx = NULL; float *out_avx_fma = NULL; in = (float *)malloc(sizeof(float) * TEST_COUNT); ref = (float *)malloc(sizeof(float) * TEST_COUNT); out_avx = (float *)malloc(sizeof(float) * TEST_COUNT); out_avx_fma = (float *)malloc(sizeof(float) * TEST_COUNT); if ((in == NULL) || (ref == NULL) || (out_avx == NULL) || (out_avx_fma == NULL)) { printf("alloc failed\n"); return 0; } printf("test start\n"); float scalar = 0; float diff_avx = 0; float diff_avx_fma = 0; const float TOLERANCE = 1e-3f; srand(time(0)); memset(ref, 0x0, sizeof(float) * TEST_COUNT); memset(out_avx, 0x0, sizeof(float) * TEST_COUNT); memset(out_avx_fma, 0x0, sizeof(float) * TEST_COUNT); for (int i = 0; i < TEST_COUNT; i++) { in[i] = ((float)rand()) / ((float)rand()) * 10.0f; } scalar = ((float)rand()) / ((float)rand()) * 10.0f; multiply_scalar_and_accumulate_generic(ref, in, scalar, TEST_COUNT); multiply_scalar_and_accumulate_avx(out_avx, in, scalar, TEST_COUNT); multiply_scalar_and_accumulate_avx_fma(out_avx_fma, in, scalar, TEST_COUNT); #define MAKE_ACCUMULATE (1) #if MAKE_ACCUMULATE scalar = ((float)rand()) / ((float)rand()) * 10.0f; multiply_scalar_and_accumulate_generic(ref, in, scalar, TEST_COUNT); multiply_scalar_and_accumulate_avx(out_avx, in, scalar, TEST_COUNT); multiply_scalar_and_accumulate_avx_fma(out_avx_fma, in, scalar, TEST_COUNT); #endif for (int i = 0; i < TEST_COUNT; i++) { diff_avx = fabsf(out_avx[i] - ref[i]); diff_avx_fma = fabsf(out_avx_fma[i] - ref[i]); if (diff_avx > TOLERANCE) { printf("[Err AVX] pos:%06d, %20.4f != %20.4f, avx_diff:%.4f, avx_fma_diff:%.4f\n", i, out_avx[i], ref[i], diff_avx, diff_avx_fma); } if (diff_avx_fma > TOLERANCE) { printf("[Err AVX_FMA] pos:%06d, %20.4f != %20.4f, avx_fma_diff:%.4f, avx_diff:%.4f\n", i, out_avx_fma[i], ref[i], diff_avx_fma, diff_avx); } } printf("test end\n"); return 0; }
测试结果
MAKE_ACCUMULATE == 0:仅执行一次乘积累加,标量实现、AVX先乘后加实现、FMA指令实现的计算结果无明显差异,均在精度阈值内。MAKE_ACCUMULATE == 1:执行两次乘积累加(即c经过第一次累加后已为非零值),_mm256_fmadd_ps的计算结果与标量/AVX先乘后加的结果出现超过1e-3的偏差,触发错误输出。
内容的提问来源于stack exchange,提问作者Y-Jiechao
相关产品推荐
相关产品推荐

