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

[ROCm/rccl commit: 85eb1f16bc]
Tento commit je obsažen v:
Bertan Dogancay
2025-02-26 09:48:03 -05:00
odevzdal GitHub
rodič acf5822a6c
revize a9d09c6551
4 změnil soubory, kde provedl 24 přidání a 10 odebrání
+18 -4
Zobrazit soubor
@@ -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 {