Adding timeout functionality/EnvVar to TestBed (#1044)

* Adding timeout functionality/EnvVar to TestBed
* updating timeout unit to microseconds

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

[ROCm/rccl commit: 9c0ef11ac7]
This commit is contained in:
Tim
2024-01-17 11:33:01 -05:00
committed by GitHub
parent 1d62a5f440
commit 245e757b26
6 changed files with 53 additions and 5 deletions
+45 -2
View File
@@ -393,6 +393,9 @@ namespace RcclUnitTesting
ErrCode TestBedChild::ExecuteCollectives()
{
int timeoutUs = 0;
PIPE_READ(timeoutUs);
bool useHipGraph = false;
PIPE_READ(useHipGraph);
@@ -432,7 +435,6 @@ namespace RcclUnitTesting
{
for (int localRank : localRanksToExecute)
{
CHECK_HIP(hipSetDevice(this->deviceIds[localRank]));
if (this->verbose) INFO("Capturing stream for rank %d\n", localRank);
CHECK_HIP(hipSetDevice(this->deviceIds[localRank]));
for (int i = 0; i < this->numStreamsPerGroup; i++)
@@ -659,10 +661,47 @@ namespace RcclUnitTesting
}
// Synchronize
std::vector<hipStream_t> streamsToComplete;
for (int localRank : localRanksToExecute)
{
for (int i = 0; i < this->numStreamsPerGroup; i++)
streamsToComplete.push_back(this->streams[localRank][i]);
}
int usElapsed = 0;
using namespace std::chrono;
using Clock = std::chrono::high_resolution_clock;
if (this->verbose) INFO("Starting sychronization and timing\n");
const auto start = Clock::now();
while (!streamsToComplete.empty() && usElapsed < timeoutUs)
{
for (int i = 0; i < streamsToComplete.size(); i++)
{
if (hipStreamQuery(streamsToComplete[i]) == hipSuccess)
{
streamsToComplete.erase(streamsToComplete.begin() + i);
i--;
}
}
usElapsed = duration_cast<microseconds>(Clock::now() - start).count();
}
// timed out
if (!streamsToComplete.empty())
{
if (this->verbose) INFO("Collective timed out, aborting\n");
for (int localRank : localRanksToExecute)
{
ncclCommAbort(this->comms[localRank]);
timeoutUs = -1;
}
}
// extra sync to flush GPU cache for validation later
// TODO: remove this after figuring out & fixing the exact behavior
// of fencing between kernels and at hipStreamQuery
for (int localRank : localRanksToExecute)
{
if (this->verbose) INFO("Starting synchronization for rank %d\n", localRank);
CHECK_HIP(hipSetDevice(this->deviceIds[localRank]));
for (int i = 0; i < this->numStreamsPerGroup; i++)
CHECK_HIP(hipStreamSynchronize(this->streams[localRank][i]));
}
@@ -699,6 +738,10 @@ namespace RcclUnitTesting
collArg.expected.ToString(collArg.dataType, numOutputElementsToPrint).c_str());
}
}
if (timeoutUs == -1)
return TEST_TIMEOUT;
if (this->verbose) INFO("Child %d finishes ExecuteCollectives()\n", this->childId);
return TEST_SUCCESS;
}