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

基于AVX-2指令快速计算__m256i寄存器前导零的优化方案问询

Efficiently Find First Set Bit Index in __m256i (AVX-2)

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_si256 and _mm_testz_si128 are 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_epi8 packs the highest bit of each of the 32 bytes in the __m256i register into a single 32-bit integer.
  • __lzcnt_u32 finds 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 of vec[0]).
  • The all-zero case returns UINT32_MAX/UINT64_MAX—adjust this to fit your error-handling needs.

内容的提问来源于stack exchange,提问作者Sven Hager

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.20 08:58:48