[UT] Start supporting multiple group calls and graphs (#1151)

* Start supporting multiple group calls UT
Tá an tiomantas seo le fáil i:
Bertan Dogancay
2024-04-25 11:11:16 -06:00
tiomanta ag GitHub
tuismitheoir efe99057b0
tiomantas 0ec41f1386
D'athraigh 8 comhad le 566 breiseanna agus 240 scriosta
+213 -145
Féach ar an gComhad
@@ -97,8 +97,10 @@ namespace RcclUnitTesting
case CHILD_PREPARE_DATA : status = PrepareData(); break;
case CHILD_EXECUTE_COLL : status = ExecuteCollectives(); break;
case CHILD_VALIDATE_RESULTS: status = ValidateResults(); break;
case CHILD_LAUNCH_GRAPHS : status = LaunchGraphs(); break;
case CHILD_DEALLOCATE_MEM : status = DeallocateMem(); break;
case CHILD_DESTROY_COMMS : status = DestroyComms(); break;
case CHILD_DESTROY_GRAPHS : status = DestroyGraphs(); break;
case CHILD_STOP : goto stop;
default: exit(0);
}
@@ -144,6 +146,7 @@ namespace RcclUnitTesting
PIPE_READ(id);
PIPE_READ(this->totalRanks);
PIPE_READ(this->rankOffset);
PIPE_READ(this->numGroupCalls);
PIPE_READ(this->numCollectivesInGroup);
PIPE_READ(this->useBlocking);
bool useMultiRankPerGpu;
@@ -155,16 +158,29 @@ namespace RcclUnitTesting
PIPE_READ(numGpus);
this->deviceIds.resize(numGpus);
this->streams.clear();
this->streams.resize(numGpus);
this->collArgs.resize(numGpus);
for (int i = 0; i < numGpus; i++)
this->streams.resize(this->numGroupCalls);
this->collArgs.resize(this->numGroupCalls);
for (int i = 0; i < this->numGroupCalls; i++)
{
PIPE_READ(this->deviceIds[i]);
this->collArgs[i].clear();
this->collArgs[i].resize(numCollectivesInGroup);
this->streams[i].resize(numStreamsPerGroup);
this->collArgs[i].resize(numGpus);
this->streams[i].resize(numGpus);
for (int j = 0; j < numGpus; j++)
{
//PIPE_READ(this->deviceIds[j]);
this->collArgs[i][j].clear();
this->collArgs[i][j].resize(numCollectivesInGroup[i]);
this->streams[i][j].resize(numStreamsPerGroup[i]);
}
}
for (int i = 0; i < numGpus; i++)
PIPE_READ(this->deviceIds[i]);
// Initialize graphs
this->graphs.resize(this->numGroupCalls);
this->graphExecs.resize(this->numGroupCalls);
this->graphEnabled.resize(this->numGroupCalls);
// Initialize communicators
comms.clear();
comms.resize(numGpus);
@@ -172,52 +188,57 @@ namespace RcclUnitTesting
// Initialize within a group call to avoid deadlock when using multiple ranks per child
ErrCode status = TEST_SUCCESS;
CHILD_NCCL_CALL(ncclGroupStart(), "ncclGroupStart");
for (int localRank = 0; localRank < numGpus; ++localRank)
for (int groupCallIdx = 0; groupCallIdx < this->numGroupCalls; ++groupCallIdx)
{
int const globalRank = this->rankOffset + localRank;
int const currGpu = this->deviceIds[localRank];
if (hipSetDevice(currGpu) != hipSuccess)
for (int localRank = 0; localRank < numGpus; ++localRank)
{
ERROR("Rank %d on child %d unable to switch to GPU %d\n", globalRank, this->childId, currGpu);
status = TEST_FAIL;
break;
}
int const globalRank = this->rankOffset + localRank;
int const currGpu = this->deviceIds[localRank];
for (int i = 0; i < this->numStreamsPerGroup; i++)
{
if (hipStreamCreate(&(this->streams[localRank][i])) != hipSuccess)
if (hipSetDevice(currGpu) != hipSuccess)
{
ERROR("Rank %d on child %d unable to create stream %d for GPU %d\n", globalRank, this->childId, i, currGpu);
ERROR("Rank %d on child %d unable to switch to GPU %d\n", globalRank, this->childId, currGpu);
status = TEST_FAIL;
break;
}
}
if (useMultiRankPerGpu)
{
//if (ncclCommInitRankMulti(&this->comms[localRank], this->totalRanks, id, globalRank, globalRank) != ncclSuccess)
for (int i = 0; i < this->numStreamsPerGroup[groupCallIdx]; i++)
{
ERROR("Rank %d on child %d unable to call ncclCommInitRankMulti\n", globalRank, this->childId);
status = TEST_FAIL;
break;
if (hipStreamCreate(&(this->streams[groupCallIdx][localRank][i])) != hipSuccess)
{
ERROR("Rank %d on child %d unable to create stream %d for GPU %d in group %d\n", globalRank, this->childId, i, currGpu, groupCallIdx);
status = TEST_FAIL;
break;
}
}
}
else if (this->useBlocking == false)
{
// When non-blocking communicator is desired call ncclCommInitRankConfig with appropriate flag
ncclConfig_t config = NCCL_CONFIG_INITIALIZER;
config.blocking = 0;
ncclCommInitRankConfig(&this->comms[localRank], this->totalRanks, id, globalRank, &config);
CHILD_NCCL_CALL_NON_BLOCKING("ncclCommGetAsyncErrorInitRankConfig", localRank);
}
else
{
if (ncclCommInitRank(&this->comms[localRank], this->totalRanks, id, globalRank) != ncclSuccess)
{
ERROR("Rank %d on child %d unable to call ncclCommInitRank\n", globalRank, this->childId);
status = TEST_FAIL;
break;
if (groupCallIdx == 0) {
if (useMultiRankPerGpu)
{
//if (ncclCommInitRankMulti(&this->comms[localRank], this->totalRanks, id, globalRank, globalRank) != ncclSuccess)
{
ERROR("Rank %d on child %d unable to call ncclCommInitRankMulti\n", globalRank, this->childId);
status = TEST_FAIL;
break;
}
}
else if (this->useBlocking == false)
{
// When non-blocking communicator is desired call ncclCommInitRankConfig with appropriate flag
ncclConfig_t config = NCCL_CONFIG_INITIALIZER;
config.blocking = 0;
ncclCommInitRankConfig(&this->comms[localRank], this->totalRanks, id, globalRank, &config);
CHILD_NCCL_CALL_NON_BLOCKING("ncclCommGetAsyncErrorInitRankConfig", localRank);
}
else
{
if (ncclCommInitRank(&this->comms[localRank], this->totalRanks, id, globalRank) != ncclSuccess)
{
ERROR("Rank %d on child %d unable to call ncclCommInitRank\n", globalRank, this->childId);
status = TEST_FAIL;
break;
}
}
}
}
}
@@ -256,6 +277,7 @@ namespace RcclUnitTesting
// Read values sent by parent [see TestBed::SetCollectiveArgs()]
int globalRank;
int collId;
int groupId;
ncclFunc_t funcType;
ncclDataType_t dataType;
size_t numInputElements;
@@ -265,6 +287,7 @@ namespace RcclUnitTesting
PIPE_READ(globalRank);
PIPE_READ(collId);
PIPE_READ(groupId);
PIPE_READ(funcType);
PIPE_READ(dataType);
PIPE_READ(numInputElements);
@@ -280,19 +303,19 @@ namespace RcclUnitTesting
int const localRank = globalRank - rankOffset;
CHECK_HIP(hipSetDevice(this->deviceIds[localRank]));
for (int collIdx = 0; collIdx < collArgs[localRank].size(); ++collIdx)
for (int collIdx = 0; collIdx < collArgs[groupId][localRank].size(); ++collIdx)
{
if (collId == -1 || collId == collIdx)
{
CollectiveArgs& collArg = this->collArgs[localRank][collIdx];
CollectiveArgs& collArg = this->collArgs[groupId][localRank][collIdx];
CHECK_CALL(collArg.SetArgs(globalRank, this->totalRanks,
this->deviceIds[localRank],
funcType, dataType,
numInputElements, numOutputElements,
streamIdx,
options));
if (this->verbose) INFO("Rank %d on child %d sets collective %d [%s]\n",
globalRank, this->childId, collIdx,
if (this->verbose) INFO("Rank %d on child %d sets collective %d in group %d [%s]\n",
globalRank, this->childId, collIdx, groupId,
collArg.GetDescription().c_str());
// If pre-mult scalars are provided, then create a custom reduction operator
@@ -304,8 +327,8 @@ namespace RcclUnitTesting
(ncclScalarResidence_t)options.scalarMode,
this->comms[localRank]),
"ncclRedOpCreatePreMulSum");
if (verbose) INFO("Child %d created custom redop %d for collective %d\n",
this->childId, collArg.options.redOp, collIdx);
if (verbose) INFO("Child %d created custom redop %d for group %d collective %d\n",
this->childId, collArg.options.redOp, groupId, collIdx);
}
}
}
@@ -322,11 +345,13 @@ namespace RcclUnitTesting
int collId;
bool inPlace;
bool useManagedMem;
int groupId;
PIPE_READ(globalRank);
PIPE_READ(collId);
PIPE_READ(inPlace);
PIPE_READ(useManagedMem);
PIPE_READ(groupId);
if (globalRank < this->rankOffset || (this->rankOffset + comms.size() <= globalRank))
{
@@ -336,14 +361,14 @@ namespace RcclUnitTesting
int const localRank = globalRank - rankOffset;
CHECK_HIP(hipSetDevice(this->deviceIds[localRank]));
for (int collIdx = 0; collIdx < collArgs[localRank].size(); ++collIdx)
for (int collIdx = 0; collIdx < collArgs[groupId][localRank].size(); ++collIdx)
{
if (collId == -1 || collId == collIdx)
{
CollectiveArgs& collArg = this->collArgs[localRank][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 on device %d (%s,%s) Input: %p Output %p\n",
globalRank, this->childId, collIdx, this->deviceIds[localRank],
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",
globalRank, this->childId, collIdx, groupId, this->deviceIds[localRank],
inPlace ? "in-place" : "out-of-place",
useManagedMem ? "managed" : "unmanaged",
collArg.inputGpu.ptr,
@@ -363,9 +388,11 @@ namespace RcclUnitTesting
// Read values sent by parent [see TestBed::PrepareData()]
int globalRank;
int collId;
int groupId;
CollFuncPtr prepDataFunc;
PIPE_READ(globalRank);
PIPE_READ(groupId);
PIPE_READ(collId);
PIPE_READ(prepDataFunc);
@@ -378,13 +405,13 @@ namespace RcclUnitTesting
int const localRank = globalRank - rankOffset;
CHECK_HIP(hipSetDevice(this->deviceIds[localRank]));
for (int collIdx = 0; collIdx < collArgs[localRank].size(); ++collIdx)
for (int collIdx = 0; collIdx < collArgs[groupId][localRank].size(); ++collIdx)
{
if (collId == -1 || collId == collIdx)
{
if (this->verbose) INFO("Rank %d on child %d prepares data for collective %d\n",
globalRank, this->childId, collIdx);
CHECK_CALL(this->collArgs[localRank][collIdx].PrepareData(prepDataFunc));
if (this->verbose) INFO("Rank %d on child %d prepares data for collective %d in group %d\n",
globalRank, this->childId, collIdx, groupId);
CHECK_CALL(this->collArgs[groupId][localRank][collIdx].PrepareData(prepDataFunc));
}
}
if (this->verbose) INFO("Child %d finishes PrepareData()\n", this->childId);
@@ -394,9 +421,11 @@ namespace RcclUnitTesting
ErrCode TestBedChild::ExecuteCollectives()
{
int timeoutUs = 0;
PIPE_READ(timeoutUs);
int groupId = 0;
bool useHipGraph = false;
PIPE_READ(timeoutUs);
PIPE_READ(groupId);
PIPE_READ(useHipGraph);
int numRanksToExecute, tempRank;
@@ -420,14 +449,14 @@ namespace RcclUnitTesting
}
numRanksToExecute = (int)localRanksToExecute.size();
std::vector<std::vector<hipGraph_t>> graphs;
std::vector<std::vector<hipGraphExec_t>> graphExec;
graphs.resize(numRanksToExecute);
graphExec.resize(numRanksToExecute);
this->graphs[groupId].resize(numRanksToExecute);
this->graphExecs[groupId].resize(numRanksToExecute);
this->graphEnabled[groupId].resize(numRanksToExecute);
for (int i = 0; i < numRanksToExecute; i++)
{
graphs[i].resize(this->numStreamsPerGroup);
graphExec[i].resize(this->numStreamsPerGroup);
this->graphs[groupId][i].resize(this->numStreamsPerGroup[groupId]);
this->graphExecs[groupId][i].resize(this->numStreamsPerGroup[groupId]);
this->graphEnabled[groupId][i].resize(this->numStreamsPerGroup[groupId]);
}
// Start HIP graph stream capture if requested
@@ -435,11 +464,11 @@ namespace RcclUnitTesting
{
for (int localRank : localRanksToExecute)
{
if (this->verbose) INFO("Capturing stream for rank %d\n", localRank);
if (this->verbose) INFO("Capturing stream for group %d rank %d\n", groupId, localRank);
CHECK_HIP(hipSetDevice(this->deviceIds[localRank]));
for (int i = 0; i < this->numStreamsPerGroup; i++)
for (int i = 0; i < this->numStreamsPerGroup[groupId]; i++)
{
CHECK_HIP(hipStreamBeginCapture(this->streams[localRank][i], hipStreamCaptureModeRelaxed));
CHECK_HIP(hipStreamBeginCapture(this->streams[groupId][localRank][i], hipStreamCaptureModeRelaxed));
}
}
}
@@ -448,14 +477,14 @@ namespace RcclUnitTesting
CHILD_NCCL_CALL(ncclGroupStart(), "ncclGroupStart");
// Loop over all collectives to be executed in group call
for (int collId = 0; collId < this->numCollectivesInGroup; ++collId)
for (int collId = 0; collId < this->numCollectivesInGroup[groupId]; ++collId)
{
// Loop over all local ranks
for (int localRank : localRanksToExecute)
{
CHECK_HIP(hipSetDevice(this->deviceIds[localRank]));
CollectiveArgs const& collArg = this->collArgs[localRank][collId];
CollectiveArgs const& collArg = this->collArgs[groupId][localRank][collId];
if (this->printValues && !useHipGraph)
{
@@ -464,14 +493,14 @@ namespace RcclUnitTesting
size_t const numInputBytes = numInputElementsToPrint * DataTypeToBytes(collArg.dataType);
inputCpu.AllocateCpuMem(numInputBytes);
CHECK_HIP(hipMemcpy(inputCpu.ptr, collArg.inputGpu.ptr, numInputBytes, hipMemcpyDeviceToHost));
printf("[ DEBUG ] Rank %02d Coll %d %-10s: %s\n", collArg.globalRank, collId, "Input",
printf("[ DEBUG ] Rank %02d Group %d Coll %d %-10s: %s\n", collArg.globalRank, groupId, collId, "Input",
inputCpu.ToString(collArg.dataType, numInputElementsToPrint).c_str());
inputCpu.FreeCpuMem();
int const numOutputElementsToPrint = (this->printValues < 0 ? collArg.numOutputElements : this->printValues);
size_t const numOutputBytes = numOutputElementsToPrint * DataTypeToBytes(collArg.dataType);
CHECK_HIP(hipMemcpy(collArg.outputCpu.ptr, collArg.outputGpu.ptr, numOutputBytes, hipMemcpyDeviceToHost));
printf("[ DEBUG ] Rank %02d Coll %d %-10s: %s\n", collArg.globalRank, collId, "Pre-Output",
printf("[ DEBUG ] Rank %02d Group %d Coll %d %-10s: %s\n", collArg.globalRank, groupId, collId, "Pre-Output",
collArg.outputCpu.ToString(collArg.dataType, numOutputElementsToPrint).c_str());
}
@@ -484,7 +513,7 @@ namespace RcclUnitTesting
collArg.dataType,
collArg.options.root,
this->comms[localRank],
this->streams[localRank][collArg.streamIdx]),
this->streams[groupId][localRank][collArg.streamIdx]),
"ncclBroadcast");
break;
case ncclCollReduce:
@@ -495,7 +524,7 @@ namespace RcclUnitTesting
collArg.options.redOp,
collArg.options.root,
this->comms[localRank],
this->streams[localRank][collArg.streamIdx]),
this->streams[groupId][localRank][collArg.streamIdx]),
"ncclReduce");
break;
case ncclCollAllGather:
@@ -504,7 +533,7 @@ namespace RcclUnitTesting
collArg.numInputElements,
collArg.dataType,
this->comms[localRank],
this->streams[localRank][collArg.streamIdx]),
this->streams[groupId][localRank][collArg.streamIdx]),
"ncclAllGather");
break;
case ncclCollReduceScatter:
@@ -514,7 +543,7 @@ namespace RcclUnitTesting
collArg.dataType,
collArg.options.redOp,
this->comms[localRank],
this->streams[localRank][collArg.streamIdx]),
this->streams[groupId][localRank][collArg.streamIdx]),
"ncclReduceScatter");
break;
case ncclCollAllReduce:
@@ -524,7 +553,7 @@ namespace RcclUnitTesting
collArg.dataType,
collArg.options.redOp,
this->comms[localRank],
this->streams[localRank][collArg.streamIdx]),
this->streams[groupId][localRank][collArg.streamIdx]),
"ncclAllReduce");
break;
case ncclCollGather:
@@ -534,7 +563,7 @@ namespace RcclUnitTesting
collArg.dataType,
collArg.options.root,
this->comms[localRank],
this->streams[localRank][collArg.streamIdx]),
this->streams[groupId][localRank][collArg.streamIdx]),
"ncclGather");
break;
case ncclCollScatter:
@@ -544,7 +573,7 @@ namespace RcclUnitTesting
collArg.dataType,
collArg.options.root,
this->comms[localRank],
this->streams[localRank][collArg.streamIdx]),
this->streams[groupId][localRank][collArg.streamIdx]),
"ncclScatter");
break;
case ncclCollAllToAll:
@@ -553,7 +582,7 @@ namespace RcclUnitTesting
collArg.numInputElements / collArg.totalRanks,
collArg.dataType,
this->comms[localRank],
this->streams[localRank][collArg.streamIdx]),
this->streams[groupId][localRank][collArg.streamIdx]),
"ncclAllToAll");
break;
case ncclCollAllToAllv:
@@ -565,7 +594,7 @@ namespace RcclUnitTesting
collArg.options.rdispls + (this->rankOffset + localRank)*this->totalRanks,
collArg.dataType,
this->comms[localRank],
this->streams[localRank][collArg.streamIdx]),
this->streams[groupId][localRank][collArg.streamIdx]),
"ncclAllToAllv");
break;
case ncclCollSend:
@@ -574,7 +603,7 @@ namespace RcclUnitTesting
collArg.dataType,
collArg.options.root,
this->comms[localRank],
this->streams[localRank][collArg.streamIdx]),
this->streams[groupId][localRank][collArg.streamIdx]),
"ncclSend");
break;
case ncclCollRecv:
@@ -583,7 +612,7 @@ namespace RcclUnitTesting
collArg.dataType,
collArg.options.root,
this->comms[localRank],
this->streams[localRank][collArg.streamIdx]),
this->streams[groupId][localRank][collArg.streamIdx]),
"ncclRecv");
break;
default:
@@ -624,33 +653,24 @@ namespace RcclUnitTesting
{
if (this->verbose) INFO("Ending stream capture for rank %d\n", localRank);
CHECK_HIP(hipSetDevice(this->deviceIds[localRank]));
for (int i = 0; i < this->numStreamsPerGroup; i++)
for (int i = 0; i < this->numStreamsPerGroup[groupId]; i++)
{
CHECK_HIP(hipStreamEndCapture(this->streams[localRank][i], &graphs[localRank][i]));
CHECK_HIP(hipStreamEndCapture(this->streams[groupId][localRank][i], &this->graphs[groupId][localRank][i]));
if (this->verbose)
{
size_t numNodes;
hipGraphNode_t* nodes;
CHECK_HIP(hipGraphGetNodes(graphs[localRank][i], nodes, &numNodes));
INFO("Graph for rank %d stream %d has %lu nodes\n", localRank, i, numNodes);
}
// if (this->verbose)
// {
// size_t numNodes;
// hipGraphNode_t* nodes;
// CHECK_HIP(hipGraphGetNodes(graphs[localRank][i], nodes, &numNodes));
// INFO("Graph for rank %d stream %d has %lu nodes\n", localRank, i, numNodes);
// }
}
if (this->verbose) INFO("Instantiating executable graph for rank %d\n", localRank);
for (int i = 0; i < this->numStreamsPerGroup; i++)
if (this->verbose) INFO("Instantiating executable graph for group %d rank %d\n", groupId, localRank);
for (int i = 0; i < this->numStreamsPerGroup[groupId]; i++)
{
CHECK_HIP(hipGraphInstantiate(&graphExec[localRank][i], graphs[localRank][i], NULL, NULL, 0));
}
}
for (int localRank : localRanksToExecute)
{
if (this->verbose) INFO("Launch graph for rank %d\n", localRank);
CHECK_HIP(hipSetDevice(this->deviceIds[localRank]));
for (int i = 0; i < this->numStreamsPerGroup; i++)
{
CHECK_HIP(hipGraphLaunch(graphExec[localRank][i], this->streams[localRank][i]));
CHECK_HIP(hipGraphInstantiate(&this->graphExecs[groupId][localRank][i], this->graphs[groupId][localRank][i], NULL, NULL, 0));
graphEnabled[groupId][localRank][i] = true;
}
}
}
@@ -664,8 +684,8 @@ namespace RcclUnitTesting
std::vector<hipStream_t> streamsToComplete;
for (int localRank : localRanksToExecute)
{
for (int i = 0; i < this->numStreamsPerGroup; i++)
streamsToComplete.push_back(this->streams[localRank][i]);
for (int i = 0; i < this->numStreamsPerGroup[groupId]; i++)
streamsToComplete.push_back(this->streams[groupId][localRank][i]);
}
int usElapsed = 0, timedout = 0;
using namespace std::chrono;
@@ -701,40 +721,25 @@ namespace RcclUnitTesting
// of fencing between kernels and at hipStreamQuery
for (int localRank : localRanksToExecute)
{
if (this->verbose) INFO("Starting synchronization for rank %d\n", localRank);
for (int i = 0; i < this->numStreamsPerGroup; i++)
CHECK_HIP(hipStreamSynchronize(this->streams[localRank][i]));
}
// Destroy graphs
if (useHipGraph)
{
for (int localRank : localRanksToExecute)
{
if (this->verbose) INFO("Destroying graphs for rank %d\n", localRank);
CHECK_HIP(hipSetDevice(this->deviceIds[localRank]));
for (int i = 0; i < this->numStreamsPerGroup; i++)
{
CHECK_HIP(hipGraphDestroy(graphs[localRank][i]));
CHECK_HIP(hipGraphExecDestroy(graphExec[localRank][i]));
}
}
if (this->verbose) INFO("Starting synchronization for group %d rank %d\n", groupId, localRank);
for (int i = 0; i < this->numStreamsPerGroup[groupId]; i++)
CHECK_HIP(hipStreamSynchronize(this->streams[groupId][localRank][i]));
}
if (this->printValues)
{
for (int collId = 0; collId < this->numCollectivesInGroup; ++collId)
for (int collId = 0; collId < this->numCollectivesInGroup[groupId]; ++collId)
for (int localRank : localRanksToExecute)
{
CollectiveArgs const& collArg = this->collArgs[localRank][collId];
CollectiveArgs const& collArg = this->collArgs[groupId][localRank][collId];
CHECK_HIP(hipSetDevice(this->deviceIds[localRank]));
int numOutputElementsToPrint = (this->printValues < 0 ? collArg.numOutputElements : this->printValues);
size_t const numOutputBytes = numOutputElementsToPrint * DataTypeToBytes(collArg.dataType);
CHECK_HIP(hipMemcpy(collArg.outputCpu.ptr, collArg.outputGpu.ptr, numOutputBytes, hipMemcpyDeviceToHost));
printf("[ DEBUG ] Rank %02d Coll %d %-10s: %s\n", collArg.globalRank, collId, "Output",
printf("[ DEBUG ] Rank %02d Group %d Coll %d %-10s: %s\n", collArg.globalRank, groupId, collId, "Output",
collArg.outputCpu.ToString(collArg.dataType, numOutputElementsToPrint).c_str());
printf("[ DEBUG ] Rank %02d Coll %d %-10s: %s\n", collArg.globalRank, collId, "Expected",
printf("[ DEBUG ] Rank %02d Group %d Coll %d %-10s: %s\n", collArg.globalRank, groupId, collId, "Expected",
collArg.expected.ToString(collArg.dataType, numOutputElementsToPrint).c_str());
}
}
@@ -752,8 +757,9 @@ namespace RcclUnitTesting
ErrCode TestBedChild::ValidateResults()
{
// Read values sent by parent [see TestBed::ValidateResults()]
int globalRank, collId;
int globalRank, groupId, collId;
PIPE_READ(globalRank);
PIPE_READ(groupId);
PIPE_READ(collId);
if (this->verbose) INFO("Child %d begins ValidateResults()\n", this->childId);
@@ -767,15 +773,15 @@ namespace RcclUnitTesting
CHECK_HIP(hipSetDevice(this->deviceIds[localRank]));
ErrCode status = TEST_SUCCESS;
for (int collIdx = 0; collIdx < collArgs[localRank].size(); ++collIdx)
for (int collIdx = 0; collIdx < collArgs[groupId][localRank].size(); ++collIdx)
{
if (collId == -1 || collId == collIdx)
{
if (this->verbose) INFO("Rank %d on child %d validating collective %d results\n",
globalRank, this->childId, collIdx);
if (this->collArgs[localRank][collIdx].ValidateResults() != TEST_SUCCESS)
if (this->verbose) INFO("Rank %d on child %d validating collective %d in group %d results\n",
globalRank, this->childId, collIdx, groupId);
if (this->collArgs[groupId][localRank][collIdx].ValidateResults() != TEST_SUCCESS)
{
ERROR("Rank %d Collective %d output does not match expected\n", globalRank, collIdx);
ERROR("Rank %d Group %d Collective %d output does not match expected\n", globalRank, groupId, collIdx);
status = TEST_FAIL;
}
}
@@ -785,13 +791,35 @@ namespace RcclUnitTesting
return status;
}
ErrCode TestBedChild::LaunchGraphs()
{
int groupId;
PIPE_READ(groupId);
if (this->verbose) INFO("Child %d begins LaunchGraphs for group %d\n", this->childId, groupId);
for (int localRank = 0; localRank < this->deviceIds.size(); ++localRank) {
CHECK_HIP(hipSetDevice(this->deviceIds[localRank]));
for (int streamIdx = 0; streamIdx < this->numStreamsPerGroup[groupId]; ++streamIdx)
{
if (this->verbose) INFO("Launch graph for group %d rank %d stream %d\n", groupId, localRank, streamIdx);
CHECK_HIP(hipGraphLaunch(this->graphExecs[groupId][localRank][streamIdx], this->streams[groupId][localRank][streamIdx]));
}
}
if (this->verbose) INFO("Child %d finishes LaunchGraphs for group %d\n", this->childId, groupId);
return TEST_SUCCESS;
}
ErrCode TestBedChild::DeallocateMem()
{
if (this->verbose) INFO("Child %d begins DeallocateMem\n", this->childId);
// Read values sent by parent [see TestBed::DeallocateMem()]
int globalRank, collId;
int globalRank, groupId, collId;
PIPE_READ(globalRank);
PIPE_READ(groupId);
PIPE_READ(collId);
if (globalRank < this->rankOffset || (this->rankOffset + comms.size() <= globalRank))
@@ -802,15 +830,15 @@ namespace RcclUnitTesting
int const localRank = globalRank - rankOffset;
CHECK_HIP(hipSetDevice(this->deviceIds[localRank]));
for (int collIdx = 0; collIdx < collArgs[localRank].size(); ++collIdx)
for (int collIdx = 0; collIdx < collArgs[groupId][localRank].size(); ++collIdx)
{
CollectiveArgs& collArg = this->collArgs[localRank][collIdx];
CollectiveArgs& collArg = this->collArgs[groupId][localRank][collIdx];
if (collId == -1 || collId == collIdx)
{
if (this->verbose)
{
INFO("Child %d release memory for collective %d (Input: %p Output %p\n",
this->childId, collIdx, collArg.inputGpu.ptr, collArg.outputGpu.ptr);
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);
}
CHECK_CALL(collArg.DeallocateMem());
@@ -819,8 +847,8 @@ namespace RcclUnitTesting
{
CHILD_NCCL_CALL(ncclRedOpDestroy(collArg.options.redOp, this->comms[localRank]),
"ncclRedOpDestroy");
if (verbose) INFO("Child %d destroys custom redop %d for collective %d\n",
this->childId, collArg.options.redOp, collIdx);
if (verbose) INFO("Child %d destroys custom redop %d for collective %d in group %d\n",
this->childId, collArg.options.redOp, collIdx, groupId);
}
}
if (this->verbose) INFO("Child %d finishes DeallocateMem\n", this->childId);
@@ -852,11 +880,14 @@ namespace RcclUnitTesting
{
CHILD_NCCL_CALL(ncclCommDestroy(this->comms[i]), "ncclCommDestroy");
}
for (int i = 0; i < this->streams.size(); ++i)
for (int i = 0; i < this->numGroupCalls; ++i)
{
for (int j = 0; j < this->numStreamsPerGroup; j++)
for (int j = 0; j < this->streams[i].size(); ++j)
{
CHECK_HIP(hipStreamDestroy(this->streams[i][j]));
for (int k = 0; k < this->streams[i][j].size(); ++k)
{
CHECK_HIP(hipStreamDestroy(this->streams[i][j][k]));
}
}
}
this->comms.clear();
@@ -864,4 +895,41 @@ namespace RcclUnitTesting
if (this->verbose) INFO("Child %d finishes DestroyComms\n", this->childId);
return TEST_SUCCESS;
}
ErrCode TestBedChild::DestroyGraphs()
{
if (this->verbose) INFO("Child %d begins DestroyGraphs\n", this->childId);
int groupId;
PIPE_READ(groupId);
// Release graphs
for (int localRank = 0; localRank < this->deviceIds.size(); ++localRank)
{
CHECK_HIP(hipSetDevice(this->deviceIds[localRank]));
for (int streamIdx = 0; streamIdx < this->numStreamsPerGroup[groupId]; ++streamIdx)
{
if (graphEnabled[groupId][localRank][streamIdx])
{
if (this->verbose) INFO("Destroying graphs for group %d rank %d stream %d\n", groupId, localRank, streamIdx);
CHECK_HIP(hipGraphDestroy(this->graphs[groupId][localRank][streamIdx]));
CHECK_HIP(hipGraphExecDestroy(this->graphExecs[groupId][localRank][streamIdx]));
}
}
}
for (int localRank = 0; localRank < this->deviceIds.size(); ++localRank)
{
for (int i = 0; i < this->numStreamsPerGroup[groupId]; ++i)
CHECK_HIP(hipStreamSynchronize(this->streams[groupId][localRank][i]));
}
this->graphs[groupId].clear();
this->graphExecs[groupId].clear();
this->graphEnabled[groupId].clear();
if (this->verbose) INFO("Child %d finishes DestroyGraphs\n", this->childId);
return TEST_SUCCESS;
}
}