Skip to content

Commit 916d33e

Browse files
committed
Fix S20 multi-context deadlock: enable all blocks to participate in signaling
1 parent 49ba9b0 commit 916d33e

4 files changed

Lines changed: 167 additions & 86 deletions

File tree

bindings/ir/flagcx_device_scalar_ir_impl.h

Lines changed: 0 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -266,9 +266,6 @@ flagcxIntraBarrierSyncS(const void *commOpaque, flagcxCoopKind_t coopKind,
266266
flagcxDevBarrier<flagcxTeamTagIntra, flagcxCoopAny> bar(coop, *comm, team,
267267
index, multimem);
268268
bar.sync(order);
269-
if (comm->_gridBarrierState) {
270-
flagcxGridSync(comm->_gridBarrierState);
271-
}
272269
}
273270

274271
/* ================================================================
@@ -312,9 +309,6 @@ flagcxInterBarrierSyncS(const void *netOpaque, flagcxCoopKind_t coopKind,
312309
flagcxDevBarrier<flagcxTeamTagInter, flagcxCoopAny> bar(coop, *net, team,
313310
index);
314311
bar.sync(order, fence);
315-
if (net->_gridBarrierState) {
316-
flagcxGridSync(net->_gridBarrierState);
317-
}
318312
}
319313

320314
/* ================================================================
@@ -358,9 +352,6 @@ flagcxWorldBarrierSyncS(const void *netOpaque, flagcxCoopKind_t coopKind,
358352
flagcxDevBarrier<flagcxTeamTagWorld, flagcxCoopAny> bar(
359353
coop, flagcxTeamTagWorld{}, *net, index, multimem);
360354
bar.sync(order, fence);
361-
if (net->_gridBarrierState) {
362-
flagcxGridSync(net->_gridBarrierState);
363-
}
364355
}
365356

366357
/* ================================================================

flagcx/adaptor/device_api/default_dev_api_backend.cc

Lines changed: 24 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -897,10 +897,23 @@ static flagcxResult_t defaultDevApiCommGetDevicePtr(flagcxDevComm_t devComm,
897897
flagcxDevComm hostCopy(*devComm);
898898
hostCopy._netContexts = nullptr;
899899

900-
// Step 1: Copy flagcxDevComm to device
900+
// Step 1: Allocate grid sync state (2 x unsigned int, zero-initialized)
901901
void *dPtr = nullptr;
902902
void *netDevPtr = nullptr;
903+
void *gridSyncPtr = nullptr;
903904
flagcxResult_t res = flagcxSuccess;
905+
{
906+
size_t gsSize = 2 * sizeof(unsigned int);
907+
FLAGCXCHECKGOTO(deviceAdaptor->deviceMalloc(&gridSyncPtr, gsSize,
908+
flagcxMemDevice, NULL),
909+
res, fail);
910+
FLAGCXCHECKGOTO(deviceAdaptor->deviceMemset(gridSyncPtr, 0, gsSize,
911+
flagcxMemDevice, NULL),
912+
res, fail);
913+
}
914+
hostCopy._gridBarrierState = (unsigned int *)gridSyncPtr;
915+
916+
// Step 2: Copy flagcxDevComm to device
904917
FLAGCXCHECKGOTO(deviceAdaptor->deviceMalloc(&dPtr, sizeof(flagcxDevComm),
905918
flagcxMemDevice, NULL),
906919
res, fail);
@@ -909,7 +922,7 @@ static flagcxResult_t defaultDevApiCommGetDevicePtr(flagcxDevComm_t devComm,
909922
flagcxMemcpyHostToDevice, NULL, NULL),
910923
res, fail);
911924

912-
// Step 2: Allocate + construct net array on device
925+
// Step 3: Allocate + construct net array on device
913926
if (hostCopy._contextCount > 0 && flagcxDevNetSizeOf() > 0) {
914927
size_t netArraySize = hostCopy._contextCount * flagcxDevNetSizeOf();
915928
FLAGCXCHECKGOTO(deviceAdaptor->deviceMalloc(&netDevPtr, netArraySize,
@@ -927,12 +940,16 @@ static flagcxResult_t defaultDevApiCommGetDevicePtr(flagcxDevComm_t devComm,
927940

928941
devComm->cachedDevicePtr = dPtr;
929942
devComm->cachedNetContextsPtr = netDevPtr;
943+
devComm->cachedGridBarrierPtr = gridSyncPtr;
930944
*devPtr = dPtr;
931945
pthread_mutex_unlock(&devComm->cachedPtrMutex);
932946
return flagcxSuccess;
933947

934948
fail:
935949
pthread_mutex_unlock(&devComm->cachedPtrMutex);
950+
if (gridSyncPtr) {
951+
deviceAdaptor->deviceFree(gridSyncPtr, flagcxMemDevice, NULL);
952+
}
936953
if (netDevPtr) {
937954
deviceAdaptor->deviceFree(netDevPtr, flagcxMemDevice, NULL);
938955
}
@@ -949,10 +966,15 @@ static flagcxResult_t defaultDevApiCommFreeDevicePtr(flagcxDevComm_t devComm) {
949966
pthread_mutex_lock(&devComm->cachedPtrMutex);
950967
void *ptr = devComm->cachedDevicePtr;
951968
void *netPtr = devComm->cachedNetContextsPtr;
969+
void *gridPtr = devComm->cachedGridBarrierPtr;
952970
devComm->cachedDevicePtr = nullptr;
953971
devComm->cachedNetContextsPtr = nullptr;
972+
devComm->cachedGridBarrierPtr = nullptr;
954973
pthread_mutex_unlock(&devComm->cachedPtrMutex);
955974

975+
if (gridPtr) {
976+
FLAGCXCHECK(deviceAdaptor->deviceFree(gridPtr, flagcxMemDevice, NULL));
977+
}
956978
if (netPtr) {
957979
FLAGCXCHECK(deviceAdaptor->deviceFree(netPtr, flagcxMemDevice, NULL));
958980
}

0 commit comments

Comments
 (0)