CUDA:利用warp内所有线程O(1)时间提取32位掩码置位索引
利用CUDA Warp线程常数时间提取32位掩码的置位索引
要在warp内以常数时间(无循环)提取32位掩码的置位索引,可借助CUDA的warp级内置指令实现,核心思路是让每个线程负责自身对应的位,通过前缀和计算该位在输出列表中的位置,再同步写入结果。
实现步骤与代码示例
用共享数组存储最终的索引列表,具体代码如下:
__shared__ int valid_indices[32]; // 存储结果,最多容纳32个索引 const uint32_t bitmask = __ballot_sync(0xffffffff, isValid); const int tid = threadIdx.x % 32; // 获取warp内线程ID(范围0-31) // 1. 判断当前线程对应的位是否被置位 int has_bit = (bitmask >> tid) & 1; // 2. 计算当前位之前的置位总数(前缀和),即该索引在结果列表中的位置 int rank = __popc(bitmask & ((1U << tid) - 1)); // 3. 若当前位有效,将线程ID写入结果列表的对应位置 if (has_bit) { valid_indices[rank] = tid; } // 同步warp内所有线程,确保结果写入完成 __syncwarp(); // 线程0负责输出结果 if (tid == 0) { int valid_count = __popc(bitmask); printf("Valid indices: "); for (int i = 0; i < valid_count; ++i) { printf("%d ", valid_indices[i]); } printf("\n"); }
关键逻辑说明
__popc指令:快速计算整数中置位的数量,这里通过bitmask & ((1U << tid) - 1)保留当前线程位之前的所有位,再用__popc得到这些位的置位总数,也就是当前有效位在结果列表中的索引。- 无循环保证:所有操作均为单条或固定数目的指令,没有遍历掩码的循环,完全符合常数时间要求。
- Warp同步:
__syncwarp()确保所有线程完成结果写入后,再由线程0读取输出,避免数据竞争。
示例验证
当掩码为0b100010001(十进制273)时:
- 线程0、4、8对应的
has_bit为1 - 线程0的
rank为0,线程4的rank为1,线程8的rank为2 - 最终
valid_indices数组的前3位会被填入0、4、8,与预期一致
内容的提问来源于stack exchange,提问作者Johan
相关产品推荐
相关产品推荐

