Merge remote-tracking branch 'nccl/master' into develop
This commit is contained in:
+451
-118
@@ -49,7 +49,7 @@
|
||||
#endif
|
||||
|
||||
const char* ncclFuncStr[NCCL_NUM_FUNCTIONS+2] = { "Broadcast", "Reduce", "AllGather", "ReduceScatter", "AllReduce", "SendRecv", "AllToAllPivot" };
|
||||
const char* ncclAlgoStr[NCCL_NUM_ALGORITHMS] = { "Tree", "Ring", "CollNet" };
|
||||
const char* ncclAlgoStr[NCCL_NUM_ALGORITHMS] = { "Tree", "Ring", "CollNetDirect", "CollNetChain" };
|
||||
const char* ncclProtoStr[NCCL_NUM_PROTOCOLS] = { "LL", "LL128", "Simple" };
|
||||
const char* ncclDevRedOpStr[ncclNumDevRedOps] = { "Sum", "Prod", "Max", "Min", "PreMulSum", "SumPostDiv" };
|
||||
const char *ncclTypeStr[ncclNumTypes] = {"_i8", "_u8", "_i32", "_u32", "_i64", "_u64", "_f16", "_f32", "_f64", "_b16"};
|
||||
@@ -57,6 +57,8 @@ const char *ncclTypeStr[ncclNumTypes] = {"_i8", "_u8", "_i32", "_u32", "_i64", "
|
||||
NCCL_PARAM(GroupCudaStream, "GROUP_CUDA_STREAM", NCCL_GROUP_CUDA_STREAM);
|
||||
|
||||
NCCL_PARAM(CheckPointers, "CHECK_POINTERS", 0);
|
||||
NCCL_PARAM(CommBlocking, "COMM_BLOCKING", 0);
|
||||
|
||||
struct allocationTracker allocTracker[MAX_ALLOC_TRACK_NGPU] = {};
|
||||
|
||||
static uint64_t hashUniqueId(ncclUniqueId const &id) {
|
||||
@@ -90,17 +92,10 @@ pthread_mutex_t initLock = PTHREAD_MUTEX_INITIALIZER;
|
||||
static bool initialized = false;
|
||||
static size_t maxLocalSizeBytes = 0;
|
||||
|
||||
bool ncclMainExited = false;
|
||||
|
||||
static void atexitHandler() {
|
||||
ncclMainExited = true;
|
||||
}
|
||||
|
||||
static ncclResult_t ncclInit() {
|
||||
if (__atomic_load_n(&initialized, __ATOMIC_ACQUIRE)) return ncclSuccess;
|
||||
pthread_mutex_lock(&initLock);
|
||||
if (!initialized) {
|
||||
atexit(atexitHandler);
|
||||
initEnv();
|
||||
initGdrCopy();
|
||||
maxLocalSizeBytes = ncclKernMaxLocalSize();
|
||||
@@ -299,46 +294,11 @@ void ncclCommPushCudaGdrFree(struct ncclComm* comm, void* handle) {
|
||||
comm->destructorHead = dtor;
|
||||
}
|
||||
|
||||
void commZombieCleanup(struct ncclComm* comm) {
|
||||
ncclMemoryStackDestruct(&comm->memScoped);
|
||||
ncclMemoryStackDestruct(&comm->memPermanent);
|
||||
|
||||
struct ncclComm* intraComm0 = comm->intraComm0;
|
||||
if (0 == ncclAtomicRefCountDecrement(&intraComm0->intraRefs)) {
|
||||
// Wait for all service threads to be done. We could not
|
||||
// do it earlier because it could have blocked and prevented
|
||||
// other ranks in the process to call ncclCommDestroy
|
||||
comm = intraComm0;
|
||||
while (comm != nullptr) {
|
||||
if (comm->proxyState.thread) pthread_join(comm->proxyState.thread, nullptr);
|
||||
struct ncclComm* next = comm->intraNext;
|
||||
free(comm);
|
||||
comm = next;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
static void* commZombieMain(void* arg) {
|
||||
ncclResult_t result = ncclSuccess;
|
||||
struct ncclComm* comm = (struct ncclComm*)arg;
|
||||
while (comm->persistentRefs != 0) {
|
||||
struct ncclCommCallback* cb = ncclIntruQueueMpscDequeueAll(&comm->callbackQueue, /*waitSome=*/true);
|
||||
while (cb != nullptr) {
|
||||
struct ncclCommCallback* next = cb->next;
|
||||
NCCLCHECKGOTO(cb->fn(comm, cb), result, ignore); // may reclaim memory of cb
|
||||
ignore:
|
||||
cb = next;
|
||||
}
|
||||
}
|
||||
commZombieCleanup(comm);
|
||||
return arg;
|
||||
}
|
||||
|
||||
static ncclResult_t commFree(ncclComm_t comm) {
|
||||
if (comm == NULL)
|
||||
return ncclSuccess;
|
||||
|
||||
// First stop all threads before we free anything.
|
||||
// Stop all threads before we free anything.
|
||||
NCCLCHECK(ncclProxyDestroy(comm));
|
||||
|
||||
delete[] comm->userRedOps;
|
||||
@@ -371,9 +331,12 @@ static ncclResult_t commFree(ncclComm_t comm) {
|
||||
#endif
|
||||
|
||||
free(comm->peerInfo);
|
||||
ncclTopoFree(comm->topo);
|
||||
for (int n=0; n<comm->nNodes; n++) free(comm->nodeRanks[n].localRankToRank);
|
||||
free(comm->nodeRanks);
|
||||
if (comm->topo)
|
||||
ncclTopoFree(comm->topo);
|
||||
if (comm->nodeRanks) {
|
||||
for (int n=0; n<comm->nNodes; n++) free(comm->nodeRanks[n].localRankToRank);
|
||||
free(comm->nodeRanks);
|
||||
}
|
||||
free(comm->rankToNode);
|
||||
free(comm->rankToLocalRank);
|
||||
|
||||
@@ -386,10 +349,10 @@ static ncclResult_t commFree(ncclComm_t comm) {
|
||||
if (comm->doneEvent != NULL)
|
||||
CUDACHECK(hipEventDestroy(comm->doneEvent));
|
||||
|
||||
NCCLCHECK(ncclStrongStreamDestruct(&comm->hostStream));
|
||||
NCCLCHECK(ncclStrongStreamDestruct(&comm->deviceStream));
|
||||
|
||||
NCCLCHECK(ncclCudaHostFree((void *)comm->abortFlag));
|
||||
if (comm->initState == ncclSuccess) {
|
||||
NCCLCHECK(ncclStrongStreamDestruct(&comm->hostStream));
|
||||
NCCLCHECK(ncclStrongStreamDestruct(&comm->deviceStream));
|
||||
}
|
||||
|
||||
struct ncclDestructor* dtor = comm->destructorHead;
|
||||
while (dtor != nullptr) {
|
||||
@@ -398,16 +361,34 @@ static ncclResult_t commFree(ncclComm_t comm) {
|
||||
}
|
||||
CUDACHECK(hipStreamDestroy(comm->sideStream));
|
||||
|
||||
ncclMemoryStackDestruct(&comm->memScoped);
|
||||
ncclMemoryStackDestruct(&comm->memPermanent);
|
||||
|
||||
commPoison(comm); // Important that this does not interfere with anything used below.
|
||||
|
||||
if (comm->persistentRefs == 0) {
|
||||
commZombieCleanup(comm);
|
||||
if (comm->initState == ncclSuccess) {
|
||||
struct ncclComm* intraComm0 = comm->intraComm0;
|
||||
if (0 == ncclAtomicRefCountDecrement(&intraComm0->intraRefs)) {
|
||||
// Wait for all service threads to be done. We could not
|
||||
// do it earlier because it could have blocked and prevented
|
||||
// other ranks in the process to call ncclCommDestroy
|
||||
comm = intraComm0;
|
||||
while (comm != nullptr) {
|
||||
if (comm->proxyState.thread) pthread_join(comm->proxyState.thread, nullptr);
|
||||
struct ncclComm* next = comm->intraNext;
|
||||
free(comm);
|
||||
comm = next;
|
||||
}
|
||||
}
|
||||
} else if (comm->proxyState.thread) {
|
||||
pthread_join(comm->proxyState.thread, nullptr);
|
||||
ncclCudaHostFree((void *)comm->abortFlag);
|
||||
free(comm);
|
||||
} else {
|
||||
// Spawn a thread to listen for remaining messages from graph cleanup.
|
||||
pthread_t zombie;
|
||||
pthread_create(&zombie, nullptr, commZombieMain, comm);
|
||||
pthread_detach(zombie);
|
||||
ncclCudaHostFree((void *)comm->abortFlag);
|
||||
free(comm);
|
||||
}
|
||||
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
@@ -445,6 +426,26 @@ static ncclResult_t dmaBufSupported(struct ncclComm* comm) {
|
||||
return ncclInternalError;
|
||||
}
|
||||
|
||||
ncclResult_t ncclCommEnsureReady(ncclComm_t comm) {
|
||||
/* comm must be ready, or error will be reported */
|
||||
ncclResult_t ret = ncclSuccess;
|
||||
|
||||
if (*comm->abortFlag) {
|
||||
ncclGroupJobAbort();
|
||||
} else {
|
||||
NCCLCHECK(ncclCommGetAsyncError(comm, &ret));
|
||||
if (ret != ncclSuccess) {
|
||||
/* if ret is not ncclInProgress, we just keep it. */
|
||||
WARN("Attempt to use communicator before the previous operation returned ncclSuccess\n");
|
||||
if (ret == ncclInProgress) ret = ncclInvalidArgument;
|
||||
goto exit;
|
||||
}
|
||||
}
|
||||
|
||||
exit:
|
||||
return ret;
|
||||
}
|
||||
|
||||
static ncclResult_t commAlloc(ncclComm_t* comret, int ndev, int rank, int virtualId) {
|
||||
if (ndev < 1) {
|
||||
WARN("invalid device count (%d) requested", ndev);
|
||||
@@ -456,7 +457,19 @@ static ncclResult_t commAlloc(ncclComm_t* comret, int ndev, int rank, int virtua
|
||||
}
|
||||
|
||||
struct ncclComm* comm;
|
||||
NCCLCHECK(ncclCalloc(&comm, 1));
|
||||
/* Cuurently we calloc comm in ncclCommInitRankDev for async function support.
|
||||
* This 'if' structure is designed to consider the case where commAlloc is called
|
||||
* in other cases except ncclCommInitRankDev. */
|
||||
if (*comret == NULL) {
|
||||
/* user requests a new communicator */
|
||||
NCCLCHECK(ncclCalloc(&comm, 1));
|
||||
NCCLCHECK(ncclCudaHostCalloc((uint32_t**)&comm->abortFlag, 1));
|
||||
NCCLCHECK(ncclCommSetAsyncError(comm, ncclInProgress));
|
||||
} else {
|
||||
/* We already allocated a communicator in ncclCommInitRankDev. */
|
||||
comm = *comret;
|
||||
}
|
||||
|
||||
ncclMemoryStackConstruct(&comm->memPermanent);
|
||||
ncclMemoryStackConstruct(&comm->memScoped);
|
||||
comm->destructorHead = nullptr;
|
||||
@@ -485,10 +498,6 @@ static ncclResult_t commAlloc(ncclComm_t* comret, int ndev, int rank, int virtua
|
||||
CUDACHECK(hipStreamCreateWithFlags(&comm->sideStream, hipStreamNonBlocking));
|
||||
comm->checkPointers = ncclParamCheckPointers() == 1 ? true : false;
|
||||
comm->dmaBufSupport = (dmaBufSupported(comm) == ncclSuccess) ? true : false;
|
||||
comm->fatalError = ncclSuccess;
|
||||
|
||||
NCCLCHECK(ncclCudaHostCalloc((uint32_t**)&comm->abortFlag, 1));
|
||||
*comm->abortFlag = 0;
|
||||
|
||||
#ifdef ENABLE_COLLTRACE
|
||||
NCCLCHECK(ncclCudaHostCalloc((uint32_t **)&comm->collTraceTail, 1));
|
||||
@@ -572,8 +581,9 @@ static ncclResult_t devCommSetup(ncclComm_t comm) {
|
||||
tmpCommAndChans.channels[c].ring = comm->channels[c].ring;
|
||||
tmpCommAndChans.channels[c].ring.userRanks = comm->channels[c].devRingUserRanks;
|
||||
tmpCommAndChans.channels[c].tree = comm->channels[c].tree;
|
||||
tmpCommAndChans.channels[c].collnetChain = comm->channels[c].collnetChain;
|
||||
tmpCommAndChans.channels[c].collnetDirect = comm->channels[c].collnetDirect;
|
||||
tmpCommAndChans.channels[c].binTree = comm->channels[c].binTree;
|
||||
tmpCommAndChans.channels[c].collTree = comm->channels[c].collTree;
|
||||
tmpCommAndChans.channels[c].workFifoDone = &comm->workFifoDone[c];
|
||||
|
||||
if (comm->channels[c].ring.userRanks != nullptr) {
|
||||
@@ -682,6 +692,8 @@ NCCL_PARAM(BuffSize, "BUFFSIZE", -2);
|
||||
NCCL_PARAM(LlBuffSize, "LL_BUFFSIZE", -2);
|
||||
NCCL_PARAM(Ll128BuffSize, "LL128_BUFFSIZE", -2);
|
||||
|
||||
NCCL_PARAM(P2pNetChunkSize, "P2P_NET_CHUNKSIZE", (1 << 17)); /* 128 kB */
|
||||
|
||||
static ncclResult_t computeBuffSizes(struct ncclComm* comm) {
|
||||
int cpuArch, cpuVendor, cpuModel;
|
||||
NCCLCHECK(ncclTopoCpuType(comm->topo, &cpuArch, &cpuVendor, &cpuModel));
|
||||
@@ -694,12 +706,15 @@ static ncclResult_t computeBuffSizes(struct ncclComm* comm) {
|
||||
for (int p=0; p<NCCL_NUM_PROTOCOLS; p++) {
|
||||
comm->buffSizes[p] = envs[p] != -2 ? envs[p] : defaults[p];
|
||||
}
|
||||
|
||||
comm->p2pNetChunkSize = ncclParamP2pNetChunkSize();
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
NCCL_PARAM(GraphDumpFileRank, "GRAPH_DUMP_FILE_RANK", 0);
|
||||
NCCL_PARAM(CollNetNodeThreshold, "COLLNET_NODE_THRESHOLD", 2);
|
||||
NCCL_PARAM(NvbPreconnect, "NVB_PRECONNECT", 0);
|
||||
NCCL_PARAM(AllocP2pNetLLBuffers, "NCCL_ALLOC_P2P_NET_LL_BUFFERS", 0);
|
||||
|
||||
static ncclResult_t initTransportsRank(struct ncclComm* comm, ncclUniqueId* commId) {
|
||||
// We use 2 AllGathers
|
||||
@@ -825,6 +840,8 @@ static ncclResult_t initTransportsRank(struct ncclComm* comm, ncclUniqueId* comm
|
||||
allXgmi &= isXGMI;
|
||||
}
|
||||
}
|
||||
// Initialize num P2P LL buffers for this communicator
|
||||
comm->allocP2pNetLLBuffers = ncclParamAllocP2pNetLLBuffers() == 1;
|
||||
|
||||
if (comm->rank == ncclParamGraphDumpFileRank()) {
|
||||
struct ncclTopoGraph* graphs[3] = { &ringGraph, &treeGraph, &collNetGraph };
|
||||
@@ -856,8 +873,8 @@ static ncclResult_t initTransportsRank(struct ncclComm* comm, ncclUniqueId* comm
|
||||
int pattern;
|
||||
int nChannels;
|
||||
int sameChannels;
|
||||
float speedIntra;
|
||||
float speedInter;
|
||||
float bwIntra;
|
||||
float bwInter;
|
||||
int typeIntra;
|
||||
int typeInter;
|
||||
};
|
||||
@@ -900,22 +917,22 @@ static ncclResult_t initTransportsRank(struct ncclComm* comm, ncclUniqueId* comm
|
||||
allGather3Data[rank].tree.pattern = treeGraph.pattern;
|
||||
allGather3Data[rank].tree.nChannels = treeGraph.nChannels;
|
||||
allGather3Data[rank].tree.sameChannels = treeGraph.sameChannels;
|
||||
allGather3Data[rank].tree.speedIntra = treeGraph.speedIntra;
|
||||
allGather3Data[rank].tree.speedInter = treeGraph.speedInter;
|
||||
allGather3Data[rank].tree.bwIntra = treeGraph.bwIntra;
|
||||
allGather3Data[rank].tree.bwInter = treeGraph.bwInter;
|
||||
allGather3Data[rank].tree.typeIntra = treeGraph.typeIntra;
|
||||
allGather3Data[rank].tree.typeInter = treeGraph.typeInter;
|
||||
allGather3Data[rank].ring.pattern = ringGraph.pattern;
|
||||
allGather3Data[rank].ring.nChannels = ringGraph.nChannels;
|
||||
allGather3Data[rank].ring.sameChannels = ringGraph.sameChannels;
|
||||
allGather3Data[rank].ring.speedIntra = ringGraph.speedIntra;
|
||||
allGather3Data[rank].ring.speedInter = ringGraph.speedInter;
|
||||
allGather3Data[rank].ring.bwIntra = ringGraph.bwIntra;
|
||||
allGather3Data[rank].ring.bwInter = ringGraph.bwInter;
|
||||
allGather3Data[rank].ring.typeIntra = ringGraph.typeIntra;
|
||||
allGather3Data[rank].ring.typeInter = ringGraph.typeInter;
|
||||
allGather3Data[rank].collNet.pattern = collNetGraph.pattern;
|
||||
allGather3Data[rank].collNet.nChannels = collNetGraph.nChannels;
|
||||
allGather3Data[rank].collNet.sameChannels = collNetGraph.sameChannels;
|
||||
allGather3Data[rank].collNet.speedIntra = collNetGraph.speedIntra;
|
||||
allGather3Data[rank].collNet.speedInter = collNetGraph.speedInter;
|
||||
allGather3Data[rank].collNet.bwIntra = collNetGraph.bwIntra;
|
||||
allGather3Data[rank].collNet.bwInter = collNetGraph.bwInter;
|
||||
allGather3Data[rank].collNet.typeIntra = collNetGraph.typeIntra;
|
||||
allGather3Data[rank].collNet.typeInter = collNetGraph.typeInter;
|
||||
allGather3Data[rank].collNetSupport = comm->collNetSupport;
|
||||
@@ -990,20 +1007,20 @@ static ncclResult_t initTransportsRank(struct ncclComm* comm, ncclUniqueId* comm
|
||||
// Make sure we align all ranks so that the tuning is consistent across ranks
|
||||
treeGraph.nChannels = std::min(allGather3Data[i].tree.nChannels, treeGraph.nChannels);
|
||||
treeGraph.sameChannels = std::min(allGather3Data[i].tree.sameChannels, treeGraph.sameChannels);
|
||||
treeGraph.speedIntra = std::min(allGather3Data[i].tree.speedIntra, treeGraph.speedIntra);
|
||||
treeGraph.speedInter = std::min(allGather3Data[i].tree.speedInter, treeGraph.speedInter);
|
||||
treeGraph.bwIntra = std::min(allGather3Data[i].tree.bwIntra, treeGraph.bwIntra);
|
||||
treeGraph.bwInter = std::min(allGather3Data[i].tree.bwInter, treeGraph.bwInter);
|
||||
treeGraph.typeIntra = std::max(allGather3Data[i].tree.typeIntra, treeGraph.typeIntra);
|
||||
treeGraph.typeInter = std::max(allGather3Data[i].tree.typeInter, treeGraph.typeInter);
|
||||
ringGraph.nChannels = std::min(allGather3Data[i].ring.nChannels, ringGraph.nChannels);
|
||||
ringGraph.sameChannels = std::min(allGather3Data[i].ring.sameChannels, ringGraph.sameChannels);
|
||||
ringGraph.speedIntra = std::min(allGather3Data[i].ring.speedIntra, ringGraph.speedIntra);
|
||||
ringGraph.speedInter = std::min(allGather3Data[i].ring.speedInter, ringGraph.speedInter);
|
||||
ringGraph.bwIntra = std::min(allGather3Data[i].ring.bwIntra, ringGraph.bwIntra);
|
||||
ringGraph.bwInter = std::min(allGather3Data[i].ring.bwInter, ringGraph.bwInter);
|
||||
ringGraph.typeIntra = std::max(allGather3Data[i].ring.typeIntra, ringGraph.typeIntra);
|
||||
ringGraph.typeInter = std::max(allGather3Data[i].ring.typeInter, ringGraph.typeInter);
|
||||
collNetGraph.nChannels = std::min(allGather3Data[i].collNet.nChannels, collNetGraph.nChannels);
|
||||
collNetGraph.sameChannels = std::min(allGather3Data[i].collNet.sameChannels, collNetGraph.sameChannels);
|
||||
collNetGraph.speedIntra = std::min(allGather3Data[i].collNet.speedIntra, collNetGraph.speedIntra);
|
||||
collNetGraph.speedInter = std::min(allGather3Data[i].collNet.speedInter, collNetGraph.speedInter);
|
||||
collNetGraph.bwIntra = std::min(allGather3Data[i].collNet.bwIntra, collNetGraph.bwIntra);
|
||||
collNetGraph.bwInter = std::min(allGather3Data[i].collNet.bwInter, collNetGraph.bwInter);
|
||||
collNetGraph.typeIntra = std::max(allGather3Data[i].collNet.typeIntra, collNetGraph.typeIntra);
|
||||
collNetGraph.typeInter = std::max(allGather3Data[i].collNet.typeInter, collNetGraph.typeInter);
|
||||
comm->collNetSupport = std::min(allGather3Data[i].collNetSupport, comm->collNetSupport);
|
||||
@@ -1145,16 +1162,38 @@ static ncclResult_t initTransportsRank(struct ncclComm* comm, ncclUniqueId* comm
|
||||
NCCLCHECKGOTO(ncclTransportCollNetCheck(comm, collNetSetupFail), ret, collnet_cleanup);
|
||||
TRACE(NCCL_INIT, "rank %d Connected inter-node CollNet", rank);
|
||||
|
||||
// Connect intra-node CollNet
|
||||
char line[1024];
|
||||
line[0]='\0';
|
||||
for (int c=0; c<comm->nChannels; c++) {
|
||||
struct ncclTree* chain = &comm->channels[c].collnetChain;
|
||||
snprintf(line+strlen(line), 1023-strlen(line), " [%d] %d->%d->%d",
|
||||
c, chain->down[0], rank, chain->up);
|
||||
}
|
||||
line[1023] = '\0';
|
||||
INFO(NCCL_INIT, "Collnet Chains %s", line);
|
||||
// Connect Collnet + chain
|
||||
for (int c=0; c<comm->nChannels; c++) {
|
||||
struct ncclChannel* channel = comm->channels+c;
|
||||
NCCLCHECKGOTO(ncclTransportP2pConnect(comm, c, 1, &channel->collnetChain.up, 1, channel->collnetChain.down, 0), ret, collnet_cleanup);
|
||||
}
|
||||
NCCLCHECKGOTO(ncclTransportP2pSetup(comm, &collNetGraph, 0), ret, collnet_cleanup);
|
||||
for (int c=0; c<comm->nChannels; c++) {
|
||||
struct ncclChannel* channel = comm->channels+c;
|
||||
NCCLCHECKGOTO(ncclTransportP2pConnect(comm, c, 1, channel->collnetChain.down, 1, &channel->collnetChain.up, 1), ret, collnet_cleanup);
|
||||
}
|
||||
NCCLCHECKGOTO(ncclTransportP2pSetup(comm, &collNetGraph, 1), ret, collnet_cleanup);
|
||||
INFO(NCCL_INIT, "Connected collnet + chain");
|
||||
|
||||
// Connect intra-node CollNet + Direct
|
||||
int highestTransportType0, highestTransportType1;
|
||||
for (int c=0; c<comm->nChannels; c++) {
|
||||
struct ncclChannel* channelRecv = comm->channels+c;
|
||||
NCCLCHECKGOTO(ncclTransportP2pConnect(comm, c, NCCL_MAX_DIRECT_ARITY, channelRecv->collTree.up, NCCL_MAX_DIRECT_ARITY, channelRecv->collTree.down, 0), ret, collnet_cleanup);
|
||||
NCCLCHECKGOTO(ncclTransportP2pConnect(comm, c, NCCL_MAX_DIRECT_ARITY, channelRecv->collnetDirect.up, NCCL_MAX_DIRECT_ARITY, channelRecv->collnetDirect.down, 0), ret, collnet_cleanup);
|
||||
}
|
||||
NCCLCHECKGOTO(ncclTransportP2pSetup(comm, &collNetGraph, 0, &highestTransportType0), ret, collnet_cleanup);
|
||||
for (int c=0; c<comm->nChannels; c++) {
|
||||
struct ncclChannel* channelSend = comm->channels+c;
|
||||
NCCLCHECKGOTO(ncclTransportP2pConnect(comm, c, NCCL_MAX_DIRECT_ARITY, channelSend->collTree.down, NCCL_MAX_DIRECT_ARITY, channelSend->collTree.up, 1), ret, collnet_cleanup);
|
||||
NCCLCHECKGOTO(ncclTransportP2pConnect(comm, c, NCCL_MAX_DIRECT_ARITY, channelSend->collnetDirect.down, NCCL_MAX_DIRECT_ARITY, channelSend->collnetDirect.up, 1), ret, collnet_cleanup);
|
||||
}
|
||||
NCCLCHECKGOTO(ncclTransportP2pSetup(comm, &collNetGraph, 1, &highestTransportType1), ret, collnet_cleanup);
|
||||
|
||||
@@ -1331,6 +1370,8 @@ collnet_cleanup:
|
||||
}
|
||||
}
|
||||
|
||||
NCCLCHECKGOTO(devCommSetup(comm), ret, affinity_restore);
|
||||
|
||||
/* Local intra-node barrier */
|
||||
NCCLCHECK(bootstrapBarrier(comm->bootstrap, comm->localRankToRank, comm->localRank, comm->localRanks, comm->localRankToRank[0]));
|
||||
|
||||
@@ -1358,9 +1399,15 @@ struct ncclCommInitRankAsyncJob {
|
||||
int virtualId;
|
||||
};
|
||||
|
||||
struct ncclCommFinalizeAsyncJob {
|
||||
struct ncclAsyncJob base;
|
||||
ncclComm_t comm;
|
||||
};
|
||||
|
||||
static ncclResult_t ncclCommInitRankFunc(struct ncclAsyncJob* job_) {
|
||||
struct ncclCommInitRankAsyncJob* job = (struct ncclCommInitRankAsyncJob*)job_;
|
||||
ncclComm_t* newcomm = job->newcomm;
|
||||
ncclComm_t comm = *newcomm;
|
||||
int nranks = job->nranks;
|
||||
ncclUniqueId commId = job->commId; // C++ struct assignment
|
||||
int myrank = job->myrank;
|
||||
@@ -1375,60 +1422,86 @@ static ncclResult_t ncclCommInitRankFunc(struct ncclAsyncJob* job_) {
|
||||
TRACE(NCCL_INIT, "Setting cudaLimitStackSize to %zi", maxLocalSizeBytes);
|
||||
//CUDACHECKIGNORE(hipDeviceSetLimit(hipLimitStackSize, maxLocalSizeBytes));
|
||||
}
|
||||
*newcomm = NULL;
|
||||
NCCLCHECKGOTO(commAlloc(newcomm, nranks, myrank, virtualId), res, cleanup);
|
||||
NCCLCHECKGOTO(initTransportsRank(*newcomm, &commId), res, cleanup);
|
||||
NCCLCHECKGOTO(devCommSetup(*newcomm), res, cleanup);
|
||||
|
||||
// update communicator state
|
||||
comm->initState = ncclSuccess;
|
||||
|
||||
INFO(NCCL_INIT,"comm %p rank %d nranks %d cudaDev %d busId %lx localSize %ld used %ld bytes - Init COMPLETE", *newcomm, myrank, nranks, (*newcomm)->cudaDev, (*newcomm)->busId, ncclKernLocalSize(ncclGetKernelIndex(*newcomm)), allocTracker[(*newcomm)->cudaDev].totalAllocSize);
|
||||
TRACE_CALL("ncclCommInitRank(%p,%d,0x%llx,%d,%d)", *newcomm, nranks, (unsigned long long)hashUniqueId(commId), myrank, (*newcomm)->cudaDev);
|
||||
return ncclSuccess;
|
||||
cleanup:
|
||||
if ((*newcomm) && (*newcomm)->bootstrap) bootstrapAbort((*newcomm)->bootstrap);
|
||||
*newcomm = NULL;
|
||||
comm->initState = res;
|
||||
return res;
|
||||
}
|
||||
|
||||
static ncclResult_t parseCommConfig(ncclComm_t comm, ncclConfig_t *config) {
|
||||
ncclResult_t ret = ncclSuccess;
|
||||
|
||||
/* first set configuration */
|
||||
if (config) {
|
||||
comm->blocking = config->blocking;
|
||||
} else {
|
||||
/* default setting of communicator */
|
||||
comm->blocking = 1;
|
||||
}
|
||||
|
||||
return ret;
|
||||
}
|
||||
|
||||
static void ncclCommInitRankUndo(struct ncclAsyncJob* job_) {
|
||||
struct ncclCommInitRankAsyncJob* job = (struct ncclCommInitRankAsyncJob*)job_;
|
||||
ncclCommDestroy(*job->newcomm);
|
||||
*job->newcomm = nullptr;
|
||||
}
|
||||
|
||||
static ncclResult_t ncclCommInitRankDev(ncclComm_t* newcomm, int nranks, ncclUniqueId commId, int myrank, int cudaDev, int virtualId) {
|
||||
static ncclResult_t ncclCommInitRankDev(ncclComm_t* newcomm, int nranks, ncclUniqueId commId, int myrank, int cudaDev, ncclConfig_t *config, int virtualId) {
|
||||
ncclResult_t res;
|
||||
ncclComm_t comm = NULL;
|
||||
struct ncclCommInitRankAsyncJob *job = NULL;
|
||||
char* env = getenv("NCCL_COMM_ID");
|
||||
if (env && myrank == 0) {
|
||||
INFO(NCCL_ENV, "NCCL_COMM_ID set by environment to %s", env);
|
||||
NCCLCHECKGOTO(bootstrapCreateRoot(&commId, true), res, end);
|
||||
NCCLCHECKGOTO(bootstrapCreateRoot(&commId, true), res, fail);
|
||||
}
|
||||
|
||||
NCCLCHECKGOTO(ncclInit(), res, end);
|
||||
NCCLCHECKGOTO(ncclInit(), res, fail);
|
||||
if (myrank == 0) showVersion();
|
||||
|
||||
memset(allocTracker+cudaDev, 0, sizeof(struct allocationTracker));
|
||||
// Make sure the CUDA runtime is initialized.
|
||||
CUDACHECKGOTO(hipFree(NULL), res, end);
|
||||
CUDACHECKGOTO(hipFree(NULL), res, fail);
|
||||
|
||||
NCCLCHECKGOTO(PtrCheck(newcomm, "CommInitRank", "newcomm"), res, end);
|
||||
NCCLCHECKGOTO(PtrCheck(newcomm, "CommInitRank", "newcomm"), res, fail);
|
||||
if (nranks < 1 || myrank < 0 || myrank >= nranks) {
|
||||
WARN("Invalid rank requested : %d/%d", myrank, nranks);
|
||||
res = ncclInvalidArgument;
|
||||
goto end;
|
||||
goto fail;
|
||||
}
|
||||
|
||||
struct ncclCommInitRankAsyncJob *job;
|
||||
NCCLCHECKGOTO(ncclCalloc(&job, 1), res, end);
|
||||
NCCLCHECKGOTO(ncclCalloc(&comm, 1), res, fail);
|
||||
NCCLCHECKGOTO(ncclCudaHostCalloc((uint32_t**)&comm->abortFlag, 1), res, fail);
|
||||
// set up comm state and abortFlag only
|
||||
*comm->abortFlag = 0;
|
||||
NCCLCHECKGOTO(parseCommConfig(comm, config), res, fail);
|
||||
/* start with ncclInternalError and will be changed to ncclSuccess if init succeeds. */
|
||||
comm->initState = ncclInternalError;
|
||||
*newcomm = comm;
|
||||
|
||||
NCCLCHECKGOTO(ncclCalloc(&job, 1), res, fail);
|
||||
job->newcomm = newcomm;
|
||||
job->nranks = nranks;
|
||||
job->commId = commId; // C++ struct assignment
|
||||
job->myrank = myrank;
|
||||
job->cudaDev = cudaDev;
|
||||
job->virtualId = virtualId;
|
||||
NCCLCHECKGOTO(ncclAsyncLaunch(&job->base, ncclCommInitRankFunc, ncclCommInitRankUndo, free), res, end);
|
||||
NCCLCHECKGOTO(ncclAsyncLaunch(&job->base, ncclCommInitRankFunc, NULL, free, comm), res, fail);
|
||||
|
||||
end:
|
||||
exit:
|
||||
return ncclGroupErrCheck(res);
|
||||
fail:
|
||||
goto exit;
|
||||
}
|
||||
|
||||
NCCL_API(ncclResult_t, ncclCommInitRank, ncclComm_t* newcomm, int nranks, ncclUniqueId commId, int myrank);
|
||||
@@ -1440,7 +1513,7 @@ ncclResult_t ncclCommInitRank(ncclComm_t* newcomm, int nranks, ncclUniqueId comm
|
||||
|
||||
int cudaDev;
|
||||
CUDACHECK(hipGetDevice(&cudaDev));
|
||||
NCCLCHECK(ncclCommInitRankDev(newcomm, nranks, commId, myrank, cudaDev, -1));
|
||||
NCCLCHECK(ncclCommInitRankDev(newcomm, nranks, commId, myrank, cudaDev, NULL, -1));
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
@@ -1449,7 +1522,7 @@ ncclResult_t ncclCommInitRankMulti(ncclComm_t* newcomm, int nranks, ncclUniqueId
|
||||
NVTX3_FUNC_RANGE_IN(nccl_domain);
|
||||
int cudaDev;
|
||||
CUDACHECK(hipGetDevice(&cudaDev));
|
||||
NCCLCHECK(ncclCommInitRankDev(newcomm, nranks, commId, myrank, cudaDev, virtualId));
|
||||
NCCLCHECK(ncclCommInitRankDev(newcomm, nranks, commId, myrank, cudaDev, NULL, virtualId));
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
@@ -1457,51 +1530,164 @@ ncclResult_t ncclCommInitRankMulti(ncclComm_t* newcomm, int nranks, ncclUniqueId
|
||||
NCCL_API(ncclResult_t, ncclCommInitAll, ncclComm_t* comms, int ndev, const int* devlist);
|
||||
ncclResult_t ncclCommInitAll(ncclComm_t* comms, int ndev, const int* devlist) {
|
||||
NVTX3_FUNC_RANGE_IN(nccl_domain);
|
||||
|
||||
ncclResult_t ret = ncclSuccess;
|
||||
int totalnDev;
|
||||
int *gpuFlags = NULL;
|
||||
// Load the CUDA driver and dlsym hooks (can fail on old drivers)
|
||||
(void) rocmLibraryInit();
|
||||
if (ncclParamDmaBufEnable()) (void) rocmLibraryInit();
|
||||
|
||||
NCCLCHECK(PtrCheck(comms, "CommInitAll", "comms"));
|
||||
NCCLCHECKGOTO(PtrCheck(comms, "CommInitAll", "comms"), ret, fail);
|
||||
if (ndev < 0) {
|
||||
WARN("Invalid device count requested : %d", ndev);
|
||||
return ncclInvalidArgument;
|
||||
ret = ncclInvalidArgument;
|
||||
goto fail;
|
||||
}
|
||||
|
||||
CUDACHECKGOTO(hipGetDeviceCount(&totalnDev), ret, fail);
|
||||
if (devlist) {
|
||||
NCCLCHECKGOTO(ncclCalloc(&gpuFlags, totalnDev), ret, fail);
|
||||
for (int i = 0; i < ndev; ++i) {
|
||||
/* invalid device check. */
|
||||
if (devlist[i] < 0 || devlist[i] >= totalnDev) {
|
||||
ret = ncclUnhandledCudaError;
|
||||
goto fail;
|
||||
}
|
||||
|
||||
/* duplicate device check. */
|
||||
if (gpuFlags[devlist[i]] != 0) {
|
||||
ret = ncclInvalidUsage;
|
||||
goto fail;
|
||||
}
|
||||
|
||||
gpuFlags[devlist[i]] = 1;
|
||||
}
|
||||
free(gpuFlags);
|
||||
}
|
||||
|
||||
ncclUniqueId uniqueId;
|
||||
NCCLCHECK(ncclGetUniqueId(&uniqueId));
|
||||
NCCLCHECK(ncclGroupStart());
|
||||
NCCLCHECKGOTO(ncclGetUniqueId(&uniqueId), ret, fail);
|
||||
NCCLCHECKGOTO(ncclGroupStart(), ret, fail);
|
||||
for (int i=0; i<ndev; i++) {
|
||||
// Ignore return codes .. we need to call ncclGroupEnd to clean up anyway
|
||||
ncclCommInitRankDev(comms+i, ndev, uniqueId, i, devlist ? devlist[i] : i, -1);
|
||||
ncclCommInitRankDev(comms+i, ndev, uniqueId, i, devlist ? devlist[i] : i, NULL, -1);
|
||||
}
|
||||
NCCLCHECK(ncclGroupEnd());
|
||||
return ncclSuccess;
|
||||
NCCLCHECKGOTO(ncclGroupEnd(), ret, fail);
|
||||
|
||||
exit:
|
||||
return ret;
|
||||
fail:
|
||||
if (gpuFlags) free(gpuFlags);
|
||||
goto exit;
|
||||
}
|
||||
|
||||
static ncclResult_t commDestroy(ncclComm_t comm) {
|
||||
// Try and prevent a double free of the comm struct (user error)
|
||||
if (comm->rank == -1 || comm->nRanks <= 0 || comm->cudaDev == -1 || comm->busId == -1) {
|
||||
WARN("comm %p has already been destroyed", comm);
|
||||
ncclResult_t ncclCommSetAsyncError(ncclComm_t comm, ncclResult_t nextState) {
|
||||
if (nextState < 0 || nextState >= ncclNumResults || comm == NULL) {
|
||||
WARN("ncclCommSetAsyncError: error comm %p sets state %d", comm, nextState);
|
||||
return ncclInvalidArgument;
|
||||
}
|
||||
|
||||
__atomic_store_n(&comm->asyncResult, nextState, __ATOMIC_RELEASE);
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
NCCL_API(ncclResult_t, ncclCommInitRankConfig, ncclComm_t* comm, int nranks, ncclUniqueId commId, int myrank, ncclConfig_t *config);
|
||||
ncclResult_t ncclCommInitRankConfig(ncclComm_t *newcomm, int nranks, ncclUniqueId commId, int myrank, ncclConfig_t *config) {
|
||||
NVTX3_FUNC_RANGE_IN(nccl_domain);
|
||||
int cudaDev;
|
||||
ncclResult_t ret = ncclSuccess;
|
||||
ncclConfig_t internalConfig = NCCL_CONFIG_INITIALIZER;
|
||||
ncclConfig_t *internalConfigPtr;
|
||||
size_t realSize;
|
||||
int blockingEnv;
|
||||
|
||||
NCCLCHECK(ncclGroupStartInternal());
|
||||
internalConfigPtr = &internalConfig;
|
||||
if (config) {
|
||||
memcpy((void*)&realSize, (void*)config, sizeof(size_t));
|
||||
realSize = realSize > sizeof(ncclConfig_t) ? sizeof(ncclConfig_t) : realSize;
|
||||
memcpy((void*)internalConfigPtr, (void*)config, realSize);
|
||||
if (internalConfigPtr->magic != 0xcafebeef) {
|
||||
WARN("ncclConfig_t argument not initialized via NCCL_CONFIG_INITIALIZER");
|
||||
ret = ncclInvalidArgument;
|
||||
goto exit;
|
||||
}
|
||||
}
|
||||
|
||||
/* check input config attributes */
|
||||
if (internalConfigPtr->blocking != 0 && internalConfigPtr->blocking != 1) {
|
||||
WARN("Invalid config blocking attribute value %d", internalConfigPtr->blocking);
|
||||
ret = ncclInvalidArgument;
|
||||
goto exit;
|
||||
}
|
||||
|
||||
/* overwrite configuration from env variable. */
|
||||
blockingEnv = ncclParamCommBlocking();
|
||||
if (blockingEnv != 0 && blockingEnv != 1) {
|
||||
WARN("Invalid NCCL_COMM_BLOCKING value %d", blockingEnv);
|
||||
}
|
||||
if (blockingEnv == 1) internalConfigPtr->blocking = blockingEnv;
|
||||
|
||||
if (ncclParamDmaBufEnable()) (void) rocmLibraryInit();
|
||||
CUDACHECKGOTO(hipGetDevice(&cudaDev), ret, exit);
|
||||
NCCLCHECKGOTO(ncclCommInitRankDev(newcomm, nranks, commId, myrank, cudaDev, internalConfigPtr, -1), ret, fail);
|
||||
|
||||
exit:
|
||||
ncclGroupErrCheck(ret);
|
||||
NCCLCHECK(ncclGroupEndInternal());
|
||||
if (newcomm && *newcomm && !(*newcomm)->blocking) (void) ncclCommGetAsyncError(*newcomm, &ret);
|
||||
return ret;
|
||||
fail:
|
||||
if (newcomm && *newcomm && !(*newcomm)->blocking) (void) ncclCommSetAsyncError(*newcomm, ret);
|
||||
goto exit;
|
||||
}
|
||||
|
||||
static ncclResult_t commDestroySync(struct ncclAsyncJob* job_) {
|
||||
struct ncclCommFinalizeAsyncJob* job = (struct ncclCommFinalizeAsyncJob*) job_;
|
||||
ncclComm_t comm = job->comm;
|
||||
int savedDevice;
|
||||
#ifdef ENABLE_TRACE
|
||||
int rank = comm->rank;
|
||||
#endif
|
||||
CUDACHECK(hipGetDevice(&savedDevice));
|
||||
int commDevice = comm->cudaDev;
|
||||
ncclResult_t ret;
|
||||
|
||||
CUDACHECKGOTO(hipGetDevice(&savedDevice), ret, fail);
|
||||
if (savedDevice != commDevice) {
|
||||
CUDACHECKGOTO(hipSetDevice(commDevice), ret, fail);
|
||||
}
|
||||
|
||||
TRACE(NCCL_INIT, "Destroying comm %p rank %d abortFlag %d asyncResult %d", comm, comm->rank, *comm->abortFlag, comm->asyncResult);
|
||||
|
||||
if (comm->initState == ncclSuccess) {
|
||||
NCCLCHECKGOTO(ncclStrongStreamSynchronize(&comm->hostStream), ret, fail);
|
||||
NCCLCHECKGOTO(ncclStrongStreamSynchronize(&comm->deviceStream), ret, fail);
|
||||
}
|
||||
NCCLCHECKGOTO(ncclCommPollCallbacks(comm, false), ret, fail);
|
||||
// And keep polling until all graphs referencing us die.
|
||||
while (comm->persistentRefs != 0) {
|
||||
NCCLCHECKGOTO(ncclCommPollCallbacks(comm, /*waitSome=*/true), ret, fail);
|
||||
}
|
||||
|
||||
if (savedDevice != commDevice) {
|
||||
CUDACHECKGOTO(hipSetDevice(savedDevice), ret, fail);
|
||||
}
|
||||
|
||||
exit:
|
||||
return ret;
|
||||
fail:
|
||||
goto exit;
|
||||
}
|
||||
|
||||
static ncclResult_t commCleanup(ncclComm_t comm) {
|
||||
int savedDevice;
|
||||
int commDevice = comm->cudaDev;
|
||||
|
||||
CUDACHECK(hipGetDevice(&savedDevice));
|
||||
if (savedDevice != commDevice) {
|
||||
CUDACHECK(hipSetDevice(commDevice));
|
||||
}
|
||||
|
||||
TRACE(NCCL_INIT, "Destroying comm %p rank %d abortFlag %d fatalError %d", comm, comm->rank, *comm->abortFlag, comm->fatalError);
|
||||
|
||||
NCCLCHECK(ncclStrongStreamSynchronize(&comm->hostStream));
|
||||
NCCLCHECK(ncclStrongStreamSynchronize(&comm->deviceStream));
|
||||
NCCLCHECK(ncclCommPollCallbacks(comm));
|
||||
|
||||
NCCLCHECK(commFree(comm));
|
||||
|
||||
if (savedDevice != commDevice)
|
||||
@@ -1521,6 +1707,125 @@ static ncclResult_t commDestroy(ncclComm_t comm) {
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
static ncclResult_t commFinalize(ncclComm_t comm, bool userCalled) {
|
||||
ncclResult_t ret = ncclSuccess;
|
||||
struct ncclCommFinalizeAsyncJob *job = NULL;
|
||||
|
||||
comm->finalizeCalled = true;
|
||||
/* launch async thread to finalize comm. */
|
||||
NCCLCHECKGOTO(ncclCalloc(&job, 1), ret, fail);
|
||||
job->comm = comm;
|
||||
|
||||
if (userCalled) {
|
||||
NCCLCHECKGOTO(ncclAsyncLaunch(&job->base, commDestroySync, NULL, free, comm), ret, fail);
|
||||
} else {
|
||||
NCCLCHECKGOTO(commDestroySync(&job->base), ret, fail);
|
||||
free(job);
|
||||
}
|
||||
|
||||
exit:
|
||||
return ncclGroupErrCheck(ret);
|
||||
fail:
|
||||
if (job) free(job);
|
||||
goto exit;
|
||||
}
|
||||
|
||||
NCCL_API(ncclResult_t, ncclCommFinalize, ncclComm_t comm);
|
||||
ncclResult_t ncclCommFinalize(ncclComm_t comm) {
|
||||
NVTX3_FUNC_RANGE_IN(nccl_domain);
|
||||
ncclResult_t ret = ncclSuccess;
|
||||
|
||||
NCCLCHECK(ncclGroupStartInternal());
|
||||
if (comm == NULL) goto exit;
|
||||
|
||||
/* wait comm ready before finalize. */
|
||||
NCCLCHECKGOTO(ncclCommEnsureReady(comm), ret, fail);
|
||||
|
||||
/* prevent double finalize. */
|
||||
if (comm->finalizeCalled) {
|
||||
ret = ncclInvalidArgument;
|
||||
goto fail;
|
||||
}
|
||||
|
||||
/* finalize comm. */
|
||||
ret = commFinalize(comm, true);
|
||||
|
||||
exit:
|
||||
ncclGroupErrCheck(ret);
|
||||
NCCLCHECK(ncclGroupEndInternal());
|
||||
if (comm && !comm->blocking) { NCCLCHECK(ncclCommGetAsyncError(comm, &ret)) };
|
||||
return ret;
|
||||
fail:
|
||||
if (comm && !comm->blocking) (void) ncclCommSetAsyncError(comm, ret);
|
||||
goto exit;
|
||||
}
|
||||
|
||||
static ncclResult_t commReclaim(ncclComm_t comm) {
|
||||
ncclResult_t ret = ncclSuccess;
|
||||
ncclResult_t state;
|
||||
int curRank; /* Debug info */
|
||||
|
||||
NCCLCHECKGOTO(ncclCommGetAsyncError(comm, &state), ret, fail);
|
||||
TRACE(NCCL_INIT, "commReclaim: reclaim comm %p rank %d state %d", comm, comm->rank, state);
|
||||
if (state == ncclSuccess && *comm->abortFlag == 0 && comm->finalizeCalled == false) {
|
||||
/* user does not call ncclCommFinalize and this is a normal comm destroy. ncclCommDestroy
|
||||
* should be nonblocking until last call of ncclCommDestroy. */
|
||||
NCCLCHECKGOTO(commFinalize(comm, false), ret, fail);
|
||||
}
|
||||
|
||||
if (comm->initState != ncclSuccess) {
|
||||
/* if init errors happen, no finalize thread should have been launched. Main thread can reclaim
|
||||
* everything since no NCCL kernel was issued. */
|
||||
struct ncclCommFinalizeAsyncJob job;
|
||||
|
||||
job.comm = comm;
|
||||
curRank = comm->rank;
|
||||
/* comm aborts, commDestroySync should not be blocked. */
|
||||
if ((ret = commDestroySync((struct ncclAsyncJob*) &job)) != ncclSuccess) {
|
||||
WARN("commReclaim: comm %p (rank = %d) in abort, error %d", comm, curRank, ret);
|
||||
}
|
||||
|
||||
if ((ret = commCleanup(comm)) != ncclSuccess) {
|
||||
WARN("commReclaim: cleanup comm %p rank %d failed in destroy/abort, error %d", comm, curRank, ret);
|
||||
}
|
||||
} else {
|
||||
int curRankCnt;
|
||||
int intraRanks = comm->intraRanks;
|
||||
ncclComm_t intracomm0 = comm->intraComm0;
|
||||
int *finalizeRankCnt = &intracomm0->finalizeRankCnt;
|
||||
|
||||
assert(intracomm0 != NULL && finalizeRankCnt != NULL);
|
||||
curRankCnt = __atomic_add_fetch(finalizeRankCnt, 1, __ATOMIC_ACQ_REL);
|
||||
if (curRankCnt == intraRanks) {
|
||||
ncclComm_t curIntraComm;
|
||||
ncclComm_t nextIntraComm = intracomm0;
|
||||
|
||||
while (nextIntraComm) {
|
||||
curIntraComm = nextIntraComm;
|
||||
curRank = curIntraComm->rank;
|
||||
nextIntraComm = nextIntraComm->intraNext;
|
||||
|
||||
if (comm->finalizeCalled == false) {
|
||||
struct ncclCommFinalizeAsyncJob job;
|
||||
job.comm = curIntraComm;
|
||||
/* every comm aborts, commDestroySync should not be blocked. */
|
||||
if ((ret = commDestroySync((struct ncclAsyncJob*) &job)) != ncclSuccess)
|
||||
WARN("commReclaim: comm %p (rank = %d) in abort, error %d", curIntraComm, curRank, ret);
|
||||
}
|
||||
|
||||
if ((ret = commCleanup(curIntraComm)) != ncclSuccess) {
|
||||
WARN("commReclaim: cleanup comm %p rank %d failed in destroy/abort, error %d", curIntraComm, curRank, ret);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
exit:
|
||||
return ret;
|
||||
fail:
|
||||
goto exit;
|
||||
}
|
||||
|
||||
NCCL_API(ncclResult_t, ncclCommDestroy, ncclComm_t comm);
|
||||
ncclResult_t ncclCommDestroy(ncclComm_t comm) {
|
||||
NVTX3_FUNC_RANGE_IN(nccl_domain);
|
||||
@@ -1530,9 +1835,18 @@ ncclResult_t ncclCommDestroy(ncclComm_t comm) {
|
||||
int rank = comm->rank, nranks = comm->nRanks, cudaDev = comm->cudaDev;
|
||||
int64_t busId = comm->busId;
|
||||
TRACE(NCCL_INIT, "comm %p rank %d nRanks %d cudaDev %d busId %lx", comm, rank, nranks, cudaDev, busId);
|
||||
// Try and prevent a double free of the comm struct (user error)
|
||||
if (comm->rank == -1 || comm->nRanks == -1 || comm->cudaDev == -1 || comm->busId == -1) {
|
||||
WARN("comm %p has already been destroyed", comm);
|
||||
return ncclInvalidArgument;
|
||||
}
|
||||
|
||||
NCCLCHECK(commDestroy(comm));
|
||||
/* init thread must be joined before we destroy the comm. */
|
||||
NCCLCHECK(ncclCommEnsureReady(comm));
|
||||
|
||||
NCCLCHECK(commReclaim(comm));
|
||||
INFO(NCCL_INIT,"comm %p rank %d nranks %d cudaDev %d busId %lx - Destroy COMPLETE", comm, rank, nranks, cudaDev, busId);
|
||||
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
@@ -1548,9 +1862,13 @@ ncclResult_t ncclCommAbort(ncclComm_t comm) {
|
||||
|
||||
// Ask anything that might still be running on the device to quit
|
||||
*comm->abortFlag = 1;
|
||||
/* init thread must be joined before we destroy the comm,
|
||||
* and we should ignore the init error here. */
|
||||
ncclCommEnsureReady(comm);
|
||||
|
||||
//NCCLCHECK(commDestroy(comm));
|
||||
(void) commReclaim(comm);
|
||||
INFO(NCCL_INIT,"comm %p rank %d nranks %d cudaDev %d busId %lx - Abort COMPLETE", comm, rank, nranks, cudaDev, busId);
|
||||
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
@@ -1564,6 +1882,7 @@ const char* ncclGetErrorString(ncclResult_t code) {
|
||||
case ncclInvalidArgument : return "invalid argument";
|
||||
case ncclInvalidUsage : return "invalid usage";
|
||||
case ncclRemoteError : return "remote process exited or there was a network error";
|
||||
case ncclInProgress : return "NCCL operation in progress";
|
||||
default : return "unknown result code";
|
||||
}
|
||||
}
|
||||
@@ -1580,15 +1899,21 @@ NCCL_API(ncclResult_t, ncclCommGetAsyncError, ncclComm_t comm, ncclResult_t *asy
|
||||
ncclResult_t ncclCommGetAsyncError(ncclComm_t comm, ncclResult_t *asyncError) {
|
||||
NCCLCHECK(PtrCheck(comm, "ncclGetAsyncError", "comm"));
|
||||
NCCLCHECK(PtrCheck(asyncError, "ncclGetAsyncError", "asyncError"));
|
||||
*asyncError = comm->fatalError;
|
||||
|
||||
*asyncError = __atomic_load_n(&comm->asyncResult, __ATOMIC_ACQUIRE);
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
NCCL_API(ncclResult_t, ncclCommCount, const ncclComm_t comm, int* count);
|
||||
ncclResult_t ncclCommCount(const ncclComm_t comm, int* count) {
|
||||
NVTX3_FUNC_RANGE_IN(nccl_domain);
|
||||
|
||||
NCCLCHECK(PtrCheck(comm, "CommCount", "comm"));
|
||||
NCCLCHECK(PtrCheck(count, "CommCount", "count"));
|
||||
|
||||
/* init thread must be joined before we access the attributes of comm. */
|
||||
NCCLCHECK(ncclCommEnsureReady(comm));
|
||||
|
||||
*count = comm->nRanks;
|
||||
return ncclSuccess;
|
||||
}
|
||||
@@ -1596,8 +1921,12 @@ ncclResult_t ncclCommCount(const ncclComm_t comm, int* count) {
|
||||
NCCL_API(ncclResult_t, ncclCommCuDevice, const ncclComm_t comm, int* devid);
|
||||
ncclResult_t ncclCommCuDevice(const ncclComm_t comm, int* devid) {
|
||||
NVTX3_FUNC_RANGE_IN(nccl_domain);
|
||||
|
||||
NCCLCHECK(PtrCheck(comm, "CommCuDevice", "comm"));
|
||||
NCCLCHECK(PtrCheck(devid, "CommCuDevice", "devid"));
|
||||
|
||||
NCCLCHECK(ncclCommEnsureReady(comm));
|
||||
|
||||
*devid = comm->cudaDev;
|
||||
return ncclSuccess;
|
||||
}
|
||||
@@ -1605,8 +1934,12 @@ ncclResult_t ncclCommCuDevice(const ncclComm_t comm, int* devid) {
|
||||
NCCL_API(ncclResult_t, ncclCommUserRank, const ncclComm_t comm, int* rank);
|
||||
ncclResult_t ncclCommUserRank(const ncclComm_t comm, int* rank) {
|
||||
NVTX3_FUNC_RANGE_IN(nccl_domain);
|
||||
|
||||
NCCLCHECK(PtrCheck(comm, "CommUserRank", "comm"));
|
||||
NCCLCHECK(PtrCheck(rank, "CommUserRank", "rank"));
|
||||
|
||||
NCCLCHECK(ncclCommEnsureReady(comm));
|
||||
|
||||
*rank = comm->rank;
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user