AVX2 64位逐元素整数乘法溢出int64_t最大值的钳位方案
问题描述
基于Haswell指令集实现的__m256i类型64位整数逐元素乘法函数mul64_haswell_mul存在整数溢出问题:当a2.position[0-4]的计算结果超过int64_t最大值时会返回错误数值,本次测试场景下溢出结果的真实值为14618374452099416064。需要添加溢出处理逻辑:逐元素乘法计算结果若大于int64_t最大值,直接将对应位置的数值钳位为int64_t最大值。
原有问题代码如下:
union sseUnion { int64_t position[4]; btSimdFloat4 mVec256; }; // vector operator * : multiply element by element __m256i mul64_haswell_mul(__m256i a, __m256i b) { // instruction does not exist. Split into 32-bit multiplies __m256i bswap = _mm256_shuffle_epi32(b, 0xB1); // swap H<->L __m256i prodlh = _mm256_mullo_epi32(a, bswap); // 32 bit L*H products __m256i zero = _mm256_setzero_si256(); // 0 __m256i prodlh2 = _mm256_hadd_epi32(prodlh, zero); // a0Lb0H+a0Hb0L,a1Lb1H+a1Hb1L,0,0 __m256i prodlh3 = _mm256_shuffle_epi32(prodlh2, 0x73); // 0, a0Lb0H+a0Hb0L, 0, a1Lb1H+a1Hb1L __m256i prodll = _mm256_mul_epu32(a, b); // a0Lb0L,a1Lb1L, 64 bit unsigned products __m256i prod = _mm256_add_epi64(prodll, prodlh3); // a0Lb0L+(a0Lb0H+a0Hb0L)<<32, a1Lb1L+(a1Lb1H+a1Hb1L)<<32 return prod; } int main() { sseUnion _sseUnion; _sseUnion.mVec256 = _mm256_set_epi64x(1000000, 1000000, 1000000, 1000000); sseUnion a2; a2.mVec256 = _mm256_setr_epi64x(401000000, 401000000, 401000000, 401000000); a2.mVec256 = _mm256_add_epi64(_sseUnion.mVec256, a2.mVec256); a2.mVec256 = mul64_haswell_mul(_sseUnion.mVec256, a2.mVec256); a2.mVec256 = mul64_haswell_mul(_sseUnion.mVec256, a2.mVec256); printf("%d", a2.mVec256.m256i_i64[0]); }
实现方案
核心逻辑是在原有乘法计算完成后,通过符号位判定正溢出场景,无分支替换溢出值为INT64_MAX,整个流程仅新增5条AVX2指令,性能损耗极低,完全适配Haswell架构。
修正点说明
- 修复原union的类型不匹配问题:原代码用浮点向量类型
btSimdFloat4存储整数向量__m256i存在类型别名未定义行为,替换为__m256i类型。 - 新增溢出判定逻辑:两个非负64位整数相乘,若结果符号位为1(即结果为负),说明乘积超过
INT64_MAX,触发钳位。 - 修复原printf格式符错误:
%d仅能输出32位int,输出int64_t需使用%lld或标准库PRId64宏。
完整可运行代码
#include <immintrin.h> #include <stdint.h> #include <stdio.h> #define INT64_MAX 0x7FFFFFFFFFFFFFFFLL union sseUnion { int64_t position[4]; __m256i mVec256; }; // 带正溢出钳位的64位整数逐元素乘法 __m256i mul64_haswell_mul(__m256i a, __m256i b) { // 保留原有32位拆分计算64位乘积逻辑 __m256i bswap = _mm256_shuffle_epi32(b, 0xB1); __m256i prodlh = _mm256_mullo_epi32(a, bswap); __m256i zero = _mm256_setzero_si256(); __m256i prodlh2 = _mm256_hadd_epi32(prodlh, zero); __m256i prodlh3 = _mm256_shuffle_epi32(prodlh2, 0x73); __m256i prodll = _mm256_mul_epu32(a, b); __m256i prod = _mm256_add_epi64(prodll, prodlh3); // 正溢出判定:两个操作数均非负,但乘积为负 __m256i sign_a = _mm256_srai_epi64(a, 63); __m256i sign_b = _mm256_srai_epi64(b, 63); __m256i sign_prod = _mm256_srai_epi64(prod, 63); __m256i overflow_mask = _mm256_andnot_si256(_mm256_or_si256(sign_a, sign_b), sign_prod); // 溢出位置钳位为INT64_MAX const __m256i i64_max_vec = _mm256_set1_epi64x(INT64_MAX); prod = _mm256_blendv_epi8(prod, i64_max_vec, overflow_mask); return prod; } int main() { sseUnion _sseUnion; _sseUnion.mVec256 = _mm256_set_epi64x(1000000, 1000000, 1000000, 1000000); sseUnion a2; a2.mVec256 = _mm256_setr_epi64x(401000000, 401000000, 401000000, 401000000); a2.mVec256 = _mm256_add_epi64(_sseUnion.mVec256, a2.mVec256); a2.mVec256 = mul64_haswell_mul(_sseUnion.mVec256, a2.mVec256); a2.mVec256 = mul64_haswell_mul(_sseUnion.mVec256, a2.mVec256); printf("%lld", a2.position[0]); return 0; }
验证结果
测试场景下两次乘法后真实乘积为14618374452099416064,大于INT64_MAX(9223372036854775807),函数会正确将对应位置钳位为INT64_MAX,输出结果为9223372036854775807,符合需求。
内容的提问来源于stack exchange,提问作者张文阳
相关产品推荐
相关产品推荐

