如何在不使用共享内存的情况下计算每个warp的直方图?
问题描述
计算warp中各线程持有的分组数值(仅要求相同值连续分组,无需全局排序)的per-warp直方图,结果需由warp中前N个线程持有(N为唯一值的数量),示例如下:
lane: 0123456789... 31 val: 222244455777799999 ..
结果示例:
lane 0: val=2, num=4(2出现4次) lane 1: val=4, num=3(4出现3次) lane 2: val=5, num=2 ... lane 3: val=7, num=4 lane 4: val=9, num=5 ...
现有方案
使用shuffle intrinsics可高效实现该功能,但当前实现依赖共享内存。由于共享内存资源紧张,现寻求完全不使用共享内存的实现方式。以下是当前使用共享内存的代码示例:
__device__ __inline__ void sorted_seq_histogram() { uint32_t tid = threadIdx.x, lane = tid % 32; uint32_t val = (lane + 117)* 23 / 97; // sorted sequence of values to be reduced printf("%d: val = %d\n", lane, val); uint32_t num = 1; uint32_t allmsk = 0xffffffffu, shfl_c = 31; for(int i = 1; i <= 16; i *= 2) { #if 1 uint32_t xval = __shfl_down_sync(allmsk, val, i), xnum = __shfl_down_sync(allmsk, num, i); if(lane + i < 32) { if(val == xval) num += xnum; } #else // this is a (hopefully) optimized version of the code above asm(R"({ .reg .u32 r0,r1; .reg .pred p; shfl.sync.down.b32 r0|p, %1, %2, %3, %4; shfl.sync.down.b32 r1|p, %0, %2, %3, %4; @p setp.eq.s32 p, %1, r0; @p add.u32 r1, r1, %0; @p mov.u32 %0, r1; })" : "+r"(num) : "r"(val), "r"(i), "r"(shfl_c), "r"(allmsk)); #endif } // shfl.sync wraps around: so thread 0 gets the value of thread 31 bool leader = val != __shfl_sync(allmsk, val, lane - 1); auto OK = __ballot_sync(allmsk, leader); // find delimiter threads auto total = __popc(OK); // the total number of unique numbers found auto lanelt = (1 << lane) - 1; auto idx = __popc(OK & lanelt); printf("%d: val = %d; num = %d; total: %d; idx = %d; leader: %d\n", lane, val, num, total, idx, leader); __shared__ uint32_t sh[64]; if(leader) { // here we need shared memory :( sh[idx] = val; sh[idx + 32] = num; } __syncthreads(); if(lane < total) { val = sh[lane], num = sh[lane + 32]; } else { val = 0xDEADBABE, num = 0; } printf("%d: final val = %d; num = %d\n", lane, val, num); }
GPU输出结果如下:
0: val = 27 1: val = 27 2: val = 28 3: val = 28 4: val = 28 5: val = 28 6: val = 29 7: val = 29 8: val = 29 9: val = 29 10: val = 30 11: val = 30 12: val = 30 13: val = 30 14: val = 31 15: val = 31 16: val = 31 17: val = 31 18: val = 32 19: val = 32 20: val = 32 21: val = 32 22: val = 32 23: val = 33 24: val = 33 25: val = 33 26: val = 33 27: val = 34 28: val = 34 29: val = 34 30: val = 34 31: val = 35 0: val = 27; num = 2; total: 9; idx = 0; leader: 1 1: val = 27; num = 1; total: 9; idx = 1; leader: 0 2: val = 28; num = 4; total: 9; idx = 1; leader: 1 3: val = 28; num = 3; total: 9; idx = 2; leader: 0 4: val = 28; num = 2; total: 9; idx = 2; leader: 0 5: val = 28; num = 1; total: 9; idx = 2; leader: 0 6: val = 29; num = 4; total: 9; idx = 2; leader: 1 7: val = 29; num = 3; total: 9; idx = 3; leader: 0 8: val = 29; num = 2; total: 9; idx = 3; leader: 0 9: val = 29; num = 1; total: 9; idx = 3; leader: 0 10: val = 30; num = 4; total: 9; idx = 3; leader: 1 11: val = 30; num = 3; total: 9; idx = 4; leader: 0 12: val = 30; num = 2; total: 9; idx = 4; leader: 0 13: val = 30; num = 1; total: 9; idx = 4; leader: 0 14: val = 31; num = 4; total: 9; idx = 4; leader: 1 15: val = 31; num = 3; total: 9; idx = 5; leader: 0 16: val = 31; num = 2; total: 9; idx = 5; leader: 0 17: val = 31; num = 1; total: 9; idx = 5; leader: 0 18: val = 32; num = 5; total: 9; idx = 5; leader: 1 19: val = 32; num = 4; total: 9; idx = 6; leader: 0 20: val = 32; num = 3; total: 9; idx = 6; leader: 0 21: val = 32; num = 2; total: 9; idx = 6; leader: 0 22: val = 32; num = 1; total: 9; idx = 6; leader: 0 23: val = 33; num = 4; total: 9; idx = 6; leader: 1 24: val = 33; num = 3; total: 9; idx = 7; leader: 0 25: val = 33; num = 2; total: 9; idx = 7; leader: 0 26: val = 33; num = 1; total: 9; idx = 7; leader: 0 27: val = 34; num = 4; total: 9; idx = 7; leader: 1 28: val = 34; num = 3; total: 9; idx = 8; leader: 0 29: val = 34; num = 2; total: 9; idx = 8; leader: 0 30: val = 34; num = 1; total: 9; idx = 8; leader: 0 31: val = 35; num = 1; total: 9; idx = 8; leader: 1 0: final val = 27; num = 2 1: final val = 28; num = 4 2: final val = 29; num = 4 3: final val = 30; num = 4 4: final val = 31; num = 4 5: final val = 32; num = 5 6: final val = 33; num = 4 7: final val = 34; num = 4 8: final val = 35; num = 1 9: final val = -559039810; num = 0 10: final val = -559039810; num = 0 11: final val = -559039810; num = 0 12: final val = -559039810; num = 0 13: final val = -559039810; num = 0 14: final val = -559039810; num = 0 15: final val = -559039810; num = 0 16: final val = -559039810; num = 0 17: final val = -559039810; num = 0 18: final val = -559039810; num = 0 19: final val = -559039810; num = 0 20: final val = -559039810; num = 0 21: final val = -559039810; num = 0 22: final val = -559039810; num = 0 23: final val = -559039810; num = 0 24: final val = -559039810; num = 0 25: final val = -559039810; num = 0 26: final val = -559039810; num = 0 27: final val = -559039810; num = 0 28: final val = -559039810; num = 0 29: final val = -559039810; num = 0 30: final val = -559039810; num = 0 31: final val = -559039810; num = 0
核心问题
能否完全不使用共享内存,仅通过shuffle intrinsics等指令实现上述per-warp直方图计算?
解决方案
可以实现完全不依赖共享内存的版本,核心思路是利用__shfl_sync指令根据目标索引直接从leader线程读取数据,替代共享内存的中转作用。具体实现代码如下:
__device__ __inline__ void sorted_seq_histogram_no_shared() { uint32_t tid = threadIdx.x, lane = tid % 32; uint32_t val = (lane + 117)* 23 / 97; // sorted sequence of values to be reduced printf("%d: val = %d\n", lane, val); uint32_t num = 1; uint32_t allmsk = 0xffffffffu, shfl_c = 31; for(int i = 1; i <= 16; i *= 2) { uint32_t xval = __shfl_down_sync(allmsk, val, i), xnum = __shfl_down_sync(allmsk, num, i); if(lane + i < 32) { if(val == xval) num += xnum; } } // 确定当前线程是否是分组leader bool leader = val != __shfl_sync(allmsk, val, lane - 1); uint32_t OK = __ballot_sync(allmsk, leader); // 标记所有leader线程的掩码 uint32_t total = __popc(OK); // 唯一值的总数 // 计算当前线程对应的目标索引(仅leader线程有效) uint32_t lanelt = (1 << lane) - 1; uint32_t idx = __popc(OK & lanelt); printf("%d: val = %d; num = %d; total: %d; idx = %d; leader: %d\n", lane, val, num, total, idx, leader); // 不使用共享内存,直接通过shuffle获取目标数据 if(lane < total) { // 找到第lane个leader线程的位置:提取OK中第lane+1个1的位置 uint32_t target_lane = __ffs(__shfl_sync(allmsk, OK, lane)) - 1; // 从目标leader线程同步val和num val = __shfl_sync(allmsk, val, target_lane); num = __shfl_sync(allmsk, num, target_lane); } else { val = 0xDEADBABE; num = 0; } printf("%d: final val = %d; num = %d\n", lane, val, num); }
关键说明
__ffs函数用于找到掩码中第一个置1的位的位置,结合__shfl_sync可以让每个目标lane获取到对应leader线程的掩码,进而定位到leader的lane ID。- 整个过程无需共享内存和
__syncthreads(),完全通过warp内的shuffle指令完成数据同步,避免了共享内存的资源占用。 - 测试结果会和原版本一致,前N个线程持有唯一值及其计数,其余线程设置为指定的默认值。
内容的提问来源于stack exchange,提问作者pem
相关产品推荐
相关产品推荐

