SSE4.2 _mm_cmpistrm/_mm_cmpestrm指令返回错误结果排查
使用SSE4.2 _mm_cmpistrm计算数组交集时结果错误的排查与修复
问题描述
尝试用SSE4.2的_mm_cmpistrm指令计算两个uint16_t数组的交集,但返回结果不符合预期。
原代码
#include <nmmintrin.h> #include <cstdint> #include <cstdio> void test(uint16_t *a, uint16_t *b) { __m128i v_a = _mm_loadu_si128((__m128i*)a); __m128i v_b = _mm_loadu_si128((__m128i*)b); __m128i res_v1 = _mm_cmpistrm(v_a, v_b, _SIDD_UWORD_OPS | _SIDD_CMP_EQUAL_ANY | _SIDD_BIT_MASK); uint32_t mask_a = _mm_extract_epi32(res_v1, 0); __m128i res_v2 = _mm_cmpistrm(v_b, v_a, _SIDD_UWORD_OPS | _SIDD_CMP_EQUAL_ANY | _SIDD_BIT_MASK); uint32_t mask_b = _mm_extract_epi32(res_v2, 0); printf("match a:"); for (int i = 0; i < 8; i++) { if (mask_a & (1 << (7 - i))) { printf(" %d", a[i]); } } putchar('\n'); printf("match b:"); for (int i = 0; i < 8; i++) { if (mask_b & (1 << (7 - i))) { printf(" %d", b[i]); } } putchar('\n'); } int main() { uint16_t a[] = {13, 18, 19, 24, 97, 104, 456, 1024}; uint16_t b[] = {11, 17, 18, 19, 24, 58, 104, 456}; test(a, b); }
预期输出
match a: 18 19 24 104 456 match b: 18 19 24 104 456
实际输出
match a: 13 18 24 97 104 match b: 17 18 24 58 104
问题分析与修复
这是代码逻辑错误,和Intel CPU无关。
错误原因
_mm_cmpistrm使用_SIDD_BIT_MASK模式时,返回的掩码位与源操作数1的元素对应关系为:
- 源操作数1的第0个元素 → 掩码的第0位(最低位)
- 源操作数1的第1个元素 → 掩码的第1位
- ...
- 源操作数1的第7个元素 → 掩码的第7位
而你的代码中用1 << (7 - i)来检查掩码位,相当于把元素索引和掩码位的对应关系完全反转,导致错误匹配了数组元素。
修复后的代码
将掩码检查的位计算改为1 << i,让数组第i个元素对应掩码的第i位:
#include <nmmintrin.h> #include <cstdint> #include <cstdio> void test(uint16_t *a, uint16_t *b) { __m128i v_a = _mm_loadu_si128((__m128i*)a); __m128i v_b = _mm_loadu_si128((__m128i*)b); __m128i res_v1 = _mm_cmpistrm(v_a, v_b, _SIDD_UWORD_OPS | _SIDD_CMP_EQUAL_ANY | _SIDD_BIT_MASK); uint32_t mask_a = _mm_extract_epi32(res_v1, 0); __m128i res_v2 = _mm_cmpistrm(v_b, v_a, _SIDD_UWORD_OPS | _SIDD_CMP_EQUAL_ANY | _SIDD_BIT_MASK); uint32_t mask_b = _mm_extract_epi32(res_v2, 0); printf("match a:"); for (int i = 0; i < 8; i++) { // 修复:将(1 << (7 - i))改为(1 << i) if (mask_a & (1 << i)) { printf(" %d", a[i]); } } putchar('\n'); printf("match b:"); for (int i = 0; i < 8; i++) { // 修复:将(1 << (7 - i))改为(1 << i) if (mask_b & (1 << i)) { printf(" %d", b[i]); } } putchar('\n'); } int main() { uint16_t a[] = {13, 18, 19, 24, 97, 104, 456, 1024}; uint16_t b[] = {11, 17, 18, 19, 24, 58, 104, 456}; test(a, b); }
验证结果
修复后运行代码,将得到预期的输出:
match a: 18 19 24 104 456 match b: 18 19 24 104 456
内容的提问来源于stack exchange,提问作者zelin
相关产品推荐
相关产品推荐

