Merge remote-tracking branch 'nccl/master' into develop
[ROCm/rccl commit: e1a835910e]
This commit is contained in:
@@ -222,6 +222,7 @@ struct bootstrapState {
|
||||
struct ncclSocket ringSendSocket;
|
||||
union ncclSocketAddress* peerCommAddresses;
|
||||
union ncclSocketAddress* peerProxyAddresses;
|
||||
uint64_t* peerProxyAddressesUDS;
|
||||
struct unexConn* unexpectedConnections;
|
||||
int cudaDev;
|
||||
int rank;
|
||||
@@ -300,6 +301,7 @@ ncclResult_t bootstrapInit(struct ncclBootstrapHandle* handle, struct ncclComm*
|
||||
|
||||
// Create the service proxy
|
||||
NCCLCHECK(ncclCalloc(&state->peerProxyAddresses, nranks));
|
||||
NCCLCHECK(ncclCalloc(&state->peerProxyAddressesUDS, nranks));
|
||||
|
||||
// proxy is aborted through a message; don't set abortFlag
|
||||
NCCLCHECK(ncclCalloc(&proxySocket, 1));
|
||||
@@ -307,7 +309,13 @@ ncclResult_t bootstrapInit(struct ncclBootstrapHandle* handle, struct ncclComm*
|
||||
NCCLCHECK(ncclSocketListen(proxySocket));
|
||||
NCCLCHECK(ncclSocketGetAddr(proxySocket, state->peerProxyAddresses+rank));
|
||||
NCCLCHECK(bootstrapAllGather(state, state->peerProxyAddresses, sizeof(union ncclSocketAddress)));
|
||||
NCCLCHECK(ncclProxyInit(comm, proxySocket, state->peerProxyAddresses));
|
||||
// cuMem UDS support
|
||||
// Make sure we create a unique UDS socket name
|
||||
uint64_t randId;
|
||||
NCCLCHECK(getRandomData(&randId, sizeof(randId)));
|
||||
state->peerProxyAddressesUDS[rank] = getPidHash()+randId;
|
||||
NCCLCHECK(bootstrapAllGather(state, state->peerProxyAddressesUDS, sizeof(*state->peerProxyAddressesUDS)));
|
||||
NCCLCHECK(ncclProxyInit(comm, proxySocket, state->peerProxyAddresses, state->peerProxyAddressesUDS));
|
||||
|
||||
TRACE(NCCL_INIT, "rank %d nranks %d - DONE", rank, nranks);
|
||||
|
||||
@@ -360,8 +368,6 @@ ncclResult_t bootstrapSplit(struct ncclBootstrapHandle* handle, struct ncclComm*
|
||||
for (int i = 0; i < nranks; ++i) {
|
||||
comm->topParentRanks[i] = parent->topParentRanks[parentRanks[i]];
|
||||
}
|
||||
comm->proxyState = parent->sharedRes->proxyState;
|
||||
ncclAtomicRefCountIncrement(&parent->sharedRes->proxyState->refCount);
|
||||
} else {
|
||||
// Create the service proxy
|
||||
NCCLCHECKGOTO(ncclCalloc(&state->peerProxyAddresses, nranks), ret, fail);
|
||||
@@ -371,10 +377,17 @@ ncclResult_t bootstrapSplit(struct ncclBootstrapHandle* handle, struct ncclComm*
|
||||
NCCLCHECKGOTO(ncclSocketGetAddr(proxySocket, &tmpAddr), ret, fail);
|
||||
memcpy(state->peerProxyAddresses + rank, &tmpAddr, sizeof(union ncclSocketAddress));
|
||||
NCCLCHECKGOTO(bootstrapAllGather(state, state->peerProxyAddresses, sizeof(union ncclSocketAddress)), ret, fail);
|
||||
NCCLCHECKGOTO(ncclProxyInit(comm, proxySocket, state->peerProxyAddresses), ret, fail);
|
||||
// cuMem UDS support
|
||||
NCCLCHECKGOTO(ncclCalloc(&state->peerProxyAddressesUDS, nranks), ret, fail);
|
||||
// Make sure we create a unique UDS socket name
|
||||
uint64_t randId;
|
||||
NCCLCHECKGOTO(getRandomData(&randId, sizeof(randId)), ret, fail);
|
||||
state->peerProxyAddressesUDS[rank] = getPidHash()+randId;
|
||||
NCCLCHECKGOTO(bootstrapAllGather(state, state->peerProxyAddressesUDS, sizeof(*state->peerProxyAddressesUDS)), ret, fail);
|
||||
NCCLCHECKGOTO(ncclProxyInit(comm, proxySocket, state->peerProxyAddresses, state->peerProxyAddressesUDS), ret, fail);
|
||||
}
|
||||
|
||||
INFO(NCCL_INIT, "bootstrapSplit: rank %d nranks %d color %d key %d prev %d next %d - DONE", rank, nranks, color, key, prev, next);
|
||||
INFO(NCCL_INIT, "bootstrapSplit: comm %p parent %p rank %d nranks %d color %d key %d prev %d next %d - DONE", comm, parent, rank, nranks, color, key, prev, next);
|
||||
|
||||
exit:
|
||||
return ret;
|
||||
@@ -573,7 +586,7 @@ ncclResult_t bootstrapClose(void* commState) {
|
||||
struct bootstrapState* state = (struct bootstrapState*)commState;
|
||||
if (state->unexpectedConnections != NULL) {
|
||||
unexpectedFree(state);
|
||||
if (*state->abortFlag == 0) {
|
||||
if (__atomic_load_n(state->abortFlag, __ATOMIC_RELAXED) == 0) {
|
||||
WARN("Unexpected connections are not empty");
|
||||
return ncclInternalError;
|
||||
}
|
||||
@@ -597,6 +610,7 @@ ncclResult_t bootstrapAbort(void* commState) {
|
||||
NCCLCHECK(ncclSocketClose(&state->ringRecvSocket));
|
||||
free(state->peerCommAddresses);
|
||||
free(state->peerProxyAddresses);
|
||||
free(state->peerProxyAddressesUDS);
|
||||
free(state);
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user