device: update the logic for channelId assignment
This commit is contained in:
+20
-14
@@ -235,23 +235,29 @@ __forceinline__ __device__ void ncclKernelMain(struct ncclDevComm* comm, struct
|
|||||||
|
|
||||||
switch (tid/WARP_SIZE) {
|
switch (tid/WARP_SIZE) {
|
||||||
case 0:
|
case 0:
|
||||||
ncclShmem.channelId = blockIdx.x;
|
//ncclShmem.channelId = blockIdx.x;
|
||||||
/*for (int i = 0; i < num; i++) {
|
for (int i = 0; i < num; i++) {
|
||||||
if (channelMask.masks[i] & (1ull<<x)) {
|
|
||||||
y = __popcll(channelMask.masks[i] & ((1ull<<x)-1));
|
|
||||||
y = total + y;
|
|
||||||
if (blockIdx.x == y) ncclShmem.channelId = x;
|
|
||||||
}
|
|
||||||
if (WARP_SIZE < 64) {
|
|
||||||
x = WARP_SIZE + tid;
|
|
||||||
if (channelMask.masks[i] & (1ull<<x)) {
|
if (channelMask.masks[i] & (1ull<<x)) {
|
||||||
y = __popcll(channelMask.masks[i] & ((1ull<<x)-1));
|
y = __popcll(channelMask.masks[i] & ((1ull<<x)-1));
|
||||||
y = y + total;
|
y = total + y;
|
||||||
if (blockIdx.x == y) ncclShmem.channelId = x;
|
if (blockIdx.x == y) {
|
||||||
|
ncclShmem.channelId = y;
|
||||||
|
break;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
if (WARP_SIZE < 64) {
|
||||||
|
x = WARP_SIZE + tid;
|
||||||
|
if (channelMask.masks[i] & (1ull<<x)) {
|
||||||
|
y = __popcll(channelMask.masks[i] & ((1ull<<x)-1));
|
||||||
|
y = y + total;
|
||||||
|
if (blockIdx.x == y) {
|
||||||
|
ncclShmem.channelId = y;
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
total = total + __popcll(channelMask.masks[i]);
|
||||||
}
|
}
|
||||||
total = __popcll(channelMask.masks[i]);
|
|
||||||
}*/
|
|
||||||
break;
|
break;
|
||||||
case 1:
|
case 1:
|
||||||
if (tid < WARP_SIZE + NCCL_MAX_GROUPS)
|
if (tid < WARP_SIZE + NCCL_MAX_GROUPS)
|
||||||
|
|||||||
Reference in New Issue
Block a user