Use bit reversal based mapping for multi-node (#1572)

[ROCm/rccl commit: 85eb1f16bc]
Этот коммит содержится в:
Bertan Dogancay
2025-02-26 09:48:03 -05:00
коммит произвёл GitHub
родитель acf5822a6c
Коммит a9d09c6551
4 изменённых файлов: 24 добавлений и 10 удалений
+18 -4
Просмотреть файл
@@ -267,11 +267,25 @@ inline __host__ uint8_t ncclP2pChannelBaseForRound(struct ncclComm* comm, int p2
// ncclP2pChannelToPart and ncclP2pChannelForPart are inverses. The device code
// uses ncclP2pChannelToPart to determine which part "this" channel is responsible for.
inline __host__ int ncclP2pChannelForPart(int nP2pChannels, int base, int part, int nParts) {
return (base * nParts + part) & (nP2pChannels-1);
inline __host__ int ncclP2pChannelForPart(int nP2pChannels, int base, int part, int nParts, int nNodes) {
if (nNodes > 1) {
// Only works because nP2pChannels is pow2
int nChannelsLog2 = countOneBits(nP2pChannels-1);
int delta = reverseBits(part, nChannelsLog2);
return (base + delta) & (nP2pChannels-1);
} else {
return (base * nParts + part) & (nP2pChannels-1);
}
}
inline __device__ int ncclP2pChannelToPart(int nP2pChannels, int base, int channel, int nParts) {
return (channel - base * nParts) & (nParts-1);
inline __device__ int ncclP2pChannelToPart(int nP2pChannels, int base, int channel, int nParts, int nNodes) {
if (nNodes > 1) {
// Only works because nP2pChannels is pow2
int nChannelsLog2 = countOneBits(nP2pChannels-1);
int delta = (channel-base) & (nP2pChannels-1);
return reverseBits(delta, nChannelsLog2);
} else {
return (channel - base * nParts) & (nParts-1);
}
}
struct alignas(16) ncclDevWorkColl {