如何用AVX2 64位比较向量结果?_mm256_testz_si256使用异常求助
问题分析与修正
你的代码存在几个关键问题,导致_mm256_testz_si256未按预期工作:
1. 逻辑判断错误
_mm256_cmpeq_epi64会在元素相等时返回对应64位全1,不等时返回全0。而_mm256_testz_si256(a,b)的作用是判断a & b的结果是否全为0:
- 如果全为0,返回1;否则返回0。
你当前用_mm256_testz_si256(r1, r1)判断cmpeq的结果,只有当所有元素都不等(r1全0)时才会返回1,这和你要检查“所有元素相等”的需求完全相反。
更直观的做法是对两个向量做异或运算:相等的元素异或结果为0,不等则为非0。之后用_mm256_testz_si256判断异或结果是否全0,全0则说明当前4个元素都相等。
2. 不必要的const_cast
vec1和vec2是const引用,你用const_cast去掉const是不必要的,甚至可能导致未定义行为。_mm256_loadu_si256可以直接接收const指针,只需调整类型转换即可。
3. 越界访问风险
当向量长度s不是4的倍数时,最后一次循环会加载超出向量范围的元素(比如s=10,start=8时,会访问索引8-11,但向量只有0-9),导致越界。需要单独处理剩余的不足4个元素。
修正后的代码
#include <vector> #include <cstdint> #include <iostream> #include <immintrin.h> bool is_equal(const std::vector<uint64_t>& vec1, const std::vector<uint64_t>& vec2, std::size_t s) { // 先检查长度是否足够 if (vec1.size() < s || vec2.size() < s) { return false; } const uint64_t* data1 = vec1.data(); const uint64_t* data2 = vec2.data(); std::size_t start = 0; // 处理4个元素一组的批量数据 const std::size_t batch_size = 4; while (start + batch_size <= s) { __m256i v1 = _mm256_loadu_si256(reinterpret_cast<const __m256i*>(data1 + start)); __m256i v2 = _mm256_loadu_si256(reinterpret_cast<const __m256i*>(data2 + start)); __m256i xor_result = _mm256_xor_si256(v1, v2); // testz返回1表示异或结果全0,即所有元素相等 if (_mm256_testz_si256(xor_result, xor_result) == 0) { return false; } start += batch_size; } // 处理剩余不足4个的元素 while (start < s) { if (data1[start] != data2[start]) { return false; } start++; } return true; } int main() { std::vector<uint64_t> vec = {1, 2, 3, 4, 5, 6, 7, 8, 9, 0, 1, 12}; std::vector<uint64_t> vec2 = {1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12}; std::cout << std::boolalpha << is_equal(vec, vec2, vec.size()) << std::endl; return 0; }
额外说明
- 你代码中使用的
_mm256_*系列是AVX2指令,不是SSE2。如果需要兼容SSE2,应该使用__m128i和对应的_mm_*指令(一次处理2个uint64_t元素)。 - 编译时需要添加AVX2编译选项,比如GCC/Clang用
-mavx2,MSVC用/arch:AVX2。
内容的提问来源于stack exchange,提问作者oyjh
相关产品推荐
相关产品推荐

