Skip to content

Commit cfa760a

Browse files
authored
[CRL] Refactor P2P zerocopy (flagos-ai#452)
1 parent 9d0e73d commit cfa760a

13 files changed

Lines changed: 768 additions & 515 deletions

File tree

flagcx/core/group.cc

Lines changed: 9 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -283,23 +283,15 @@ static flagcxResult_t groupLaunch(struct flagcxAsyncJob *job_) {
283283
peer, comm->rank, op->args.p2pPeerSlotIdx,
284284
op->args.p2pPeerOpHash);
285285

286-
flagcxConnector *recvConn =
287-
&comm->channels[op->channelId].peers[peer]->recv[0];
288-
flagcxConnector *peerConns[] = {recvConn};
289286
int peerRanks[] = {peer};
290287
uintptr_t regOffset = 0;
291288
uintptr_t *peerRmtAddr = NULL;
292289
op->args.regBufFlag = 0;
293290
FLAGCXCHECK(flagcxP2pRegisterBuffer(
294-
comm, p2p->buff, p2p->bytes, peerConns, peerRanks, 1,
295-
/*isSender=*/false, &op->args.regBufFlag, &regOffset,
296-
&peerRmtAddr, op->args.p2pPeerSlotIdx));
297-
if (op->args.regBufFlag) {
298-
INFO(FLAGCX_REG,
299-
"flagcxGroup P2P recv reg rank %d <- %d buff %p size %zu "
300-
"offset %zu remote %p",
301-
comm->rank, peer, p2p->buff, p2p->bytes, (size_t)regOffset,
302-
peerRmtAddr ? (void *)(*peerRmtAddr) : NULL);
291+
comm, p2p->buff, p2p->bytes, peerRanks, 1,
292+
&op->args.regBufFlag, &regOffset, &peerRmtAddr));
293+
if (op->args.regBufFlag && peerRmtAddr) {
294+
op->args.p2pRmtAddr = (void *)peerRmtAddr;
303295
}
304296
} else if (op->connection->transport == TRANSPORT_NET) {
305297
op->args.chunkSize = flagcxNetChunkSize;
@@ -375,16 +367,15 @@ static flagcxResult_t groupLaunch(struct flagcxAsyncJob *job_) {
375367
"peerOpHash(%ld)]",
376368
peer, comm->rank, op->args.p2pPeerSlotIdx,
377369
op->args.p2pPeerOpHash);
378-
flagcxConnector *peerConns[] = {
379-
comm->channels[op->channelId].peers[peer]->send};
370+
// Send side: register own buffer to peer's proxy for READ mode.
371+
// The actual IPC address comes from SHM at proxy time.
380372
int peerRanks[] = {peer};
381373
uintptr_t regOffset = 0;
382374
uintptr_t *peerRmtAddr = NULL;
375+
op->args.regBufFlag = 0;
383376
FLAGCXCHECK(flagcxP2pRegisterBuffer(
384-
comm, p2p->buff, p2p->bytes, peerConns, peerRanks, 1,
385-
/*isSender=*/true, &op->args.regBufFlag, &regOffset,
386-
&peerRmtAddr, op->args.p2pSlotIdx));
387-
// peerRmtAddr is fully resolved (rmtRegAddr + peer's userOffset)
377+
comm, p2p->buff, p2p->bytes, peerRanks, 1,
378+
&op->args.regBufFlag, &regOffset, &peerRmtAddr));
388379
if (op->args.regBufFlag && peerRmtAddr) {
389380
op->args.p2pRmtAddr = (void *)peerRmtAddr;
390381
}

