Skip to content

Commit d7adb10

Browse files
authored
[CRL] Add Device API multi-FIFO support (flagos-ai#447)
1 parent 301d4df commit d7adb10

14 files changed

Lines changed: 592 additions & 168 deletions

File tree

flagcx/adaptor/flagcx_device.cc

Lines changed: 180 additions & 104 deletions
Large diffs are not rendered by default.

flagcx/adaptor/include/device_api/fallback_comm_traits.h

Lines changed: 15 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -92,7 +92,7 @@ struct CommTraits<Fallback<PlatformTag>> {
9292
// Baseline
9393
int rank, nRanks;
9494
int intraRank, intraSize;
95-
void *fifoBuffer;
95+
void *fifoBuffers[FLAGCX_DEVICE_CTA_COUNT];
9696

9797
// IPC barriers
9898
uint64_t **barrierPeers;
@@ -121,8 +121,8 @@ struct CommTraits<Fallback<PlatformTag>> {
121121
}
122122
FLAGCX_DEVICE_INLINE_DECORATOR int getRank() const { return rank; }
123123
FLAGCX_DEVICE_INLINE_DECORATOR int getSize() const { return nRanks; }
124-
FLAGCX_DEVICE_INLINE_DECORATOR void *getFifoBuffer() const {
125-
return fifoBuffer;
124+
FLAGCX_DEVICE_INLINE_DECORATOR void *getFifoBuffer(int contextId) const {
125+
return fifoBuffers[contextId];
126126
}
127127

128128
// Populate from host-side handle (deferred template avoids forward-decl)
@@ -133,7 +133,8 @@ struct CommTraits<Fallback<PlatformTag>> {
133133
dc.nRanks = di.nRanks;
134134
dc.intraRank = di.intraRank;
135135
dc.intraSize = di.intraSize;
136-
dc.fifoBuffer = di.fifoBuffer;
136+
for (int i = 0; i < di.contextCount; i++)
137+
dc.fifoBuffers[i] = di.fifoBuffers[i];
137138
dc.barrierPeers = di.barrierPeers;
138139
dc.intraBarrierEpoch = di.intraBarrierEpoch;
139140
dc.nBarriers = di.nBarriers;
@@ -268,9 +269,13 @@ struct CommTraits<Fallback<PlatformTag>> {
268269

269270
FLAGCX_DEVICE_INLINE_DECORATOR
270271
Transport(const Comm &dc, int contextIndex)
271-
: _dc(dc), fifoBuffer(dc.fifoBuffer), signalBuffer(dc.signalBuffer),
272-
shadowBuffer(dc.shadowBuffer), counterBuffer(dc.counterBuffer),
273-
signalCount(dc.signalCount), counterCount(dc.counterCount) {
272+
: _dc(dc),
273+
fifoBuffer(
274+
dc.fifoBuffers[contextIndex %
275+
((dc.contextCount > 0) ? dc.contextCount : 1)]),
276+
signalBuffer(dc.signalBuffer), shadowBuffer(dc.shadowBuffer),
277+
counterBuffer(dc.counterBuffer), signalCount(dc.signalCount),
278+
counterCount(dc.counterCount) {
274279
int cnt = (dc.contextCount > 0) ? dc.contextCount : 1;
275280
contextId = contextIndex % cnt;
276281
}
@@ -837,10 +842,10 @@ struct Barrier<Fallback<P>, flagcxTeamTagInter, Coop> {
837842

838843
// Active ctor
839844
FLAGCX_DEVICE_INLINE_DECORATOR
840-
Barrier(Coop coop, const Transport &, const Comm &dc, Team, uint32_t index,
841-
int nInterPeers)
845+
Barrier(Coop coop, const Transport &trans, const Comm &dc, Team,
846+
uint32_t index, int nInterPeers)
842847
: _coop(coop), _interSignals(dc.interSignalFlags),
843-
_fifoBuffer(dc.fifoBuffer), _nInterPeers(nInterPeers),
848+
_fifoBuffer(trans.fifoBuffer), _nInterPeers(nInterPeers),
844849
_isLeader(dc.isInterLeader), _ctaIndex(index),
845850
_epoch(dc.interBarrierEpoch) {}
846851

flagcx/adaptor/include/device_api/flagcx_device.h

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -47,7 +47,9 @@ struct flagcxDevCommInternal {
4747
// ---- Baseline (always set) ----
4848
int rank, nRanks;
4949
int intraRank, intraSize;
50-
void *fifoBuffer; // Device-accessible FIFO (from heteroComm, may be null)
50+
void *fifoBuffers[FLAGCX_DEVICE_CTA_COUNT]; // Device-accessible FIFOs (one
51+
// per context, from heteroComm,
52+
// may be null)
5153
// ---- IPC barrier layer (set if IPC barrier setup succeeds, else nullptr)
5254
// ----
5355
uint64_t *
@@ -194,8 +196,8 @@ struct flagcxDevComm {
194196
FLAGCX_DEVICE_INLINE_DECORATOR int getSize() const {
195197
return _commBase.getSize();
196198
}
197-
FLAGCX_DEVICE_INLINE_DECORATOR void *getFifoBuffer() const {
198-
return _commBase.getFifoBuffer();
199+
FLAGCX_DEVICE_INLINE_DECORATOR void *getFifoBuffer(int contextId) const {
200+
return _commBase.getFifoBuffer(contextId);
199201
}
200202
};
201203

flagcx/adaptor/include/device_api/nvidia_comm_traits.h

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -129,7 +129,8 @@ struct CommTraits<NvidiaVendor> {
129129
}
130130
FLAGCX_DEVICE_INLINE_DECORATOR int getRank() const { return _impl.rank; }
131131
FLAGCX_DEVICE_INLINE_DECORATOR int getSize() const { return _impl.nRanks; }
132-
FLAGCX_DEVICE_INLINE_DECORATOR void *getFifoBuffer() const {
132+
FLAGCX_DEVICE_INLINE_DECORATOR void *
133+
getFifoBuffer(int /*contextId*/) const {
133134
return nullptr;
134135
}
135136

flagcx/core/include/comm.h

Lines changed: 13 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -352,14 +352,24 @@ struct flagcxHeteroComm {
352352
struct flagcxRegCache regCache;
353353
uint64_t groupHash;
354354
uint64_t endMagic;
355-
// Kernel FIFO buffer for device side communication
356-
void *fifoBuffer;
355+
// Kernel FIFO buffers for device side communication (one per context)
356+
void *fifoBuffers[FLAGCX_DEVICE_CTA_COUNT];
357357
// uniRunner FIFO buffer
358358
void *uniRunnerFifoBuffer;
359359
// Device communicator (set by flagcxDevCommCreate).
360360
// Used by proxy for BarrierSignal, WaitSignal, PutValue handlers.
361361
flagcxDevComm_t devCommHandle;
362-
362+
// Inter-node signal relay — established once, shared across devComms.
363+
bool relayInitialized;
364+
bool isInterLeader;
365+
int nInterPeers;
366+
int *interPeerRanks;
367+
uint64_t *interSignalFlags;
368+
uint64_t *interSignalFlagsHost;
369+
void **signalSendComms;
370+
void **barrierRecvComms;
371+
void *barrierHandleInfo;
372+
void *netAdaptorPtr;
363373
// Async RMA proxy state (one-sided Put/Get offload thread).
364374
struct flagcxRmaProxyState *rmaProxy;
365375
};

flagcx/core/include/proxy.h

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -28,8 +28,9 @@ enum flagcxProxyOpState {
2828
};
2929

3030
struct flagcxProxyKernelState {
31-
pthread_t thread;
32-
flagcxFifo_t fifo;
31+
pthread_t threads[FLAGCX_DEVICE_CTA_COUNT];
32+
flagcxFifo_t fifos[FLAGCX_DEVICE_CTA_COUNT];
33+
int contextCount = 1;
3334
flagcxStream_t stream;
3435
int stop = 0;
3536
// Synchronization for initialization

flagcx/core/proxy.cc

Lines changed: 71 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -318,6 +318,7 @@ flagcxProxyGetPostedOps(struct flagcxProxyState *proxyState, int *added) {
318318
}
319319

320320
FLAGCX_PARAM(ProgressAppendOpFreq, "PROGRESS_APPENDOP_FREQ", 8);
321+
FLAGCX_PARAM(KernelProxyParallelism, "KERNEL_PROXY_PARALLELISM", 4);
321322

322323
inline void *flagcxProxyProgress(void *proxyState_) {
323324
struct flagcxProxyState *proxyState = (flagcxProxyState *)proxyState_;
@@ -856,6 +857,11 @@ flagcxResult_t flagcxProxyCallBlocking(struct flagcxHeteroComm *comm,
856857
goto exit;
857858
}
858859

860+
struct flagcxProxyKernelServiceArg {
861+
struct flagcxHeteroComm *comm;
862+
int contextId;
863+
};
864+
859865
flagcxResult_t flagcxProxyInit(struct flagcxHeteroComm *comm) {
860866
INFO(FLAGCX_INIT, "rank=%d flagcxProxyInit called.", comm->rank);
861867
FLAGCXCHECK(flagcxSocketInit(&comm->proxyState->listenSock,
@@ -877,21 +883,45 @@ flagcxResult_t flagcxProxyInit(struct flagcxHeteroComm *comm) {
877883
pthread_create(&comm->proxyState->progressState.thread, NULL,
878884
flagcxProxyProgress, comm->proxyState);
879885
#ifdef COMPILE_KERNEL_HOST
880-
// Initialize synchronization primitives before creating thread
886+
// Initialize synchronization primitives before creating threads
881887
pthread_mutex_init(&comm->proxyState->kernelState.initMutex, NULL);
882888
pthread_cond_init(&comm->proxyState->kernelState.initCond, NULL);
883889
comm->proxyState->kernelState.ready = 0;
884890

885-
pthread_create(&comm->proxyState->kernelState.thread, NULL,
886-
flagcxProxyKernelService, (void *)comm);
891+
int nKernelProxies = flagcxParamKernelProxyParallelism();
892+
if (nKernelProxies < 1)
893+
nKernelProxies = 1;
894+
if (nKernelProxies > FLAGCX_DEVICE_CTA_COUNT)
895+
nKernelProxies = FLAGCX_DEVICE_CTA_COUNT;
896+
comm->proxyState->kernelState.contextCount = nKernelProxies;
897+
898+
int nStarted = 0;
899+
for (int i = 0; i < nKernelProxies; i++) {
900+
flagcxProxyKernelServiceArg *arg = new flagcxProxyKernelServiceArg{comm, i};
901+
if (pthread_create(&comm->proxyState->kernelState.threads[i], NULL,
902+
flagcxProxyKernelService, arg) != 0) {
903+
WARN("flagcxProxyInit: failed to create kernel proxy thread %d", i);
904+
delete arg;
905+
break;
906+
}
907+
nStarted++;
908+
}
909+
// Adjust contextCount to the number of threads actually started so the
910+
// cond-wait below and the stop/join loop use a consistent count.
911+
comm->proxyState->kernelState.contextCount = nStarted;
887912

888-
// Wait for kernel proxy thread to finish initialization
913+
// Wait for all started kernel proxy threads to finish initialization
889914
pthread_mutex_lock(&comm->proxyState->kernelState.initMutex);
890-
while (comm->proxyState->kernelState.ready == 0) {
915+
while (comm->proxyState->kernelState.ready < nStarted) {
891916
pthread_cond_wait(&comm->proxyState->kernelState.initCond,
892917
&comm->proxyState->kernelState.initMutex);
893918
}
894919
pthread_mutex_unlock(&comm->proxyState->kernelState.initMutex);
920+
921+
if (nStarted == 0) {
922+
WARN("flagcxProxyInit: no kernel proxy threads started");
923+
return flagcxSystemError;
924+
}
895925
#endif
896926

897927
comm->proxyState->initialized = 1;
@@ -1003,8 +1033,10 @@ void *flagcxProxyService(void *args) {
10031033
pthread_mutex_unlock(&comm->proxyState->mutex);
10041034
pthread_join(comm->proxyState->progressState.thread, nullptr);
10051035
#ifdef COMPILE_KERNEL_HOST
1006-
// Stop kernel thread and cleanup its mutex/cond
1007-
pthread_join(comm->proxyState->kernelState.thread, nullptr);
1036+
// Stop all kernel threads and cleanup
1037+
for (int i = 0; i < comm->proxyState->kernelState.contextCount; i++) {
1038+
pthread_join(comm->proxyState->kernelState.threads[i], nullptr);
1039+
}
10081040
pthread_mutex_destroy(&comm->proxyState->kernelState.initMutex);
10091041
pthread_cond_destroy(&comm->proxyState->kernelState.initCond);
10101042
#endif
@@ -1056,7 +1088,11 @@ void *flagcxProxyKernelService(void *args) {
10561088
int termCount = 0;
10571089
flagcxDeviceTrigger_t ptr = NULL;
10581090
flagcxFifo_t fifo = NULL;
1059-
struct flagcxHeteroComm *comm = (struct flagcxHeteroComm *)args;
1091+
flagcxStream_t stream = NULL;
1092+
flagcxProxyKernelServiceArg *arg = (flagcxProxyKernelServiceArg *)args;
1093+
struct flagcxHeteroComm *comm = arg->comm;
1094+
int contextId = arg->contextId;
1095+
delete arg;
10601096
flagcxResult_t res = flagcxSuccess;
10611097

10621098
auto validateOneSidedPeer = [](struct flagcxHeteroComm *comm,
@@ -1079,19 +1115,19 @@ void *flagcxProxyKernelService(void *args) {
10791115
// Set device context
10801116
FLAGCXCHECKGOTO(deviceAdaptor->setDevice(comm->cudaDev), res, out);
10811117

1082-
// Create FIFO
1083-
comm->proxyState->kernelState.fifo = new flagcxFifo();
1084-
FLAGCXCHECKGOTO(comm->proxyState->kernelState.fifo->flagcxFifoInit(), res,
1085-
out);
1086-
fifo = comm->proxyState->kernelState.fifo;
1087-
// comm->fifoBuffer = (void *)comm->proxyState->kernelState.fifo->buffer;
1088-
FLAGCXCHECKGOTO(deviceAdaptor->hostGetDevicePointer(
1089-
&comm->fifoBuffer,
1090-
(void *)comm->proxyState->kernelState.fifo->buffer),
1091-
res, out);
1118+
// Create FIFO for this thread
1119+
comm->proxyState->kernelState.fifos[contextId] = new flagcxFifo();
1120+
FLAGCXCHECKGOTO(
1121+
comm->proxyState->kernelState.fifos[contextId]->flagcxFifoInit(), res,
1122+
out);
1123+
fifo = comm->proxyState->kernelState.fifos[contextId];
1124+
FLAGCXCHECKGOTO(
1125+
deviceAdaptor->hostGetDevicePointer(
1126+
&comm->fifoBuffers[contextId],
1127+
(void *)comm->proxyState->kernelState.fifos[contextId]->buffer),
1128+
res, out);
10921129

10931130
// Create a dedicated stream
1094-
flagcxStream_t stream;
10951131
FLAGCXCHECKGOTO(deviceAdaptor->streamCreate(&stream), res, out);
10961132
INFO(FLAGCX_P2P, "rank %d p2p stream %lu", comm->rank, (uintptr_t)stream);
10971133

@@ -1100,8 +1136,8 @@ void *flagcxProxyKernelService(void *args) {
11001136

11011137
// Signal that initialization is complete
11021138
pthread_mutex_lock(&comm->proxyState->kernelState.initMutex);
1103-
comm->proxyState->kernelState.ready = 1;
1104-
pthread_cond_signal(&comm->proxyState->kernelState.initCond);
1139+
comm->proxyState->kernelState.ready++;
1140+
pthread_cond_broadcast(&comm->proxyState->kernelState.initCond);
11051141
pthread_mutex_unlock(&comm->proxyState->kernelState.initMutex);
11061142

11071143
while (true) {
@@ -1341,17 +1377,21 @@ void *flagcxProxyKernelService(void *args) {
13411377
if (res != flagcxSuccess)
13421378
break;
13431379
}
1344-
// destroy stream
1345-
res = deviceAdaptor->streamSynchronize(stream);
1346-
res = deviceAdaptor->streamDestroy(stream);
1347-
// deallocate trigger structure
1348-
free(ptr);
1349-
13501380
out:
1351-
// destroy fifo
1352-
res = comm->proxyState->kernelState.fifo->flagcxFifoDestroy();
1353-
delete comm->proxyState->kernelState.fifo;
1354-
comm->fifoBuffer = NULL;
1381+
// destroy stream (only if created)
1382+
if (stream != nullptr) {
1383+
deviceAdaptor->streamSynchronize(stream);
1384+
deviceAdaptor->streamDestroy(stream);
1385+
}
1386+
// deallocate trigger structure (only if allocated)
1387+
free(ptr);
1388+
// destroy fifo (only if created)
1389+
if (comm->proxyState->kernelState.fifos[contextId] != nullptr) {
1390+
comm->proxyState->kernelState.fifos[contextId]->flagcxFifoDestroy();
1391+
delete comm->proxyState->kernelState.fifos[contextId];
1392+
comm->proxyState->kernelState.fifos[contextId] = nullptr;
1393+
}
1394+
comm->fifoBuffers[contextId] = NULL;
13551395
return NULL;
13561396
}
13571397

flagcx/flagcx.cc

Lines changed: 30 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1593,7 +1593,12 @@ flagcxResult_t flagcxCommDestroy(flagcxComm_t comm) {
15931593
}
15941594

15951595
if (!useHomoComm(comm)) {
1596-
// Destroy hetero comm
1596+
// Tear down inter-node signal relay first: drains FIFOs and closes RDMA
1597+
// connections. Must run before flagcxHeteroCommDestroy, which frees
1598+
// proxyState and heteroComm. Proxy threads are stopped inside
1599+
// flagcxCommRelayDestroy via the bootstrap barrier before any teardown.
1600+
FLAGCXCHECK(flagcxCommRelayDestroy(comm));
1601+
// Destroy hetero comm (stops/joins proxy threads, frees proxyState)
15971602
FLAGCXCHECK(flagcxHeteroCommDestroy(comm->heteroComm));
15981603
// Destroy host comm
15991604
if (useHostComm()) {
@@ -1679,11 +1684,32 @@ flagcxResult_t flagcxCommUserRank(const flagcxComm_t comm, int *rank) {
16791684
return flagcxHeteroCommUserRank(comm->heteroComm, rank);
16801685
}
16811686

1682-
flagcxResult_t flagcxCommFifoBuffer(const flagcxComm_t comm, void **buffer) {
1683-
if (comm->heteroComm->fifoBuffer == NULL) {
1687+
flagcxResult_t flagcxCommFifoBuffer(const flagcxComm_t comm, int contextId,
1688+
void **buffer) {
1689+
FLAGCXCHECK(flagcxEnsureCommReady(comm));
1690+
1691+
if (buffer == nullptr) {
1692+
return flagcxInvalidArgument;
1693+
}
1694+
1695+
if (contextId < 0 || contextId >= FLAGCX_DEVICE_CTA_COUNT) {
1696+
return flagcxInvalidArgument;
1697+
}
1698+
1699+
// FIFO buffers are only available on hetero communicators
1700+
if (useHomoComm(comm) && !useHeteroComm()) {
1701+
return flagcxNotSupported;
1702+
}
1703+
1704+
if (comm->heteroComm == nullptr) {
1705+
return flagcxNotSupported;
1706+
}
1707+
1708+
if (comm->heteroComm->fifoBuffers[contextId] == nullptr) {
16841709
return flagcxInvalidUsage;
16851710
}
1686-
*buffer = comm->heteroComm->fifoBuffer;
1711+
1712+
*buffer = comm->heteroComm->fifoBuffers[contextId];
16871713
return flagcxSuccess;
16881714
}
16891715

flagcx/include/flagcx.h

Lines changed: 0 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -262,10 +262,6 @@ flagcxResult_t flagcxCommGetDeviceNumber(const flagcxComm_t comm, int *device);
262262
/* Returns the user-ordered "rank" associated with the communicator. */
263263
flagcxResult_t flagcxCommUserRank(const flagcxComm_t comm, int *rank);
264264

265-
/* Returns `(void *)fifoBuffer` associated with the `heteroComm` of the input
266-
* communicator */
267-
flagcxResult_t flagcxCommFifoBuffer(const flagcxComm_t comm, void **buffer);
268-
269265
/*
270266
* Collective communication operations
271267
*

flagcx/include/flagcx_kernel.h

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -420,6 +420,12 @@ flagcxResult_t flagcxDevMemDestroy(flagcxComm_t comm, flagcxDevMem_t devMem);
420420
// so that cudaFree does not deadlock on device synchronization.
421421
flagcxResult_t flagcxCommCleanupIpcTable(flagcxComm_t comm);
422422

423+
// Tear down inter-node signal relay stored on heteroComm.
424+
// Must be called before flagcxHeteroCommDestroy (which frees proxyState and
425+
// heteroComm). Internally drains FIFOs and performs a cross-rank barrier
426+
// before closing RDMA connections.
427+
flagcxResult_t flagcxCommRelayDestroy(flagcxComm_t comm);
428+
423429
// Deferred device/host-pinned memory free.
424430
// Collects pointers during DevComm/DevMem cleanup.
425431
void flagcxCommDeferFree(flagcxComm_t comm, void *ptr, int memType);

0 commit comments

Comments
 (0)