Merge remote-tracking branch 'nccl/v2.19' into develop
[ROCm/rccl commit: 81ddf9de89]
This commit is contained in:
+45
-32
@@ -26,7 +26,6 @@ __thread int ncclGroupBlocking = -1; /* default mode */
|
||||
__thread bool ncclGroupJobAbortFlag = false;
|
||||
|
||||
void* ncclAsyncJobMain(void* arg);
|
||||
static ncclResult_t groupJobComplete(struct ncclGroupJob *job);
|
||||
|
||||
ncclResult_t ncclAsyncLaunch(
|
||||
struct ncclAsyncJob* job,
|
||||
@@ -93,15 +92,7 @@ ncclResult_t ncclGroupStart() {
|
||||
return ret;
|
||||
}
|
||||
|
||||
inline ncclResult_t ncclGroupStartInternal() {
|
||||
/* if previous group launch does not complete, don't launch this one. */
|
||||
if (ncclGroupJobMainPtr != NULL) {
|
||||
if (__atomic_load_n(&ncclGroupJobMainPtr->doneFlag, __ATOMIC_ACQUIRE) == false) {
|
||||
return ncclInvalidUsage;
|
||||
} else {
|
||||
NCCLCHECK(groupJobComplete(ncclGroupJobMainPtr));
|
||||
}
|
||||
}
|
||||
ncclResult_t ncclGroupStartInternal() {
|
||||
ncclGroupDepth++;
|
||||
if (mscclAvailable() && !mscclIsCaller()) {
|
||||
NCCLCHECK(mscclGroupStart());
|
||||
@@ -202,9 +193,28 @@ failure:
|
||||
return result;
|
||||
}
|
||||
|
||||
static void groupCleanup(struct ncclComm** groupCommHeadPtr, struct ncclComm** groupCommPreconnectHeadPtr, struct ncclIntruQueue<struct ncclAsyncJob, &ncclAsyncJob::next>* asyncJobsPtr, ncclResult_t* groupErrorPtr, ncclResult_t error) {
|
||||
static inline void groupResetJobState(struct ncclGroupJob* job) {
|
||||
if (job) {
|
||||
if (job->groupBlockingPtr) *job->groupBlockingPtr = -1;
|
||||
if (job->abortFlagPtr) *job->abortFlagPtr = false;
|
||||
if (job->groupErrorPtr) *job->groupErrorPtr = ncclSuccess;
|
||||
if (job->groupCommHeadPtr) *job->groupCommHeadPtr = NULL;
|
||||
if (job->groupCommPreconnectHeadPtr) *job->groupCommPreconnectHeadPtr = NULL;
|
||||
memset(job, 0, sizeof(struct ncclGroupJob));
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
static void groupCleanup(struct ncclComm** groupCommHeadPtr, struct ncclComm** groupCommPreconnectHeadPtr, struct ncclIntruQueue<struct ncclAsyncJob, &ncclAsyncJob::next>* asyncJobsPtr, ncclResult_t* groupErrorPtr, int* groupBlockingPtr, volatile bool* groupJobAbortFlagPtr, ncclResult_t error) {
|
||||
struct ncclComm* comm = *groupCommHeadPtr;
|
||||
|
||||
/* reset all thread local variables */
|
||||
*groupCommHeadPtr = NULL;
|
||||
*groupCommPreconnectHeadPtr = NULL;
|
||||
*groupErrorPtr = ncclSuccess;
|
||||
*groupBlockingPtr = -1;
|
||||
*groupJobAbortFlagPtr = false;
|
||||
|
||||
while (comm != nullptr) {
|
||||
struct ncclComm* next = comm->groupNext;
|
||||
(void) ncclGroupCommLeave(comm); // overwrites comm->groupNext
|
||||
@@ -254,16 +264,12 @@ static void groupCleanup(struct ncclComm** groupCommHeadPtr, struct ncclComm** g
|
||||
/* reset everything */
|
||||
while (!ncclIntruQueueEmpty(asyncJobsPtr)) {
|
||||
struct ncclAsyncJob* job = ncclIntruQueueDequeue(asyncJobsPtr);
|
||||
*job->abortFlag = 1;
|
||||
if (job->comm && !job->comm->config.blocking)
|
||||
(void) ncclCommSetAsyncError(job->comm, error);
|
||||
if (job->undo) job->undo(job);
|
||||
if (job->destructor) job->destructor((void*)job);
|
||||
}
|
||||
|
||||
*groupErrorPtr = ncclSuccess;
|
||||
*groupCommHeadPtr = nullptr;
|
||||
*groupCommPreconnectHeadPtr = nullptr;
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -346,9 +352,6 @@ static ncclResult_t groupLaunch(struct ncclAsyncJob *job_) {
|
||||
NCCLCHECKGOTO(doLaunches(groupCommHeadMain), ret, fail);
|
||||
}
|
||||
|
||||
/* this atomic must happen before cleanup and setting state of communicators */
|
||||
__atomic_store_n(&gjob->doneFlag, true, __ATOMIC_RELEASE);
|
||||
|
||||
while (!ncclIntruQueueEmpty(asyncJobsMain)) {
|
||||
struct ncclAsyncJob* job = ncclIntruQueueDequeue(asyncJobsMain);
|
||||
if (job->comm && !job->comm->config.blocking)
|
||||
@@ -366,16 +369,12 @@ static ncclResult_t groupLaunch(struct ncclAsyncJob *job_) {
|
||||
groupCommHeadMain = next;
|
||||
}
|
||||
|
||||
*gjob->groupErrorPtr = ncclSuccess;
|
||||
*gjob->groupCommHeadPtr = nullptr;
|
||||
*gjob->groupCommPreconnectHeadPtr = nullptr;
|
||||
|
||||
CUDACHECK(cudaSetDevice(savedDev));
|
||||
|
||||
exit:
|
||||
return ret;
|
||||
fail:
|
||||
groupCleanup(gjob->groupCommHeadPtr, gjob->groupCommPreconnectHeadPtr, gjob->asyncJobsPtr, gjob->groupErrorPtr, ret);
|
||||
groupCleanup(gjob->groupCommHeadPtr, gjob->groupCommPreconnectHeadPtr, gjob->asyncJobsPtr, gjob->groupErrorPtr, gjob->groupBlockingPtr, gjob->abortFlagPtr, ret);
|
||||
goto exit;
|
||||
}
|
||||
|
||||
@@ -402,7 +401,8 @@ ncclResult_t ncclGroupEndInternal() {
|
||||
ncclGroupJobMain.groupErrorPtr = &ncclGroupError;
|
||||
ncclGroupJobMain.asyncJobsPtr = &ncclAsyncJobs;
|
||||
ncclGroupJobMain.abortFlagPtr = &ncclGroupJobAbortFlag;
|
||||
ncclGroupJobMain.doneFlag = false;
|
||||
ncclGroupJobMain.groupBlockingPtr = &ncclGroupBlocking;
|
||||
ncclGroupJobMain.initialized = true;
|
||||
ncclGroupJobMainPtr = &ncclGroupJobMain;
|
||||
/* make sure ncclGroupBlocking has been set. */
|
||||
assert(ncclGroupBlocking == 0 || ncclGroupBlocking == 1);
|
||||
@@ -412,6 +412,7 @@ ncclResult_t ncclGroupEndInternal() {
|
||||
ncclAsyncJob* job = ncclIntruQueueHead(&ncclAsyncJobs);
|
||||
do {
|
||||
NCCLCHECKGOTO(ncclCommSetAsyncError(job->comm, ncclInProgress), ret, fail);
|
||||
job->comm->groupJob = ncclGroupJobMainPtr;
|
||||
job = job->next;
|
||||
} while (job);
|
||||
}
|
||||
@@ -420,30 +421,42 @@ ncclResult_t ncclGroupEndInternal() {
|
||||
ncclComm_t comm = ncclGroupCommHead;
|
||||
do {
|
||||
NCCLCHECKGOTO(ncclCommSetAsyncError(comm, ncclInProgress), ret, fail);
|
||||
/* link group job to communicators. */
|
||||
comm->groupJob = ncclGroupJobMainPtr;
|
||||
comm = comm->groupNext;
|
||||
} while (comm);
|
||||
}
|
||||
|
||||
ncclGroupJobMainPtr->base.func = groupLaunch;
|
||||
SYSCHECKGOTO(pthread_create(&ncclGroupJobMainPtr->base.thread, NULL, ncclAsyncJobMain, (void*)&ncclGroupJobMainPtr->base), ret, fail);
|
||||
ret = ncclInProgress;
|
||||
} else {
|
||||
/* blocking group */
|
||||
NCCLCHECKGOTO(groupLaunch(&ncclGroupJobMainPtr->base), ret, fail);
|
||||
groupResetJobState();
|
||||
groupResetJobState(ncclGroupJobMainPtr);
|
||||
}
|
||||
}
|
||||
|
||||
exit:
|
||||
return ret;
|
||||
fail:
|
||||
groupCleanup(&ncclGroupCommHead, &ncclGroupCommPreconnectHead, &ncclAsyncJobs, &ncclGroupError, ret);
|
||||
groupResetJobState();
|
||||
groupCleanup(&ncclGroupCommHead, &ncclGroupCommPreconnectHead, &ncclAsyncJobs, &ncclGroupError, &ncclGroupBlocking, &ncclGroupJobAbortFlag, ret);
|
||||
goto exit;
|
||||
}
|
||||
|
||||
void ncclGroupJobAbort() {
|
||||
ncclGroupJobAbortFlag = true;
|
||||
(void) groupJobComplete(ncclGroupJobMainPtr);
|
||||
/* reset group abort flag */
|
||||
ncclGroupJobAbortFlag = false;
|
||||
ncclResult_t ncclGroupJobComplete(struct ncclGroupJob* groupJob) {
|
||||
ncclResult_t ret = ncclSuccess;
|
||||
if (groupJob && groupJob->initialized) {
|
||||
ret = ncclAsyncJobComplete(&groupJob->base);
|
||||
groupResetJobState(groupJob);
|
||||
}
|
||||
return ret;
|
||||
}
|
||||
|
||||
ncclResult_t ncclGroupJobAbort(struct ncclGroupJob* groupJob) {
|
||||
if (groupJob && groupJob->initialized) {
|
||||
*groupJob->abortFlagPtr = true;
|
||||
NCCLCHECK(ncclGroupJobComplete(groupJob));
|
||||
}
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user