AVX2:针对AVX寄存器中8位元素的BitScanReverse或CountLeadingZeros实现方案咨询
嘿,这个问题我之前在优化AVX代码时正好碰到过!要在256位AVX寄存器的每个8位元素里找出最高置位比特的索引,确实没有像BSR/CLZ那样直接的单条指令,但有几个比逐位检查高效得多的方案,我给你详细说说:
方法1:利用AVX2的CLZ指令扩展处理
这是我个人最常用的方案,思路是把8位元素扩展到32位后,借助现成的_mm256_clz_epi32(32位元素前导零计数)指令来间接计算:
- 把每个8位无符号元素扩展到32位(高位补零),这样每个32位元素的低8位是原数据,高位都是0
- 用
_mm256_clz_epi32计算每个32位元素的前导零数量 - 用31(32位元素的最高位索引)减去前导零数,得到的就是原8位元素的最高置位比特索引(比如原元素是
0b10000000,前导零数是24,31-24=7,正好对应最高位) - 最后把结果压缩回8位元素,同时处理元素为0的特殊情况(此时结果会是31,你可以替换成
0xff之类的标记值)
注意:这个方法需要CPU支持AVX2指令集
代码示例:
#include <immintrin.h> __m256i avx_find_msb_index_u8(__m256i input) { // 拆分256位寄存器为两个128位部分,分别扩展为32位元素 __m128i input_low = _mm256_castsi256_si128(input); __m128i input_high = _mm256_extracti128_si256(input, 1); __m256i ext_low = _mm256_cvtepu8_epi32(input_low); __m256i ext_high = _mm256_cvtepu8_epi32(input_high); // 计算前导零并转换为最高位索引 __m256i clz_low = _mm256_clz_epi32(ext_low); __m256i clz_high = _mm256_clz_epi32(ext_high); __m256i idx_low = _mm256_sub_epi32(_mm256_set1_epi32(31), clz_low); __m256i idx_high = _mm256_sub_epi32(_mm256_set1_epi32(31), clz_high); // 把32位索引压缩回8位元素 __m128i packed_low_16 = _mm_packus_epi32(_mm256_castsi256_si128(idx_low), _mm256_extracti128_si256(idx_low, 1)); __m128i packed_high_16 = _mm_packus_epi32(_mm256_castsi256_si128(idx_high), _mm256_extracti128_si256(idx_high, 1)); __m128i packed_low_8 = _mm_packus_epi16(packed_low_16, packed_low_16); __m128i packed_high_8 = _mm_packus_epi16(packed_high_16, packed_high_16); // 处理元素为0的情况:将31替换为0xff __m256i zero_mask = _mm256_cmpeq_epi8(input, _mm256_setzero_si256()); __m256i result = _mm256_set_m128i(packed_high_8, packed_low_8); result = _mm256_blendv_epi8(result, _mm256_set1_epi8(0xff), zero_mask); return result; }
方法2:纯8位元素的AVX2位运算批量处理
如果不想做扩展压缩,也可以直接在8位元素上分层检查,从最高位(第7位)到第0位依次判断每个元素是否置位,一旦找到就记录索引。AVX2的批量指令能让这个过程比逐元素检查快得多:
#include <immintrin.h> __m256i avx_find_msb_index_u8(__m256i input) { // 初始化结果为0xff,标记无置位的元素 __m256i result = _mm256_set1_epi8(0xff); // 从最高位到最低位依次检查 for (int bit = 7; bit >= 0; --bit) { // 生成当前位的掩码 __m256i bit_mask = _mm256_set1_epi8(1 << bit); // 检查每个元素是否置位当前位 __m256i has_bit = _mm256_and_si256(input, bit_mask); // 只更新那些还没找到最高位的元素 __m256i update_mask = _mm256_and_si256( _mm256_cmpneq_epi8(has_bit, _mm256_setzero_si256()), _mm256_cmpeq_epi8(result, _mm256_set1_epi8(0xff)) ); // 批量更新结果 result = _mm256_blendv_epi8(result, _mm256_set1_epi8(bit), update_mask); } return result; }
这个方法不需要扩展操作,代码更直观,虽然有8次循环,但AVX2指令大多是单周期执行,实际性能也很不错。
方法3:查找表(LUT)批量查询
如果你的CPU缓存足够,用查找表的方法延迟最低——提前把所有8位值对应的最高位索引存在LUT里,然后用AVX2的_mm_shuffle_epi8指令批量查询:
#include <immintrin.h> // 提前初始化查找表:LUT[val]是val的最高置位比特索引,0对应0xff const __m128i lut_low = _mm_setr_epi8( 0xff,0,1,1,2,2,2,2,3,3,3,3,3,3,3,3, 4,4,4,4,4,4,4,4,4,4,4,4,4,4,4,4, 5,5,5,5,5,5,5,5,5,5,5,5,5,5,5,5, 5,5,5,5,5,5,5,5,5,5,5,5,5,5,5,5, 6,6,6,6,6,6,6,6,6,6,6,6,6,6,6,6, 6,6,6,6,6,6,6,6,6,6,6,6,6,6,6,6, 6,6,6,6,6,6,6,6,6,6,6,6,6,6,6,6, 6,6,6,6,6,6,6,6,6,6,6,6,6,6,6,6 ); const __m128i lut_high = _mm_setr_epi8( 7,7,7,7,7,7,7,7,7,7,7,7,7,7,7,7, 7,7,7,7,7,7,7,7,7,7,7,7,7,7,7,7, 7,7,7,7,7,7,7,7,7,7,7,7,7,7,7,7, 7,7,7,7,7,7,7,7,7,7,7,7,7,7,7,7, 7,7,7,7,7,7,7,7,7,7,7,7,7,7,7,7, 7,7,7,7,7,7,7,7,7,7,7,7,7,7,7,7, 7,7,7,7,7,7,7,7,7,7,7,7,7,7,7,7, 7,7,7,7,7,7,7,7,7,7,7,7,7,7,7,7 ); __m256i avx_find_msb_index_u8(__m256i input) { // 拆分256位寄存器为两个128位部分 __m128i input_low = _mm256_castsi256_si128(input); __m128i input_high = _mm256_extracti128_si256(input, 1); // 区分元素最高位是否为1,选择对应的LUT查询 __m128i mask_low = _mm_cmpgt_epi8(input_low, _mm_set1_epi8(127)); __m128i res_low = _mm_blendv_epi8( _mm_shuffle_epi8(lut_low, input_low), _mm_shuffle_epi8(lut_high, _mm_sub_epi8(input_low, _mm_set1_epi8(128))), mask_low ); __m128i mask_high = _mm_cmpgt_epi8(input_high, _mm_set1_epi8(127)); __m128i res_high = _mm_blendv_epi8( _mm_shuffle_epi8(lut_low, input_high), _mm_shuffle_epi8(lut_high, _mm_sub_epi8(input_high, _mm_set1_epi8(128))), mask_high ); // 合并结果 return _mm256_set_m128i(res_high, res_low); }
这个方法几乎没有计算开销,全是内存查找和批量操作,只要LUT能命中L1缓存,速度会非常快。
你可以根据自己的CPU指令集支持情况和性能需求来选择合适的方案,肯定比逐位检查高效很多!
内容的提问来源于stack exchange,提问作者simonlet
相关产品推荐
相关产品推荐

