Skip to content

Commit da42b75

Browse files
authored
Refactor group management logic (#284)
1 parent e49ddba commit da42b75

2 files changed

Lines changed: 80 additions & 67 deletions

File tree

flagcx/core/group.cc

Lines changed: 67 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -370,3 +370,70 @@ static flagcxResult_t groupLaunch(struct flagcxAsyncJob *job_) {
370370
fail:
371371
goto exit;
372372
}
373+
374+
static flagcxResult_t groupCleanup(struct flagcxAsyncJob *job_) {
375+
struct flagcxGroupJob *gjob = (struct flagcxGroupJob *)job_;
376+
struct flagcxHeteroComm *groupCommHeadMain = *gjob->groupCommHeadPtr;
377+
struct flagcxHeteroComm *groupCommPreconnectHeadMain =
378+
*gjob->groupCommPreconnectHeadPtr;
379+
struct flagcxIntruQueue<struct flagcxAsyncJob, &flagcxAsyncJob::next>
380+
*asyncJobsMain = gjob->asyncJobsPtr;
381+
382+
// clean up preconnect comms
383+
while (groupCommPreconnectHeadMain != nullptr) {
384+
struct flagcxHeteroComm *comm = groupCommPreconnectHeadMain;
385+
struct flagcxHeteroComm *next = comm->preconnectNext;
386+
comm->preconnectNext = reinterpret_cast<struct flagcxHeteroComm *>(0x1);
387+
groupCommPreconnectHeadMain = next;
388+
}
389+
390+
// clean up async jobs
391+
while (!flagcxIntruQueueEmpty(asyncJobsMain)) {
392+
struct flagcxAsyncJob *job = flagcxIntruQueueDequeue(asyncJobsMain);
393+
free(job);
394+
}
395+
396+
// clean up comms
397+
while (groupCommHeadMain != nullptr) {
398+
struct flagcxHeteroComm *comm = groupCommHeadMain;
399+
struct flagcxHeteroComm *next = comm->groupNext;
400+
(void)flagcxGroupCommLeave(comm);
401+
groupCommHeadMain = next;
402+
}
403+
404+
return flagcxSuccess;
405+
}
406+
407+
static inline void groupResetJobState() {
408+
flagcxGroupBlocking = 0;
409+
flagcxGroupJobMainPtr = NULL;
410+
flagcxGroupCommPreconnectHead = nullptr;
411+
flagcxGroupCommHead = nullptr;
412+
memset(&flagcxGroupJobMain, 0, sizeof(struct flagcxGroupJob));
413+
}
414+
415+
flagcxResult_t flagcxGroupEndInternal() {
416+
flagcxResult_t ret = flagcxSuccess;
417+
flagcxGroupDepth--;
418+
if (flagcxGroupDepth < 0)
419+
return flagcxSystemError;
420+
if (flagcxGroupDepth == 0) {
421+
if (flagcxGroupCommPreconnectHead || flagcxGroupCommHead) {
422+
flagcxGroupJobMain.groupCommHeadPtr = &flagcxGroupCommHead;
423+
flagcxGroupJobMain.groupCommPreconnectHeadPtr =
424+
&flagcxGroupCommPreconnectHead;
425+
flagcxGroupJobMain.asyncJobsPtr = &flagcxAsyncJobs;
426+
flagcxGroupJobMain.initialized = true;
427+
flagcxGroupJobMainPtr = &flagcxGroupJobMain;
428+
FLAGCXCHECKGOTO(groupLaunch(&flagcxGroupJobMainPtr->base), ret, fail);
429+
groupResetJobState();
430+
}
431+
}
432+
433+
exit:
434+
return ret;
435+
fail:
436+
groupCleanup(&flagcxGroupJobMainPtr->base);
437+
groupResetJobState();
438+
goto exit;
439+
}

flagcx/core/group.h

Lines changed: 13 additions & 67 deletions
Original file line numberDiff line numberDiff line change
@@ -10,13 +10,6 @@
1010
#include "assert.h"
1111
#include "comm.h"
1212

13-
flagcxResult_t flagcxGroupErrCheck(flagcxResult_t ret);
14-
void flagcxGroupCommJoin(struct flagcxHeteroComm *comm);
15-
void flagcxGroupCommPreconnect(struct flagcxHeteroComm *comm);
16-
flagcxResult_t flagcxGroupCommLeave(struct flagcxHeteroComm *comm);
17-
flagcxResult_t flagcxGroupJobAbort(struct flagcxGroupJob *groupJob);
18-
flagcxResult_t flagcxGroupJobComplete(struct flagcxGroupJob *groupJob);
19-
2013
typedef flagcxResult_t (*flagcxInitFunc_t)(flagcxHeteroComm_t *newcomm,
2114
int ndev, flagcxUniqueId commId,
2215
int myrank, int cudaDev);
@@ -79,69 +72,17 @@ extern __thread struct flagcxIntruQueue<struct flagcxAsyncJob,
7972
&flagcxAsyncJob::next>
8073
flagcxAsyncJobs;
8174

82-
static flagcxResult_t groupLaunch(struct flagcxAsyncJob *job_);
75+
flagcxResult_t flagcxGroupErrCheck(flagcxResult_t ret);
76+
void flagcxGroupCommJoin(struct flagcxHeteroComm *comm);
77+
void flagcxGroupCommPreconnect(struct flagcxHeteroComm *comm);
78+
flagcxResult_t flagcxGroupCommLeave(struct flagcxHeteroComm *comm);
79+
// Not implemented
80+
flagcxResult_t flagcxGroupJobAbort(struct flagcxGroupJob *groupJob);
81+
// Not implemented
82+
flagcxResult_t flagcxGroupJobComplete(struct flagcxGroupJob *groupJob);
8383
flagcxResult_t flagcxHeteroGroupStart();
8484
flagcxResult_t flagcxHeteroGroupEnd();
8585

86-
static inline void groupResetJobState() {
87-
flagcxGroupBlocking = 0;
88-
flagcxGroupJobMainPtr = NULL;
89-
flagcxGroupCommPreconnectHead = nullptr;
90-
flagcxGroupCommHead = nullptr;
91-
memset(&flagcxGroupJobMain, 0, sizeof(struct flagcxGroupJob));
92-
return;
93-
}
94-
95-
static inline flagcxResult_t groupJobComplete(struct flagcxGroupJob *job) {
96-
flagcxResult_t ret = flagcxSuccess;
97-
if (job) {
98-
ret = flagcxAsyncJobComplete(&job->base);
99-
groupResetJobState();
100-
}
101-
return ret;
102-
}
103-
104-
inline flagcxResult_t flagcxGroupStartInternal() {
105-
flagcxGroupDepth++;
106-
return flagcxSuccess;
107-
}
108-
109-
inline flagcxResult_t flagcxGroupEndInternal() {
110-
flagcxResult_t ret = flagcxSuccess;
111-
flagcxGroupDepth--;
112-
if (flagcxGroupDepth < 0)
113-
return flagcxSystemError;
114-
if (flagcxGroupDepth == 0) {
115-
116-
/**
117-
* TODO: do all jobs at Groups
118-
**/
119-
if (flagcxGroupCommPreconnectHead || flagcxGroupCommHead) {
120-
121-
flagcxGroupJobMain.groupCommHeadPtr = &flagcxGroupCommHead;
122-
flagcxGroupJobMain.groupCommPreconnectHeadPtr =
123-
&flagcxGroupCommPreconnectHead;
124-
flagcxGroupJobMain.asyncJobsPtr = &flagcxAsyncJobs;
125-
flagcxGroupJobMain.initialized = true;
126-
flagcxGroupJobMainPtr = &flagcxGroupJobMain;
127-
128-
FLAGCXCHECKGOTO(groupLaunch(&flagcxGroupJobMainPtr->base), ret, fail);
129-
groupResetJobState();
130-
}
131-
}
132-
133-
exit:
134-
return ret;
135-
fail:
136-
/**
137-
* TODO: add groupCleanup()
138-
**/
139-
// groupCleanup(&flagcxGroupCommHead, &flagcxGroupCommPreconnectHead,
140-
// &flagcxAsyncJobs, &flagcxGroupError, &flagcxGroupBlocking,
141-
// &flagcxGroupJobAbortFlag, ret);
142-
goto exit;
143-
}
144-
14586
inline flagcxResult_t flagcxGroupErrCheck(flagcxResult_t ret) {
14687
if (flagcxGroupDepth > 0) {
14788
if (ret != flagcxSuccess && ret != flagcxInProgress)
@@ -173,4 +114,9 @@ inline flagcxResult_t flagcxGroupCommLeave(struct flagcxHeteroComm *comm) {
173114
return flagcxSuccess;
174115
}
175116

117+
inline flagcxResult_t flagcxGroupStartInternal() {
118+
flagcxGroupDepth++;
119+
return flagcxSuccess;
120+
}
121+
176122
#endif

0 commit comments

Comments
 (0)