@@ -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
934948fail:
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