Adding User Buffer Registration support for Unit test (#1199)

* Adding UBR support for UT SendRecv

Signed-off-by: Tim Hu <timhu102@amd.com>

* Update test/common/TestBedChild.cpp

Co-authored-by: corey-derochie-amd <161367113+corey-derochie-amd@users.noreply.github.com>

---------

Signed-off-by: Tim Hu <timhu102@amd.com>
Co-authored-by: corey-derochie-amd <161367113+corey-derochie-amd@users.noreply.github.com>
Bu işleme şunda yer alıyor:
Tim
2024-07-30 13:39:25 -04:00
işlemeyi yapan: GitHub
ebeveyn ca5341d419
işleme a4793286c7
8 değiştirilmiş dosya ile 137 ekleme ve 26 silme
+12 -3
Dosyayı Görüntüle
@@ -371,12 +371,14 @@ namespace RcclUnitTesting
int collId;
bool inPlace;
bool useManagedMem;
bool userRegistered;
int groupId;
PIPE_READ(globalRank);
PIPE_READ(collId);
PIPE_READ(inPlace);
PIPE_READ(useManagedMem);
PIPE_READ(userRegistered);
PIPE_READ(groupId);
if (globalRank < this->rankOffset || (this->rankOffset + comms.size() <= globalRank))
@@ -392,11 +394,12 @@ namespace RcclUnitTesting
if (collId == -1 || collId == collIdx)
{
CollectiveArgs& collArg = this->collArgs[groupId][localRank][collIdx];
CHECK_CALL(collArg.AllocateMem(inPlace, useManagedMem));
if (this->verbose) INFO("Rank %d on child %d allocates memory for collective %d in group %d on device %d (%s,%s) Input: %p Output %p\n",
CHECK_CALL(collArg.AllocateMem(inPlace, useManagedMem, userRegistered));
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",
useManagedMem ? "managed" : "unmanaged",
userRegistered ? "user registered buffer" : "internal copy",
collArg.inputGpu.ptr,
collArg.outputGpu.ptr);
}
@@ -520,7 +523,7 @@ namespace RcclUnitTesting
CHECK_HIP_RANK(errCode, hipSetDevice(this->deviceIds[localRank]));
CollectiveArgs const& collArg = this->collArgs[groupId][localRank][collId];
CollectiveArgs& collArg = this->collArgs[groupId][localRank][collId];
if (this->printValues && !useHipGraph)
{
@@ -643,6 +646,8 @@ 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,
@@ -653,6 +658,8 @@ 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,
@@ -884,6 +891,8 @@ 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)