Adding interactive mode for unit tests (UT_INTERACTIVE) (#715)

Этот коммит содержится в:
gilbertlee-amd
2023-03-21 10:58:24 -06:00
коммит произвёл GitHub
родитель 9a0d4a07a6
Коммит 00c3d8d850
5 изменённых файлов: 52 добавлений и 2 удалений
+3 -1
Просмотреть файл
@@ -61,6 +61,7 @@ namespace RcclUnitTesting
printValues = GetEnvVar("UT_PRINT_VALUES", 0);
maxRanksPerGpu = GetEnvVar("UT_MAX_RANKS_PER_GPU", 1);
showTiming = GetEnvVar("UT_SHOW_TIMING", 1);
useInteractive = GetEnvVar("UT_INTERACTIVE", 0);
// Limit number of supported reduction operators to just ncclSum if only allReduce is built
#ifdef BUILD_ALLREDUCE_ONLY
@@ -170,7 +171,8 @@ namespace RcclUnitTesting
std::make_pair("UT_DATATYPES" , "List of datatypes to test"),
std::make_pair("UT_MAX_RANKS_PER_GPU", "Maximum number of ranks using the same GPU"),
std::make_pair("UT_PRINT_VALUES" , "Print array values (# of values to print, < 0 for all)"),
std::make_pair("UT_SHOW_TIMING" , "Show timing table")
std::make_pair("UT_SHOW_TIMING" , "Show timing table"),
std::make_pair("UT_INTERACTIVE" , "Run in interactive mode")
};
printf("================================================================================\n");
+1
Просмотреть файл
@@ -26,6 +26,7 @@ namespace RcclUnitTesting
int printValues; // Print out input/output/expected arrays [UT_PRINT_VALUES]
int maxRanksPerGpu; // Number of ranks using the same GPU [UT_MAX_RANKS_PER_GPU]
bool showTiming; // Show timing per case at end [UT_SHOW_TIMING]
bool useInteractive; // Run in interactive mode [UT_INTERACTIVE]
// Constructor that parses and collects environment variables
EnvVars();
+40
Просмотреть файл
@@ -90,6 +90,8 @@ namespace RcclUnitTesting
bool const useBlocking,
int const numStreamsPerGroup)
{
InteractiveWait("Starting InitComms");
// Count up the total number of GPUs to use and track child/deviceId per rank
this->numActiveChildren = deviceIdsPerProcess.size();
this->numActiveRanks = 0;
@@ -99,6 +101,7 @@ namespace RcclUnitTesting
this->rankToChildMap.clear();
this->rankToDeviceMap.clear();
if (ev.verbose) INFO("Setting up %d active child processes\n", this->numActiveChildren);
for (int childId = 0; childId < this->numActiveChildren; ++childId)
{
for (auto i = 0; i < deviceIdsPerProcess[childId].size(); ++i)
@@ -168,6 +171,7 @@ namespace RcclUnitTesting
{
PIPE_CHECK(childId);
}
InteractiveWait("Finishing InitComms");
}
void TestBed::InitComms(int const numGpus, int const numCollectivesInGroup, bool const useBlocking, int const numStreamsPerGroup)
@@ -184,6 +188,7 @@ namespace RcclUnitTesting
int const rank,
int const streamIdx)
{
InteractiveWait("Starting SetCollectiveArgs");
// Build list of ranks this applies to (-1 for rank means to set for all)
std::vector<int> rankList;
for (int i = 0; i < this->numActiveRanks; ++i)
@@ -211,6 +216,7 @@ namespace RcclUnitTesting
PIPE_WRITE(childId, optionalArgs);
PIPE_CHECK(childId);
}
InteractiveWait("Finishing SetCollectiveArgs");
}
void TestBed::AllocateMem(bool const inPlace,
@@ -218,6 +224,8 @@ namespace RcclUnitTesting
int const collId,
int const rank)
{
InteractiveWait("Starting AllocateMem");
// Build list of ranks this applies to (-1 for rank means to set for all)
std::vector<int> rankList;
for (int i = 0; i < this->numActiveRanks; ++i)
@@ -235,12 +243,14 @@ namespace RcclUnitTesting
PIPE_WRITE(childId, useManagedMem);
PIPE_CHECK(childId);
}
InteractiveWait("Finishing AllocateMem");
}
void TestBed::PrepareData(int const collId,
int const rank,
CollFuncPtr const prepDataFunc)
{
InteractiveWait("Starting PrepareData");
// Build list of ranks this applies to (-1 for rank means to set for all)
std::vector<int> rankList;
for (int i = 0; i < this->numActiveRanks; ++i)
@@ -257,10 +267,13 @@ namespace RcclUnitTesting
PIPE_WRITE(childId, prepDataFunc);
PIPE_CHECK(childId);
}
InteractiveWait("Finishing PrepareData");
}
void TestBed::ExecuteCollectives(std::vector<int> const &currentRanks, bool const useHipGraph)
{
InteractiveWait("Starting ExecuteCollectives");
int const cmd = TestBedChild::CHILD_EXECUTE_COLL;
++TestBed::NumTestsRun();
@@ -275,6 +288,7 @@ namespace RcclUnitTesting
{
if ((currentRanks.size() == 0) || (ranksPerChild[childId].size() > 0))
{
InteractiveWait("Starting ExecuteCollectives for child " + std::to_string(childId));
PIPE_WRITE(childId, cmd);
PIPE_WRITE(childId, useHipGraph);
int tempCurrentRanks = currentRanks.size();
@@ -290,10 +304,14 @@ namespace RcclUnitTesting
{
if ((currentRanks.size() == 0) || (ranksPerChild[childId].size() > 0)) PIPE_CHECK(childId);
}
InteractiveWait("Finishing ExecuteCollectives");
}
void TestBed::ValidateResults(bool& isCorrect, int const collId, int const rank)
{
InteractiveWait("Starting ValidateResults");
// Build list of ranks this applies to (-1 for rank means to set for all)
std::vector<int> rankList;
for (int i = 0; i < this->numActiveRanks; ++i)
@@ -316,10 +334,14 @@ namespace RcclUnitTesting
}
ASSERT_EQ(isCorrect, true) << "Output does not match expected";
InteractiveWait("Finishing ValidateResults");
}
void TestBed::DeallocateMem(int const collId, int const rank)
{
InteractiveWait("Starting ValidateResults");
// Build list of ranks this applies to (-1 for rank means to set for all)
std::vector<int> rankList;
for (int i = 0; i < this->numActiveRanks; ++i)
@@ -335,10 +357,14 @@ namespace RcclUnitTesting
PIPE_WRITE(childId, collId);
PIPE_CHECK(childId);
}
InteractiveWait("Finishing ValidateResults");
}
void TestBed::DestroyComms()
{
InteractiveWait("Starting DestroyComms");
int const cmd = TestBedChild::CHILD_DESTROY_COMMS;
for (int childId = 0; childId < this->numActiveChildren; ++childId)
{
@@ -353,10 +379,13 @@ namespace RcclUnitTesting
this->numActiveChildren = 0;
this->numActiveRanks = 0;
this->numCollectivesInGroup = 0;
InteractiveWait("Finishing DestroyComms");
}
void TestBed::Finalize()
{
InteractiveWait("Starting Finalize");
// Send Stop to all child processes
int const cmd = TestBedChild::CHILD_STOP;
for (int childId = 0; childId < this->numDevicesAvailable; ++childId)
@@ -368,6 +397,7 @@ namespace RcclUnitTesting
close(childList[childId]->parentReadFd);
}
this->numDevicesAvailable = 0;
InteractiveWait("Finishing Finalize");
}
TestBed::~TestBed()
@@ -562,6 +592,16 @@ namespace RcclUnitTesting
}
}
void TestBed::InteractiveWait(std::string message)
{
if (ev.useInteractive)
{
INFO("%s\n", message.c_str());
INFO("<Hit any key to continue>\n");
scanf("%*c");
}
}
int& TestBed::NumTestsRun()
{
static int numTestsRun = 0;
+3
Просмотреть файл
@@ -130,6 +130,9 @@ namespace RcclUnitTesting
std::vector<bool> const& managedMemList,
std::vector<bool> const& useHipGraphList);
// Wait for user-input if in interactive mode
void InteractiveWait(std::string message);
// Used to track total number of calls to ExecuteCollectives()
static int& NumTestsRun();
+5 -1
Просмотреть файл
@@ -431,6 +431,7 @@ namespace RcclUnitTesting
{
for (int localRank : localRanksToExecute)
{
CHECK_HIP(hipSetDevice(this->deviceIds[localRank]));
if (this->verbose) INFO("Capturing stream for rank %d\n", localRank);
for (int i = 0; i < this->numStreamsPerGroup; i++)
{
@@ -601,6 +602,7 @@ namespace RcclUnitTesting
{
for (int localRank = 0; localRank < this->comms.size(); ++localRank)
{
CHECK_HIP(hipSetDevice(this->deviceIds[localRank]));
CHILD_NCCL_CALL_NON_BLOCKING("ncclCommGetAsyncErrorGroupEnd", localRank);
}
}
@@ -617,7 +619,7 @@ namespace RcclUnitTesting
for (int localRank : localRanksToExecute)
{
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++)
{
CHECK_HIP(hipStreamEndCapture(this->streams[localRank][i], &graphs[localRank][i]));
@@ -641,6 +643,7 @@ namespace RcclUnitTesting
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]));
@@ -668,6 +671,7 @@ namespace RcclUnitTesting
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]));