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

128D数组欧氏距离计算性能优化问询(AVX512仍存瓶颈)

128维uint8_t向量的AVX512欧氏距离计算优化空间分析

我在做近似最近邻搜索(ANN)时,需要计算预加载查询点和多个邻近点的平方欧氏距离。这里的LSHPoint是array<int8_t,128>类型,但实际存储的是uint8_t的128维坐标——用int8_t只是为了LSH分桶时计算点积更方便,要是用无符号类型,点积全为正会导致所有点进入同一个桶。

标量实现代码

OptimizedDistance::scalar_distance(const LSHPoint& vec){
 float dist = 0;    
 for (size_t j = 0; j < 128; ++j){
     int diff = static_cast<int>(m_Query[j]) - static_cast<int>(vec[j]); 
     dist += diff * diff; }
 return dist;
}

128个8位差值的平方和最大为8323200,小于2^23,完全能被float精确表示,所以只用整数操作就能得到精确结果,但编译器无法自动推导这一点,会为标量循环生成不必要的慢汇编。

当前AVX512实现代码

float OptimizedDistance::avx512_distance_preloaded(const LSHPoint& vec)
{
    __m512i sum = _mm512_setzero_si512();

    __m512i v0 = _mm512_loadu_si512(reinterpret_cast<const __m512i*>(&vec[0]));
    __m512i v1 = _mm512_loadu_si512(reinterpret_cast<const __m512i*>(&vec[64]));

    __m512i diff_lo = _mm512_sub_epi32(_mm512_cvtepu8_epi32(_mm512_extracti32x4_epi32(m_AVX512Query.q0, 0)),
        _mm512_cvtepu8_epi32(_mm512_extracti32x4_epi32(v0, 0)));
    __m512i sq_lo = _mm512_mullo_epi32(diff_lo, diff_lo);
    sum = _mm512_add_epi32(sum, sq_lo);

    diff_lo = _mm512_sub_epi32(_mm512_cvtepu8_epi32(_mm512_extracti32x4_epi32(m_AVX512Query.q0, 1)),
        _mm512_cvtepu8_epi32(_mm512_extracti32x4_epi32(v0, 1)));
    sq_lo = _mm512_mullo_epi32(diff_lo, diff_lo);
    sum = _mm512_add_epi32(sum, sq_lo);

    diff_lo = _mm512_sub_epi32(_mm512_cvtepu8_epi32(_mm512_extracti32x4_epi32(m_AVX512Query.q0, 2)),
        _mm512_cvtepu8_epi32(_mm512_extracti32x4_epi32(v0, 2)));
    sq_lo = _mm512_mullo_epi32(diff_lo, diff_lo);
    sum = _mm512_add_epi32(sum, sq_lo);

    diff_lo = _mm512_sub_epi32(_mm512_cvtepu8_epi32(_mm512_extracti32x4_epi32(m_AVX512Query.q0, 3)),
        _mm512_cvtepu8_epi32(_mm512_extracti32x4_epi32(v0, 3)));
    sq_lo = _mm512_mullo_epi32(diff_lo, diff_lo);
    sum = _mm512_add_epi32(sum, sq_lo);

    diff_lo = _mm512_sub_epi32(_mm512_cvtepu8_epi32(_mm512_extracti32x4_epi32(m_AVX512Query.q1, 0)),
        _mm512_cvtepu8_epi32(_mm512_extracti32x4_epi32(v1, 0)));
    sq_lo = _mm512_mullo_epi32(diff_lo, diff_lo);
    sum = _mm512_add_epi32(sum, sq_lo);

    diff_lo = _mm512_sub_epi32(_mm512_cvtepu8_epi32(_mm512_extracti32x4_epi32(m_AVX512Query.q1, 1)),
        _mm512_cvtepu8_epi32(_mm512_extracti32x4_epi32(v1, 1)));
    sq_lo = _mm512_mullo_epi32(diff_lo, diff_lo);
    sum = _mm512_add_epi32(sum, sq_lo);

    diff_lo = _mm512_sub_epi32(_mm512_cvtepu8_epi32(_mm512_extracti32x4_epi32(m_AVX512Query.q1, 2)),
        _mm512_cvtepu8_epi32(_mm512_extracti32x4_epi32(v1, 2)));
    sq_lo = _mm512_mullo_epi32(diff_lo, diff_lo);
    sum = _mm512_add_epi32(sum, sq_lo);

    diff_lo = _mm512_sub_epi32(_mm512_cvtepu8_epi32(_mm512_extracti32x4_epi32(m_AVX512Query.q1, 3)),
        _mm512_cvtepu8_epi32(_mm512_extracti32x4_epi32(v1, 3)));
    sq_lo = _mm512_mullo_epi32(diff_lo, diff_lo);
    sum = _mm512_add_epi32(sum, sq_lo);

    __m256i sum_256 = _mm512_extracti64x4_epi64(sum, 0);
    __m256i sum_256_hi = _mm512_extracti64x4_epi64(sum, 1);
    sum_256 = _mm256_add_epi32(sum_256, sum_256_hi);
    __m128i sum_128 = _mm256_extracti128_si256(sum_256, 0);
    __m128i sum_128_hi = _mm256_extracti128_si256(sum_256, 1);
    sum_128 = _mm_add_epi32(sum_128, sum_128_hi);
    __m128i sum_64 = _mm_hadd_epi32(sum_128, sum_128);
    __m128i sum_32 = _mm_hadd_epi32(sum_64, sum_64);
    return static_cast<float>(_mm_cvtsi128_si32(sum_32));
}

