使用SIMD查找4个__m256d寄存器中最小值及其来源索引的问题
使用SIMD查找4个__m256d寄存器中最小值及其来源索引的问题
嘿,我来帮你梳理这个问题~你的思路方向是对的,但代码里有几个关键漏洞导致没法正确工作,咱们一步步拆解并修正:
你的代码存在的核心问题
_mm256_min_pd的使用误区:这个指令是逐元素比较两个寄存器,取对应位置的最小值。你现在的写法是把min依次和每个寄存器逐元素取最小,最后得到的min寄存器里,每个元素是四个寄存器对应位置的最小值,而不是16个值里的全局单个最小值。这直接导致后续的mask和索引计算完全偏离目标。- 错误的比较逻辑:你用
_mm256_cmp_pd(min, min, _CMP_EQ_OQ)是让寄存器和自身比较,得到的mask全是1,根本没法定位到全局最小值的真实位置。
正确的实现思路
要完成需求,我们需要分两步走:
- 第一步:从16个值中找出全局最小值
- 第二步:遍历所有寄存器和元素,找到所有等于全局最小值的位置(处理重复值的情况),记录对应的寄存器索引和元素索引
修正后的代码
#include <immintrin.h> #include <float.h> #include <stdio.h> int main() { // 测试数据:v1[0]和v3[2]都是1.0,是全局最小值 __m256d v1 = _mm256_set_pd(1.0, 2.0, 3.0, 4.0); // 注意:_mm256_set_pd是逆序存储,元素顺序是[4.0, 3.0, 2.0, 1.0] __m256d v2 = _mm256_set_pd(5.0, 6.0, 7.0, 8.0); // 实际存储:[8.0,7.0,6.0,5.0] __m256d v3 = _mm256_set_pd(3.0, 4.0, 1.0, 2.0); // 实际存储:[2.0,1.0,4.0,3.0] __m256d v4 = _mm256_set_pd(6.0, 5.0, 8.0, 7.0); // 实际存储:[7.0,8.0,5.0,6.0] // 第一步:找出全局最小值 // 先两两比较寄存器,得到每个寄存器内部的最小值,再比较这些最小值 __m256d temp1 = _mm256_min_pd(v1, v2); __m256d temp2 = _mm256_min_pd(v3, v4); __m256d all_mins = _mm256_min_pd(temp1, temp2); // 现在all_mins里是四个位置的最小值,再在这个寄存器内部找最小 __m128d low = _mm256_extractf128_pd(all_mins, 0); __m128d high = _mm256_extractf128_pd(all_mins, 1); low = _mm_min_pd(low, high); double global_min = _mm_cvtsd_f64(_mm_min_sd(low, _mm_shuffle_pd(low, low, 1))); // 第二步:遍历所有寄存器和元素,定位最小值的来源 __m256d* regs[4] = {&v1, &v2, &v3, &v4}; const char* reg_names[4] = {"v1", "v2", "v3", "v4"}; printf("全局最小值为:%lf\n", global_min); printf("所有来源位置:\n"); for (int reg_idx = 0; reg_idx < 4; reg_idx++) { __m256d curr_reg = *regs[reg_idx]; // 比较当前寄存器的每个元素和全局最小值,生成mask __m256d cmp_result = _mm256_cmp_pd(curr_reg, _mm256_set1_pd(global_min), _CMP_EQ_OQ); int mask = _mm256_movemask_pd(cmp_result); // 遍历mask的每一位,检查哪个元素匹配 for (int elem_idx = 0; elem_idx < 4; elem_idx++) { if (mask & (1 << elem_idx)) { // 注意:_mm256_set_pd的元素是逆序存储的,实际元素索引和传入顺序相反 // 比如_mm256_set_pd(a,b,c,d)存储的是d,c,b,a,所以elem_idx 0对应d,3对应a int actual_elem_idx = 3 - elem_idx; printf("- 来自%s的第%d个元素(传入_set_pd时的第%d个值)\n", reg_names[reg_idx], elem_idx, actual_elem_idx); } } } return 0; }
关键细节说明
- 关于
_mm256_set_pd的存储顺序:这个函数是按逆序把值存入寄存器的,比如_mm256_set_pd(1.0,2.0,3.0,4.0)实际在寄存器里的顺序是[4.0,3.0,2.0,1.0],所以在定位元素时要注意索引的转换。 - 重复值处理:代码会遍历所有元素,找出所有等于全局最小值的位置,不会遗漏重复的情况。
- 高效的全局最小值查找:通过两两比较寄存器缩小范围,再在单个寄存器内提取128位部分继续比较,最后得到全局最小值,比逐个元素遍历更高效。
备注:内容来源于stack exchange,提问作者holmessh
相关产品推荐
相关产品推荐

