基于AVX-2指令快速计算__m256i寄存器前导零的优化方案问询
Great question! Calculating the first set bit index (or equivalently, the number of leading zeros) for a 256-bit AVX-2 register doesn’t have to involve tedious non-zero word hunting. We can leverage AVX-2 and BMI1 instructions to build clean, high-performance solutions. Below are two optimized approaches, plus how to extend this to arbitrary-length vectors.
Approach 1: 64-bit Chunk-Based Check
This method breaks the 256-bit register into 64-bit chunks, using fast zero-check instructions to quickly narrow down to the first non-zero chunk, then uses BMI1's lzcnt to get the leading zero count for that chunk.
#include <immintrin.h> #include <stdint.h> #include <limits.h> uint32_t find_first_set_bit_in_m256i(__m256i vec) { uint32_t total_leading_zeros = 0; // Check high 128 bits first if (!_mm256_testz_si256(vec, vec)) { __m128i high_half = _mm256_extracti128_si256(vec, 1); // Check high 64 bits of the 128-bit chunk if (!_mm_testz_si128(high_half, high_half)) { uint64_t chunk = _mm_extract_epi64(high_half, 1); total_leading_zeros = __lzcnt64(chunk); } else { uint64_t chunk = _mm_extract_epi64(high_half, 0); total_leading_zeros = 64 + __lzcnt64(chunk); } } else { __m128i low_half = _mm256_extracti128_si256(vec, 0); // Check high 64 bits of the low 128-bit chunk if (!_mm_testz_si128(low_half, low_half)) { uint64_t chunk = _mm_extract_epi64(low_half, 1); total_leading_zeros = 128 + __lzcnt64(chunk); } else { uint64_t chunk = _mm_extract_epi64(low_half, 0); total_leading_zeros = 192 + __lzcnt64(chunk); } } // Handle all-zero case if (total_leading_zeros == 256) { return UINT32_MAX; } // Convert leading zeros to first set bit index (bit 255 is highest) return 255 - total_leading_zeros; }
How it works:
_mm256_testz_si256and_mm_testz_si128are single-cycle instructions that check if a register is all zeros.__lzcnt64(BMI1) computes leading zeros for a 64-bit value in a single cycle.- We accumulate leading zeros from higher chunks to get the total, then subtract from 255 to get the index of the first set bit (since bit 255 is the highest position).
Approach 2: Byte Mask + Leading Zero Check
This uses _mm256_movemask_epi8 to create a 32-bit mask where each bit represents the highest bit of a byte in the __m256i register. We then find the highest set bit in this mask to locate the first non-zero byte, then compute the leading zero count within that byte.
#include <immintrin.h> #include <stdint.h> #include <limits.h> uint32_t find_first_set_bit_in_m256i(__m256i vec) { int byte_mask = _mm256_movemask_epi8(vec); if (byte_mask == 0) { return UINT32_MAX; // All zeros } // Find position of highest set bit in the mask (0 = lowest byte, 31 = highest byte) unsigned int mask_lz = __lzcnt_u32((unsigned int)byte_mask); int highest_byte_pos = 31 - mask_lz; // Extract the highest non-zero byte __m256i shifted = _mm256_srli_si256(vec, highest_byte_pos); uint8_t target_byte = _mm_extract_epi8(_mm256_castsi256_si128(shifted), 0); // Calculate leading zeros within the byte (adjust for 32-bit lzcnt) int byte_lz = __lzcnt_u32(target_byte) - 24; uint32_t total_leading_zeros = highest_byte_pos * 8 + byte_lz; return 255 - total_leading_zeros; }
How it works:
_mm256_movemask_epi8packs the highest bit of each of the 32 bytes in the __m256i register into a single 32-bit integer.__lzcnt_u32finds the leading zeros in the mask to quickly locate the highest byte with a set bit.- We shift the register to bring that byte to the lowest position, extract it, and compute its leading zeros to get the final count.
Extending to Arbitrary-Length Vectors
For longer vectors (e.g., a std::vector<__m256i>), iterate from the highest-value register to the lowest, checking each for non-zero values. Once a non-zero register is found, compute its internal first set bit index and add the offset from higher registers.
#include <vector> #include <immintrin.h> #include <stdint.h> #include <limits.h> // Reuse the find_first_set_bit_in_m256i function from above uint64_t find_first_set_in_long_vector(const std::vector<__m256i>& vec) { // Assumes vec[0] is the highest 256 bits, vec.back() is the lowest for (size_t i = 0; i < vec.size(); ++i) { __m256i reg = vec[i]; if (!_mm256_testz_si256(reg, reg)) { uint32_t inner_idx = find_first_set_bit_in_m256i(reg); // Calculate total index: offset from higher registers + inner index return i * 256 + inner_idx; } } // All zeros in the entire vector return UINT64_MAX; }
Key Notes:
- Compile with AVX-2 and BMI1 support: use
-mavx2 -mbmi(GCC/Clang) or/arch:AVX2(MSVC). - Adjust the vector iteration order if your long vector stores lower bits first (start from
vec.back()instead ofvec[0]). - The all-zero case returns
UINT32_MAX/UINT64_MAX—adjust this to fit your error-handling needs.
内容的提问来源于stack exchange,提问作者Sven Hager