问题

现在的疑问是:即便已经用了AVX512,这个128维数组的欧氏距离计算还能进一步提速吗?


优化方案分析

当然还有不少优化空间,以下是几个关键方向:

1. 优化数据加载与扩展步骤

当前代码反复调用_mm512_extracti32x4_epi32和_mm512_cvtepu8_epi32,产生了额外指令开销。可以一次性完成8位到32位的扩展,避免重复操作:

  • 对查询向量,在类初始化阶段就预扩展成8个__m512i的32位整数向量(每个对应16个8位元素的扩展结果),而非存储原始8位向量再反复提取扩展。
  • 对输入的vec,直接批量扩展子块,减少提取操作。

示例简化思路:

// 类初始化时完成查询向量的预扩展
// m_QueryExt[0] = _mm512_cvtepu8_epi32(_mm_loadu_si128((__m128i*)&m_Query[0]));
// m_QueryExt[1] = _mm512_cvtepu8_epi32(_mm_loadu_si128((__m128i*)&m_Query[16]));
// ... 共8个扩展向量

// 计算时直接使用预扩展的查询向量
__m512i vec_ext0 = _mm512_cvtepu8_epi32(_mm_loadu_si128((__m128i*)&vec[0]));
__m512i diff = _mm512_sub_epi32(m_QueryExt[0], vec_ext0);
sum = _mm512_add_epi32(sum, _mm512_mullo_epi32(diff, diff));

2. 简化水平求和步骤

当前求和阶段用了多次提取和hadd操作,AVX512提供了专门的水平求和指令可以一步完成:

  • 使用_mm512_reduce_add_epi32直接对__m512i向量的所有32位元素求和,替代原来的多步操作,大幅减少指令数和延迟。

修改后的求和代码:

int total_sum = _mm512_reduce_add_epi32(sum);
return static_cast<float>(total_sum);

3. 提升指令级并行(ILP)

当前代码按串行顺序处理每个16元素块,没有充分利用CPU的并行能力。可以将多个块的计算重叠,比如先批量加载扩展所有子块,再依次执行相减、平方、累加操作,让CPU同时处理多个独立任务:

// 先批量扩展所有vec子块
__m512i vec_ext0 = _mm512_cvtepu8_epi32(_mm_loadu_si128((__m128i*)&vec[0]));
__m512i vec_ext1 = _mm512_cvtepu8_epi32(_mm_loadu_si128((__m128i*)&vec[16]));
__m512i vec_ext2 = _mm512_cvtepu8_epi32(_mm_loadu_si128((__m128i*)&vec[32]));
__m512i vec_ext3 = _mm512_cvtepu8_epi32(_mm_loadu_si128((__m128i*)&vec[48]));

// 并行计算相减与平方
__m512i diff0 = _mm512_sub_epi32(m_QueryExt[0], vec_ext0);
__m512i diff1 = _mm512_sub_epi32(m_QueryExt[1], vec_ext1);
__m512i sq0 = _mm512_mullo_epi32(diff0, diff0);
__m512i sq1 = _mm512_mullo_epi32(diff1, diff1);

// 累加结果
sum = _mm512_add_epi32(sum, sq0);
sum = _mm512_add_epi32(sum, sq1);
// 后续子块同理

4. 内存对齐优化

当前使用_mm512_loadu_si512加载未对齐内存,如果能让LSHPoint的内存地址对齐到64字节(AVX512向量长度),可以改用_mm512_load_si512,减少内存加载延迟。可通过编译器属性(如__attribute__((aligned(64))))或对齐内存分配实现。

5. 利用无符号运算特性

由于实际数据是uint8_t,相减时可尝试使用_mm512_sub_epu32等无符号运算指令,部分CPU对无符号运算的调度效率更高,具体收益需实际测试验证。


内容的提问来源于stack exchange,提问作者s. heller

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 16:42:02