128D数组欧氏距离计算性能优化问询(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

