Adjustment for UT Sendrecv (#1400)

Enabled UT sendrecv to same rank and refactor UBR call

[ROCm/rccl commit: fd9924cfe7]
Tá an tiomantas seo le fáil i:
Tim
2024-10-30 15:13:53 -04:00
tiomanta ag GitHub
tuismitheoir e1c20e7f24
tiomantas e346e19065
D'athraigh 2 comhad le 83 breiseanna agus 56 scriosta
+6 -6
Féach ar an gComhad
@@ -395,6 +395,8 @@ namespace RcclUnitTesting
{
CollectiveArgs& collArg = this->collArgs[groupId][localRank][collIdx];
CHECK_CALL(collArg.AllocateMem(inPlace, useManagedMem, userRegistered));
if (collArg.userRegistered && (collArg.funcType == ncclCollSend || collArg.funcType == ncclCollRecv))
CHILD_NCCL_CALL(ncclCommRegister(this->comms[localRank], collArg.inputGpu.ptr, collArg.numInputBytesAllocated, &(collArg.commRegHandle)),"ncclCommRegister");
if (this->verbose) INFO("Rank %d on child %d allocates memory for collective %d in group %d on device %d (%s,%s,%s) Input: %p Output %p\n",
globalRank, this->childId, collIdx, groupId, this->deviceIds[localRank],
inPlace ? "in-place" : "out-of-place",
@@ -646,8 +648,6 @@ namespace RcclUnitTesting
"ncclAllToAllv");
break;
case ncclCollSend:
if (collArg.userRegistered)
CHILD_NCCL_CALL_RANK(errCode, ncclCommRegister(this->comms[localRank], collArg.inputGpu.ptr, collArg.numInputBytesAllocated, &(collArg.commRegHandle)),"ncclCommRegister");
CHILD_NCCL_CALL_RANK(errCode, ncclSend(
collArg.inputGpu.ptr,
collArg.numInputElements,
@@ -658,8 +658,6 @@ namespace RcclUnitTesting
"ncclSend");
break;
case ncclCollRecv:
if (collArg.userRegistered)
CHILD_NCCL_CALL_RANK(errCode, ncclCommRegister(this->comms[localRank], collArg.outputGpu.ptr, collArg.numOutputBytesAllocated, &(collArg.commRegHandle)), "ncclCommRegister");
CHILD_NCCL_CALL_RANK(errCode, ncclRecv(
collArg.outputGpu.ptr,
collArg.numOutputElements,
@@ -891,8 +889,6 @@ namespace RcclUnitTesting
for (int collIdx = 0; collIdx < collArgs[groupId][localRank].size(); ++collIdx)
{
CollectiveArgs& collArg = this->collArgs[groupId][localRank][collIdx];
if (collArg.userRegistered && (collArg.funcType == ncclCollSend || collArg.funcType == ncclCollRecv))
CHILD_NCCL_CALL(ncclCommDeregister(this->comms[localRank], collArg.commRegHandle), "ncclCommDeregister");
if (collId == -1 || collId == collIdx)
{
if (this->verbose)
@@ -900,6 +896,10 @@ namespace RcclUnitTesting
INFO("Child %d release memory for collective %d in group %d (Input: %p Output %p\n",
this->childId, collIdx, groupId, collArg.inputGpu.ptr, collArg.outputGpu.ptr);
}
if (collArg.userRegistered && (collArg.funcType == ncclCollSend || collArg.funcType == ncclCollRecv))
{
CHILD_NCCL_CALL(ncclCommDeregister(this->comms[localRank], collArg.commRegHandle), "ncclCommDeregister");
}
CHECK_CALL(collArg.DeallocateMem());
}