You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用SIMD查找4个__m256d寄存器中最小值及其来源索引的问题

使用SIMD查找4个__m256d寄存器中最小值及其来源索引的问题

嘿,我来帮你梳理这个问题~你的思路方向是对的,但代码里有几个关键漏洞导致没法正确工作,咱们一步步拆解并修正:

你的代码存在的核心问题

  1. _mm256_min_pd的使用误区:这个指令是逐元素比较两个寄存器,取对应位置的最小值。你现在的写法是把min依次和每个寄存器逐元素取最小,最后得到的min寄存器里,每个元素是四个寄存器对应位置的最小值,而不是16个值里的全局单个最小值。这直接导致后续的mask和索引计算完全偏离目标。
  2. 错误的比较逻辑:你用_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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.22 11:57:57