Skip to content

Commit 1693a0b

Browse files
authored
[CRL] Change Net/P2P buffer size and chunk size to environment variable (flagos-ai#316)
1 parent 08051f6 commit 1693a0b

8 files changed

Lines changed: 65 additions & 46 deletions

File tree

flagcx/core/group.cc

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -246,9 +246,9 @@ static flagcxResult_t groupLaunch(struct flagcxAsyncJob *job_) {
246246
.peers[peer]
247247
->recv[0]
248248
.proxyConn.connection;
249-
op->args.chunkSize = CHUNKSIZE;
250-
op->args.chunkSteps = (p2p->bytes + CHUNKSIZE - 1) / (CHUNKSIZE);
251-
op->args.sendStepMask = MAXSTEPS - 1;
249+
op->args.chunkSize = flagcxNetChunkSize;
250+
op->args.chunkSteps = (p2p->bytes + flagcxNetChunkSize - 1) / (flagcxNetChunkSize);
251+
op->args.sendStepMask = FLAGCX_NET_MAX_STEPS - 1;
252252
op->args.deviceFuncRelaxedOrdering = deviceFuncRelaxedOrdering;
253253
op->stream = p2p->stream;
254254
if (op->connection->transport == TRANSPORT_P2P) {
@@ -345,9 +345,9 @@ static flagcxResult_t groupLaunch(struct flagcxAsyncJob *job_) {
345345
.peers[peer]
346346
->send[0]
347347
.proxyConn.connection;
348-
op->args.chunkSize = CHUNKSIZE;
349-
op->args.chunkSteps = (p2p->bytes + CHUNKSIZE - 1) / (CHUNKSIZE);
350-
op->args.sendStepMask = MAXSTEPS - 1;
348+
op->args.chunkSize = flagcxNetChunkSize;
349+
op->args.chunkSteps = (p2p->bytes + flagcxNetChunkSize - 1) / (flagcxNetChunkSize);
350+
op->args.sendStepMask = FLAGCX_NET_MAX_STEPS - 1;
351351
op->args.deviceFuncRelaxedOrdering = deviceFuncRelaxedOrdering;
352352
op->stream = p2p->stream;
353353
if (op->connection->transport == TRANSPORT_P2P) {

flagcx/core/include/net.h

Lines changed: 5 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -16,11 +16,10 @@
1616

1717
typedef char flagcxNetHandle_t[FLAGCX_NET_HANDLE_MAXSIZE];
1818

19-
#define REGMRBUFFERSIZE (64ULL * 1024 * 1024)
20-
#define CHUNKSIZE (4ULL * 1024 * 1024)
19+
extern int64_t flagcxNetBufferSize;
20+
extern int64_t flagcxNetChunkSize;
2121
#define FLAGCX_MAX_NET_SIZE_BYTES (1 * 1024 * 1024 * 1024 * 1024L)
22-
#define MAXSTEPS (REGMRBUFFERSIZE / CHUNKSIZE)
23-
static_assert((MAXSTEPS & (MAXSTEPS - 1)) == 0, "send step must a power of 2");
22+
#define FLAGCX_NET_MAX_STEPS 16
2423

2524
flagcxResult_t flagcxNetInit(struct flagcxHeteroComm *comm);
2625
int flagcxNetVersion(struct flagcxHeteroComm *comm);
@@ -61,7 +60,7 @@ struct sendNetResources {
6160
flagcxNetDeviceType netDeviceType;
6261
flagcxNetDeviceHandle_t *netDeviceHandle;
6362
flagcxStream_t cpStream;
64-
flagcxEvent_t cpEvents[MAXSTEPS];
63+
flagcxEvent_t cpEvents[FLAGCX_NET_MAX_STEPS];
6564
};
6665

6766
struct recvNetResources {
@@ -96,7 +95,7 @@ struct recvNetResources {
9695
flagcxNetDeviceType netDeviceType;
9796
flagcxNetDeviceHandle_t *netDeviceHandle;
9897
flagcxStream_t cpStream;
99-
flagcxEvent_t cpEvents[MAXSTEPS];
98+
flagcxEvent_t cpEvents[FLAGCX_NET_MAX_STEPS];
10099
};
101100

102101
enum flagcxIbCommState {

flagcx/core/include/p2p.h

Lines changed: 8 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -8,13 +8,13 @@
88
#include "shmutils.h"
99
#include "transport.h"
1010
#include <stddef.h>
11-
#define FLAGCX_P2P_BUFFERSIZE \
12-
(64ULL * 1024 * 1024) // 64MB buffer for P2P transfers
13-
#define FLAGCX_P2P_CHUNKSIZE (4ULL * 1024 * 1024) // 4MB chunk size
14-
#define FLAGCX_P2P_STEPS \
15-
(FLAGCX_P2P_BUFFERSIZE / FLAGCX_P2P_CHUNKSIZE) // 16 steps
16-
#define FLAGCX_P2P_MAX_OPS \
17-
32 // Maximum number of concurrent P2P operation pairs
11+
12+
extern int64_t flagcxP2PBufferSize;
13+
extern int64_t flagcxP2PChunkSize;
14+
15+
#define FLAGCX_P2P_MAX_STEPS 16
16+
#define FLAGCX_P2P_MAX_OPS \
17+
(FLAGCX_P2P_MAX_STEPS * 2) // Maximum number of concurrent P2P operation pairs
1818
#define FLAGCX_P2P_IPC_HANDLE_SIZE 64
1919

2020
#ifdef __cplusplus
@@ -79,7 +79,7 @@ struct flagcxP2pShmProxyInfo {
7979
// Device side
8080
char *recvFifo;
8181
flagcxStream_t stream;
82-
flagcxEvent_t events[FLAGCX_P2P_STEPS];
82+
flagcxEvent_t events[FLAGCX_P2P_MAX_STEPS];
8383
};
8484

8585
struct flagcxP2pResources {

flagcx/core/include/proxy.h

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -86,7 +86,7 @@ struct flagcxProxySubArgs {
8686
};
8787

8888
struct flagcxProxyArgs {
89-
struct flagcxProxySubArgs subs[MAXSTEPS];
89+
struct flagcxProxySubArgs subs[FLAGCX_NET_MAX_STEPS];
9090
proxyProgressFunc_t progress;
9191
int nsubs;
9292
int done;

flagcx/core/init.cc

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111
#include "flagcx.h"
1212
#include "group.h"
1313
#include "net.h"
14+
#include "p2p.h"
1415
#include "topo.h"
1516
#include "transport.h"
1617
#include "type.h"
@@ -252,6 +253,11 @@ static flagcxResult_t initTransportsRank(flagcxHeteroComm_t comm,
252253
return flagcxInternalError;
253254
}
254255

256+
FLAGCX_PARAM(P2PBufferSize, "P2P_BUFFER_SIZE", 64L * 1024 * 1024); // default value to 64MB
257+
FLAGCX_PARAM(P2PChunkSize, "P2P_CHUNK_SIZE", 4L * 1024 * 1024); // default value to 4MB
258+
FLAGCX_PARAM(NetBufferSize, "NET_BUFFER_SIZE", 64L * 1024 * 1024); // default value to 64MB
259+
FLAGCX_PARAM(NetChunkSize, "NET_CHUNK_SIZE", 4L * 1024 * 1024); // default value to 4MB
260+
255261
static flagcxResult_t flagcxCommInitRankFunc(struct flagcxAsyncJob *job_) {
256262
struct flagcxCommInitRankAsyncJob *job =
257263
(struct flagcxCommInitRankAsyncJob *)job_;
@@ -327,6 +333,13 @@ static flagcxResult_t flagcxCommInitRankFunc(struct flagcxAsyncJob *job_) {
327333
}
328334
}
329335

336+
flagcxNetBufferSize = flagcxParamNetBufferSize();
337+
flagcxNetChunkSize = flagcxParamNetChunkSize();
338+
flagcxP2PBufferSize = flagcxParamP2PBufferSize();
339+
flagcxP2PChunkSize = flagcxParamP2PChunkSize();
340+
assert((flagcxNetBufferSize + flagcxNetChunkSize - 1) / flagcxNetChunkSize <= FLAGCX_NET_MAX_STEPS);
341+
assert((flagcxP2PBufferSize + flagcxP2PChunkSize - 1) / flagcxP2PChunkSize <= FLAGCX_P2P_MAX_STEPS);
342+
330343
FLAGCXCHECK(flagcxNetInit(comm));
331344
INFO(FLAGCX_INIT, "Using network %s", comm->netAdaptor->name);
332345
if (env && (strcmp(env, "TRUE") == 0 || strcmp(env, "True") == 0)) {

flagcx/core/net.cc

Lines changed: 12 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,10 @@
66

77
#include <errno.h>
88
#include <string.h>
9+
#include <string>
10+
11+
int64_t flagcxNetBufferSize;
12+
int64_t flagcxNetChunkSize;
913

1014
static pthread_mutex_t netLock = PTHREAD_MUTEX_INITIALIZER;
1115
// Use adaptor system for all network types
@@ -143,12 +147,12 @@ flagcxResult_t flagcxProxySend(sendNetResources *resources, void *data,
143147
int stepMask = args->sendStepMask;
144148

145149
if (args->waitCopy < args->chunkSteps &&
146-
args->waitCopy - args->transmitted < MAXSTEPS) {
150+
args->waitCopy - args->transmitted < FLAGCX_NET_MAX_STEPS) {
147151
int step = args->waitCopy & stepMask;
148152
args->subs[step].stepSize =
149153
std::min(args->chunkSize, size - args->totalCopySize);
150154
if (!args->regBufFlag) {
151-
args->subs[step].stepBuff = resources->buffers[0] + (CHUNKSIZE * step);
155+
args->subs[step].stepBuff = resources->buffers[0] + (flagcxNetChunkSize * step);
152156
if (resources->netAdaptor == getUnifiedNetAdaptor(IBRC)) {
153157
FLAGCXCHECK(deviceAdaptor->deviceMemcpy(
154158
args->subs[step].stepBuff, (char *)data + args->totalCopySize,
@@ -164,7 +168,7 @@ flagcxResult_t flagcxProxySend(sendNetResources *resources, void *data,
164168
resources->cpStream));
165169
} else {
166170
args->subs[step].stepBuff =
167-
(void *)((char *)data + (CHUNKSIZE * args->waitCopy));
171+
(void *)((char *)data + (flagcxNetChunkSize * args->waitCopy));
168172
}
169173
args->totalCopySize += args->subs[step].stepSize;
170174
args->waitCopy++;
@@ -226,17 +230,17 @@ flagcxResult_t flagcxProxyRecv(recvNetResources *resources, void *data,
226230
if (args->copied < args->chunkSteps) {
227231
int stepMask = args->sendStepMask;
228232
if (args->posted < args->chunkSteps &&
229-
args->posted - args->copied < MAXSTEPS) {
233+
args->posted - args->copied < FLAGCX_NET_MAX_STEPS) {
230234
int tags[8] = {0};
231235
void *req = NULL;
232236
args->subs[args->posted & stepMask].stepSize =
233237
std::min(args->chunkSize, size - args->totalPostSize);
234238
if (!args->regBufFlag) {
235239
args->subs[args->posted & stepMask].stepBuff =
236-
resources->buffers[0] + CHUNKSIZE * (args->posted & stepMask);
240+
resources->buffers[0] + flagcxNetChunkSize * (args->posted & stepMask);
237241
} else {
238242
args->subs[args->posted & stepMask].stepBuff =
239-
(void *)((char *)data + CHUNKSIZE * args->posted);
243+
(void *)((char *)data + flagcxNetChunkSize * args->posted);
240244
}
241245
resources->netAdaptor->irecv(
242246
resources->netRecvComm, 1,
@@ -340,7 +344,7 @@ flagcxResult_t flagcxProxyRecv(recvNetResources *resources, void *data,
340344
}
341345

342346
flagcxResult_t flagcxSendProxyFree(sendNetResources *resources) {
343-
for (int s = 0; s < MAXSTEPS; s++) {
347+
for (int s = 0; s < FLAGCX_NET_MAX_STEPS; s++) {
344348
FLAGCXCHECK(deviceAdaptor->eventDestroy(resources->cpEvents[s]));
345349
}
346350
FLAGCXCHECK(deviceAdaptor->streamDestroy(resources->cpStream));
@@ -356,7 +360,7 @@ flagcxResult_t flagcxSendProxyFree(sendNetResources *resources) {
356360
}
357361

358362
flagcxResult_t flagcxRecvProxyFree(recvNetResources *resources) {
359-
for (int s = 0; s < MAXSTEPS; s++) {
363+
for (int s = 0; s < FLAGCX_NET_MAX_STEPS; s++) {
360364
FLAGCXCHECK(deviceAdaptor->eventDestroy(resources->cpEvents[s]));
361365
}
362366
FLAGCXCHECK(deviceAdaptor->streamDestroy(resources->cpStream));

flagcx/core/p2p.cc

Lines changed: 15 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,9 @@
88
#include <map>
99
#include <string.h> // for memcpy
1010

11+
int64_t flagcxP2PBufferSize;
12+
int64_t flagcxP2PChunkSize;
13+
1114
struct p2pIpcExpInfo {
1215
flagcxP2pIpcDesc ipcDesc;
1316
bool legacyIpcCap;
@@ -72,7 +75,7 @@ flagcxResult_t flagcxP2pProxySend(struct flagcxP2pResources *resources,
7275
slotPtr->done = 0;
7376
slotPtr->peerDone = 0;
7477
slotPtr->sendHead = 0;
75-
slotPtr->recvTail = FLAGCX_P2P_STEPS;
78+
slotPtr->recvTail = FLAGCX_P2P_MAX_STEPS;
7679
// Reset reg info for new operation
7780
regInfoPtr->copyStarted = 0;
7881
regInfoPtr->copyDone = 0;
@@ -140,7 +143,7 @@ flagcxResult_t flagcxP2pProxySend(struct flagcxP2pResources *resources,
140143
// Non-zero-copy mode: use FIFO buffer
141144
if (args->transmitted < args->chunkSteps) {
142145
if (args->copied < args->chunkSteps &&
143-
args->copied - args->transmitted < FLAGCX_P2P_STEPS) {
146+
args->copied - args->transmitted < FLAGCX_P2P_MAX_STEPS) {
144147
int step = args->copied & args->sendStepMask;
145148

146149
volatile uint64_t *recvTail = &peerSlotPtr->recvTail;
@@ -149,7 +152,7 @@ flagcxResult_t flagcxP2pProxySend(struct flagcxP2pResources *resources,
149152
args->subs[step].stepSize =
150153
std::min(args->chunkSize, size - args->totalCopySize);
151154
args->subs[step].stepBuff =
152-
resources->proxyInfo.recvFifo + (FLAGCX_P2P_CHUNKSIZE * step);
155+
resources->proxyInfo.recvFifo + (flagcxP2PChunkSize * step);
153156

154157
FLAGCXCHECK(deviceAdaptor->deviceMemcpy(
155158
args->subs[step].stepBuff, (char *)data + args->totalCopySize,
@@ -227,7 +230,7 @@ flagcxResult_t flagcxP2pProxyRecv(struct flagcxP2pResources *resources,
227230
slotPtr->done = 0;
228231
slotPtr->peerDone = 0;
229232
slotPtr->sendHead = 0;
230-
slotPtr->recvTail = FLAGCX_P2P_STEPS;
233+
slotPtr->recvTail = FLAGCX_P2P_MAX_STEPS;
231234
}
232235

233236
// Return and retry later since the slot is still in use
@@ -277,15 +280,15 @@ flagcxResult_t flagcxP2pProxyRecv(struct flagcxP2pResources *resources,
277280
// Non-zero-copy mode: use FIFO buffer
278281
if (args->transmitted < args->chunkSteps) {
279282
if (args->copied < args->chunkSteps &&
280-
args->copied - args->transmitted < FLAGCX_P2P_STEPS) {
283+
args->copied - args->transmitted < FLAGCX_P2P_MAX_STEPS) {
281284
int step = args->copied & args->sendStepMask;
282285
volatile uint64_t *sendHead = &peerSlotPtr->sendHead;
283286

284287
if (*sendHead > args->copied) {
285288
args->subs[step].stepSize =
286289
std::min(args->chunkSize, size - args->totalCopySize);
287290
args->subs[step].stepBuff =
288-
resources->proxyInfo.recvFifo + (FLAGCX_P2P_CHUNKSIZE * step);
291+
resources->proxyInfo.recvFifo + (flagcxP2PChunkSize * step);
289292

290293
FLAGCXCHECK(deviceAdaptor->deviceMemcpy(
291294
(char *)data + args->totalCopySize, args->subs[step].stepBuff,
@@ -308,7 +311,7 @@ flagcxResult_t flagcxP2pProxyRecv(struct flagcxP2pResources *resources,
308311
args->transmitted++;
309312
// Update recvTail in the shared slot
310313
volatile uint64_t *recvTail = &slotPtr->recvTail;
311-
*recvTail = args->transmitted + FLAGCX_P2P_STEPS;
314+
*recvTail = args->transmitted + FLAGCX_P2P_MAX_STEPS;
312315
}
313316
}
314317
} else {
@@ -411,7 +414,7 @@ flagcxResult_t flagcxP2pSendProxySetup(struct flagcxProxyConnection *connection,
411414
// Initialize all synchronization slots
412415
for (int i = 0; i < FLAGCX_P2P_MAX_OPS; i++) {
413416
resources->proxyInfo.shm->slots[i].sendHead = 0;
414-
resources->proxyInfo.shm->slots[i].recvTail = FLAGCX_P2P_STEPS;
417+
resources->proxyInfo.shm->slots[i].recvTail = FLAGCX_P2P_MAX_STEPS;
415418
resources->proxyInfo.shm->slots[i].opHash = -1;
416419
resources->proxyInfo.shm->slots[i].done = 1; // 1 = slot is free
417420
resources->proxyInfo.shm->slots[i].peerDone = 1; // 1 = slot is free
@@ -481,7 +484,7 @@ flagcxP2pSendProxyConnect(struct flagcxProxyConnection *connection,
481484

482485
// Create stream and events for data transfers
483486
FLAGCXCHECK(deviceAdaptor->streamCreate(&resources->proxyInfo.stream));
484-
for (int i = 0; i < FLAGCX_P2P_STEPS; i++) {
487+
for (int i = 0; i < FLAGCX_P2P_MAX_STEPS; i++) {
485488
FLAGCXCHECK(deviceAdaptor->eventCreate(&resources->proxyInfo.events[i],
486489
flagcxEventDisableTiming));
487490
}
@@ -508,7 +511,7 @@ flagcxP2pRecvProxyConnect(struct flagcxProxyConnection *connection,
508511

509512
// Create stream and events for data transfers
510513
FLAGCXCHECK(deviceAdaptor->streamCreate(&resources->proxyInfo.stream));
511-
for (int i = 0; i < FLAGCX_P2P_STEPS; i++) {
514+
for (int i = 0; i < FLAGCX_P2P_MAX_STEPS; i++) {
512515
FLAGCXCHECK(deviceAdaptor->eventCreate(&resources->proxyInfo.events[i],
513516
flagcxEventDisableTiming));
514517
}
@@ -1047,7 +1050,7 @@ flagcxResult_t flagcxP2pSendProxyFree(struct flagcxP2pResources *resources) {
10471050
if (resources == NULL)
10481051
return flagcxSuccess;
10491052

1050-
for (int s = 0; s < FLAGCX_P2P_STEPS; s++) {
1053+
for (int s = 0; s < FLAGCX_P2P_MAX_STEPS; s++) {
10511054
if (resources->proxyInfo.events[s] != NULL) {
10521055
FLAGCXCHECK(deviceAdaptor->eventDestroy(resources->proxyInfo.events[s]));
10531056
}
@@ -1068,7 +1071,7 @@ flagcxResult_t flagcxP2pRecvProxyFree(struct flagcxP2pResources *resources) {
10681071
return flagcxSuccess;
10691072

10701073
// Destroy events
1071-
for (int s = 0; s < FLAGCX_P2P_STEPS; s++) {
1074+
for (int s = 0; s < FLAGCX_P2P_MAX_STEPS; s++) {
10721075
if (resources->proxyInfo.events[s] != NULL) {
10731076
FLAGCXCHECK(deviceAdaptor->eventDestroy(resources->proxyInfo.events[s]));
10741077
}

flagcx/core/transport.cc

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -39,7 +39,7 @@ flagcxResult_t flagcxTransportP2pSetup(struct flagcxHeteroComm *comm,
3939
conn->proxyConn.connection->send = 0;
4040
conn->proxyConn.connection->transportResources = (void *)resources;
4141
if (peer != comm->rank) {
42-
struct flagcxP2pRequest req = {(size_t(FLAGCX_P2P_BUFFERSIZE)), 0};
42+
struct flagcxP2pRequest req = {(size_t(flagcxP2PBufferSize)), 0};
4343
struct flagcxP2pConnectInfo connectInfo = {0};
4444
connectInfo.rank = comm->rank;
4545
connectInfo.read = 0;
@@ -66,11 +66,11 @@ flagcxResult_t flagcxTransportP2pSetup(struct flagcxHeteroComm *comm,
6666
resources->netDev = comm->netDev;
6767
resources->netAdaptor = comm->netAdaptor;
6868
deviceAdaptor->streamCreate(&resources->cpStream);
69-
for (int s = 0; s < MAXSTEPS; s++) {
69+
for (int s = 0; s < FLAGCX_NET_MAX_STEPS; s++) {
7070
deviceAdaptor->eventCreate(&resources->cpEvents[s],
7171
flagcxEventDisableTiming);
7272
}
73-
resources->buffSizes[0] = REGMRBUFFERSIZE;
73+
resources->buffSizes[0] = flagcxNetBufferSize;
7474
if (comm->netAdaptor == getUnifiedNetAdaptor(SOCKET)) {
7575
resources->buffers[0] = (char *)malloc(resources->buffSizes[0]);
7676
if (!resources->buffers[0]) {
@@ -132,11 +132,11 @@ flagcxResult_t flagcxTransportP2pSetup(struct flagcxHeteroComm *comm,
132132
resources->netDev = comm->netDev;
133133
resources->netAdaptor = comm->netAdaptor;
134134
deviceAdaptor->streamCreate(&resources->cpStream);
135-
for (int s = 0; s < MAXSTEPS; s++) {
135+
for (int s = 0; s < FLAGCX_NET_MAX_STEPS; s++) {
136136
deviceAdaptor->eventCreate(&resources->cpEvents[s],
137137
flagcxEventDisableTiming);
138138
}
139-
resources->buffSizes[0] = REGMRBUFFERSIZE;
139+
resources->buffSizes[0] = flagcxNetBufferSize;
140140
if (comm->netAdaptor == getUnifiedNetAdaptor(SOCKET)) {
141141
resources->buffers[0] = (char *)malloc(resources->buffSizes[0]);
142142
if (!resources->buffers[0]) {

0 commit comments

Comments
 (0)