flagcx/core/include/comm.h

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -293,6 +293,8 @@ struct flagcxHeteroComm {
293293
uint64_t intraBarrierGate; // only used if this is intraComm0
294294

295295
struct flagcxProxyState *proxyState;
296+
struct flagcxProxyConnector
297+
*gproxyConn; // Array[nRanks], per-peer proxy connector for IPC reg
296298
int proxyRefCountOld; /* store proxy post-atomic-sub refcount */
297299
// Whether this communicator uses collNet
298300
int collNetSupport;

flagcx/core/include/device.h

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -120,6 +120,7 @@ struct flagcxProxyConnector {
120120
int tpRank;
121121
int tpLocalRank;
122122
int sameProcess;
123+
bool initialized;
123124
struct flagcxProxyConnection *connection;
124125
flagcxResult_t (*proxyProgress)(
125126
struct flagcxProxyState *proxyState,

flagcx/core/include/launch_kernel.h

Lines changed: 17 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -46,9 +46,11 @@ struct flagcxHostSemaphore : public flagcxSemaphore {
4646
std::unordered_map<int, int> stepInfo; // opId -> sigalId
4747
std::vector<std::pair<int, int>> signals; // [curStep, nSteps]
4848
std::vector<flagcxEvent_t> events;
49+
bool frozen; // true during execution phase
4950

5051
flagcxHostSemaphore() {
5152
counter = 0;
53+
frozen = false;
5254
stepInfo.reserve(FLAGCX_OPS_PER_SEMAPHORE);
5355
signals.reserve(FLAGCX_SIGNALS_PER_SEMAPHORE);
5456
events.reserve(FLAGCX_SIGNALS_PER_SEMAPHORE);
@@ -65,22 +67,27 @@ struct flagcxHostSemaphore : public flagcxSemaphore {
6567
return event;
6668
}
6769
void signalStart() override {
70+
frozen =
71+
true; // freeze: no more structural mutations until wait() completes
6872
for (auto it = stepInfo.begin(); it != stepInfo.end(); ++it) {
6973
__atomic_store_n(&signals[it->second].first, 0, __ATOMIC_RELEASE);
7074
}
7175
}
7276
void *getSignals() override { return nullptr; }
7377
void subCounter(int opId = 0) override {
74-
assert(stepInfo.find(opId) != stepInfo.end());
75-
__atomic_fetch_add(&signals[stepInfo[opId]].first, 1, __ATOMIC_RELEASE);
78+
auto it = stepInfo.find(opId);
79+
assert(it != stepInfo.end());
80+
int idx = it->second;
81+
__atomic_fetch_add(&signals[idx].first, 1, __ATOMIC_RELEASE);
7682
INFO(FLAGCX_PROXY,
7783
"SubCounter curStep[%d] = %d, nSteps[%d] = %d, counter %d", opId,
78-
signals[stepInfo[opId]].first, opId, signals[stepInfo[opId]].second,
79-
counter);
84+
signals[idx].first, opId, signals[idx].second, counter);
8085
}
8186
void addCounter(int opId = 0) override {
82-
if (stepInfo.find(opId) != stepInfo.end()) {
83-
__atomic_fetch_add(&signals[stepInfo[opId]].second, 1, __ATOMIC_RELEASE);
87+
assert(!frozen); // must not mutate during execution phase
88+
auto it = stepInfo.find(opId);
89+
if (it != stepInfo.end()) {
90+
__atomic_fetch_add(&signals[it->second].second, 1, __ATOMIC_RELEASE);
8491
} else {
8592
signals.emplace_back(-1, 1);
8693
stepInfo[opId] = (int)signals.size() - 1;
@@ -89,8 +96,9 @@ struct flagcxHostSemaphore : public flagcxSemaphore {
8996
}
9097
int getCounter() override { return counter; }
9198
int pollStart(int opId = 0, int step = 0) override {
92-
assert(stepInfo.find(opId) != stepInfo.end());
93-
return (signals[stepInfo[opId]].first >= step);
99+
auto it = stepInfo.find(opId);
100+
assert(it != stepInfo.end());
101+
return (signals[it->second].first >= step);
94102
}
95103
int pollEnd() override {
96104
return (__atomic_load_n(&counter, __ATOMIC_ACQUIRE) == 0);
@@ -109,6 +117,7 @@ struct flagcxHostSemaphore : public flagcxSemaphore {
109117
sched_yield();
110118
}
111119
__atomic_store_n(&counter, 0, __ATOMIC_RELEASE);
120+
frozen = false; // unfreeze: allow addCounter for next round
112121
}
113122
};
114123

flagcx/core/include/p2p.h

Lines changed: 11 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -58,9 +58,13 @@ struct flagcxP2pSyncSlot {
5858
struct p2pRegInfo {
5959
int copyDone; // Indicates if the copy operation is complete
6060
int copyStarted; // Indicates if the copy operation has started
61+
// WRITE mode: recv publishes into sender's slot
62+
int ipcRecvRegReady; // 1 = ipcRecvRmtAddr valid; recv sets
63+
uintptr_t ipcRecvRmtAddr; // Recv's buffer mapped in sender's address space
64+
// READ mode: sender publishes into receiver's slot
65+
int ipcSendRegReady; // 1 = ipcSendRmtAddr valid; sender sets
6166
uintptr_t
62-
ipcUserOffset; // Per-slot IPC offset (recv-side writes, send-side reads)
63-
int ipcRegReady; // 1 = ipcUserOffset is valid for current op
67+
ipcSendRmtAddr; // Sender's buffer mapped in receiver's address space
6468
};
6569

6670
struct flagcxP2pShm {
@@ -149,12 +153,11 @@ flagcxResult_t flagcxP2pImportShareableBuffer(struct flagcxHeteroComm *comm,
149153
struct flagcxP2pIpcDesc *ipcDesc,
150154
void **devMemPtr);
151155

152-
flagcxResult_t
153-
flagcxP2pRegisterBuffer(struct flagcxHeteroComm *comm, const void *userbuff,
154-
size_t buffSize, struct flagcxConnector **peerConns,
155-
int *peerRanks, int nPeers, bool isSender,
156-
int *regBufFlag, uintptr_t *offsetOut,
157-
uintptr_t **peerRmtAddrsOut, size_t shmRegSlotIdx);
156+
flagcxResult_t flagcxP2pRegisterBuffer(struct flagcxHeteroComm *comm,
157+
const void *userbuff, size_t buffSize,
158+
int *peerRanks, int nPeers,
159+
int *regBufFlag, uintptr_t *offsetOut,
160+
uintptr_t **peerRmtAddrsOut);
158161

159162
flagcxResult_t flagcxP2pDeregisterBuffer(struct flagcxHeteroComm *comm,
160163
struct flagcxIpcRegInfo *info);

flagcx/core/include/proxy.h

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -144,7 +144,8 @@ struct flagcxProxyArgs {
144144
uint64_t p2pPeerOpHash = -1;
145145
size_t p2pSlotIdx = 0;
146146
size_t p2pPeerSlotIdx = 0;
147-
void *p2pRmtAddr = nullptr; // Remote address for P2P zero-copy
147+
void *p2pRmtAddr = nullptr; // remote addr for zero-copy P2P (send side reads,
148+
// recv side writes)
148149

149150
union flagcxProxyOpSpecifics specifics;
150151
};
@@ -336,7 +337,8 @@ struct flagcxProxyState {
336337
pthread_mutex_t mutex;
337338
pthread_cond_t cond;
338339
union flagcxSocketAddress *peerAddresses;
339-
struct flagcxSocket peerSock;
340+
struct flagcxSocket *peerSocks; // Array[nRanks], indexed by rank
341+
int nPeerSocks; // Number of allocated peerSocks entries
340342
struct flagcxProxyOps proxyOps[MAXCHANNELS];
341343

342344
struct flagcxProxyOps *prodProgChannelHead; /*producer*/
@@ -374,6 +376,7 @@ enum proxyConnectState {
374376
struct flagcxProxyConnection {
375377
int send, transport, shared;
376378
int tpLocalRank, sameProcess;
379+
int cudaDev;
377380
struct flagcxSocket *sock;
378381
struct flagcxTransportComm *tcomm;
379382
struct flagcxProxyArgs *proxyAppend;

flagcx/core/include/reg_pool.h

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,7 @@ class flagcxRegPool {
2222
flagcxResult_t addP2pHandle(void *comm, flagcxRegItem *reg, void *handle,
2323
struct flagcxProxyConnector *proxyConn);
2424
flagcxResult_t removeRegItemP2pHandles(void *comm, flagcxRegItem *reg);
25+
flagcxResult_t removeAllP2pHandles(void *comm);
2526
flagcxResult_t registerBuffer(void *comm, void *data, size_t length);
2627
flagcxResult_t deregisterBuffer(void *comm, void *handle);
2728
std::map<uintptr_t, std::map<uintptr_t, flagcxRegItem *>> &getGlobalMap();

flagcx/core/include/register.h

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -53,6 +53,7 @@ struct flagcxIpcRegInfo {
5353
struct flagcxProxyConnector *ipcProxyconn;
5454
struct flagcxIpcImpInfo impInfo;
5555
bool handleReady;
56+
bool sameProcess; // cached at registration time for safe deregister
5657
};
5758

5859
struct flagcxRegItem {

flagcx/core/init.cc

Lines changed: 55 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@
1212
#include "group.h"
1313
#include "net.h"
1414
#include "p2p.h"
15+
#include "reg_pool.h"
1516
#include "topo.h"
1617
#include "transport.h"
1718
#include "type.h"
@@ -294,8 +295,19 @@ static flagcxResult_t flagcxCommInitRankFunc(struct flagcxAsyncJob *job_) {
294295
int nranks = comm->nRanks;
295296
for (int i = 0; i < MAXCHANNELS; i++) {
296297
FLAGCXCHECK(flagcxCalloc(&comm->channels[i].peers, nranks));
297-
for (int r = 0; r < nranks; r++)
298+
for (int r = 0; r < nranks; r++) {
298299
FLAGCXCHECK(flagcxCalloc(&comm->channels[i].peers[r], nranks));
300+
}
301+
}
302+
// Set tpRank = comm->rank for all channel connectors so local RPCs
303+
// route through peerSocks[myRank] to the local service thread
304+
for (int i = 0; i < MAXCHANNELS; i++) {
305+
for (int r = 0; r < nranks; r++) {
306+
for (int c = 0; c < FLAGCX_MAX_CONNS; c++) {
307+
comm->channels[i].peers[r]->send[c].proxyConn.tpRank = comm->rank;
308+
comm->channels[i].peers[r]->recv[c].proxyConn.tpRank = comm->rank;
309+
}
310+
}
299311
}
300312
FLAGCXCHECK(flagcxCalloc(&comm->connectSend, nranks));
301313
FLAGCXCHECK(flagcxCalloc(&comm->connectRecv, nranks));
@@ -339,6 +351,33 @@ static flagcxResult_t flagcxCommInitRankFunc(struct flagcxAsyncJob *job_) {
339351
INFO(FLAGCX_INIT, "Flagcx RuntimeProxy flag set to %d", runtimeProxy);
340352
if (!runtimeProxy) {
341353
FLAGCXCHECK(flagcxProxyInit(comm));
354+
355+
// Allocate gproxyConn array and populate peerAddresses for peer proxy
356+
// connections
357+
FLAGCXCHECK(flagcxCalloc(&comm->gproxyConn, comm->nRanks));
358+
FLAGCXCHECK(flagcxCalloc(&comm->proxyState->peerAddresses, comm->nRanks));
359+
comm->proxyState->peerAddresses[comm->rank] =
360+
comm->proxyState->listenSock.addr;
361+
FLAGCXCHECK(bootstrapAllGather(comm->bootstrap,
362+
comm->proxyState->peerAddresses,
363+
sizeof(union flagcxSocketAddress)));
364+
365+
// Pre-connect all peer sockets (including self)
366+
FLAGCXCHECK(flagcxCalloc(&comm->proxyState->peerSocks, comm->nRanks));
367+
comm->proxyState->nPeerSocks = comm->nRanks;
368+
for (int i = 0; i < comm->nRanks; i++) {
369+
FLAGCXCHECK(flagcxSocketSetFd(-1, &comm->proxyState->peerSocks[i]));
370+
}
371+
for (int i = 0; i < comm->nRanks; i++) {
372+
struct flagcxSocket *sock = &comm->proxyState->peerSocks[i];
373+
FLAGCXCHECK(flagcxSocketInit(sock, comm->proxyState->peerAddresses + i,
374+
comm->magic, flagcxSocketTypeProxy));
375+
FLAGCXCHECK(flagcxSocketConnect(sock));
376+
int ready = 0;
377+
while (!ready) {
378+
FLAGCXCHECK(flagcxSocketReady(sock, &ready));
379+
}
380+
}
342381
}
343382
}
344383

@@ -461,6 +500,8 @@ flagcxResult_t flagcxHeteroCommUserRank(const flagcxHeteroComm_t comm,
461500

462501
flagcxResult_t flagcxHeteroCommDestroy(flagcxHeteroComm_t comm) {
463502
FLAGCXCHECK(flagcxHeteroRmaProxyStop(comm));
503+
// Clean up P2P IPC handles while proxy is still alive and peerSocks valid
504+
FLAGCXCHECK(globalRegPool.removeAllP2pHandles(comm));
464505
flagcxProxyDestroy(comm);
465506
for (int i = 0; i < MAXCHANNELS; i++) {
466507
for (int r = 0; r < comm->nRanks; r++) {
@@ -477,6 +518,19 @@ flagcxResult_t flagcxHeteroCommDestroy(flagcxHeteroComm_t comm) {
477518

478519
free(comm->connectSend);
479520
free(comm->connectRecv);
521+
if (comm->gproxyConn) {
522+
// gproxyConn[i].connection is an opaque handle pointing to a
523+
// flagcxProxyConnection allocated and owned by the peer's service thread.
524+
// Do NOT free it here — the peer frees it when its service thread exits.
525+
free(comm->gproxyConn);
526+
}
527+
free(comm->proxyState->peerAddresses);
528+
if (comm->proxyState->peerSocks != NULL) {
529+
for (int i = 0; i < comm->proxyState->nPeerSocks; i++) {
530+
flagcxSocketClose(&comm->proxyState->peerSocks[i]);
531+
}
532+
free(comm->proxyState->peerSocks);
533+
}
480534
free(comm->proxyState);
481535
free(comm->tasks.peers);
482536
free(comm->tasks.p2pOrder);

0 commit comments

Comments
 (0)