replayer update v0 (#1733)

* First version of new replayer, with comments on future TODOs

* plus minor fixes for UT

* Updated format of recorder, especially in binary department, according to replayer's need

[ROCm/rccl commit: ba97c9c18b]
This commit is contained in:
Tim
2025-06-13 15:05:34 -04:00
committed by GitHub
parent 2c0cc20a76
commit 7051f217a7
8 changed files with 768 additions and 1129 deletions
+53 -25
View File
@@ -43,14 +43,21 @@ rcclApiCall::rcclApiCall(rcclCall_t type, const ncclInfo& info)://name(rcclCallS
nRanks(info.comm->nRanks),
stream(info.stream),
nTasks(info.comm->planner.nTasksP2p + info.comm->planner.nTasksColl),
globalRank(info.comm->localRankToRank[info.comm->localRank]){}
globalRank(info.comm->localRankToRank[info.comm->localRank])
{
hipMemGetAddressRange(&recvPtrBase, &recvPtrExtent, const_cast<void*>(info.recvbuff)); // should always exist for collectives
if (info.sendbuff) // ncclSend/Recv
{
hipMemGetAddressRange(&sendPtrBase, &sendPtrExtent, const_cast<void*>(info.sendbuff));
}
}
rcclApiCall::rcclApiCall(rcclCall_t type) : type(type){}
std::string siminfo_fmt = "[size : %zu, magic : %u, version : %u, estimated time : %f, timestamp : %f]";
std::string config_fmt = ", ncclConfig : [size : %zu, magic : %u, version : %u, blocking : %d, cgaClusterSize : %d, minCTA : %d, maxCTA : %d, netname : %s, splitshare : %d]";
std::string ctxt_fmt = "time : %lf, thread : %d, device : %d, captured : %d, graphID : %llu ]]"; // implicit context info
std::string ubr_fmt = "%s : [comm : %p, buff : %p, returned handle : %p, size : %zu, context : [";
std::string ubr_fmt = "%s : [comm : %p, buff : [addr : %p, base : %p, size : %zu], returned handle : %p, count : %zu, context : [";
std::string getId_fmt = "%s : [uniqueID : %llu, context : [";
std::string ubDereg_fmt = "%s : [comm : %p, handle : %p, context : [";
std::string rank_fmt = "%s : [size : %d, uniqueID : %llu, rank : %d, context : [";
@@ -62,7 +69,7 @@ std::string alloc_fmt = "%s : [returned ptr : %p, size : %zu, context : [";
std::string free_fmt = "%s : [ptr : %p, context : [";
std::string redop_fmt = "%s : [scalar : %p, datatype : %d, op : %d, residence : %d, comm : %p, context : [";
std::string redopdestroy_fmt = "%s : [op : %d, comm : %p, context : [";
std::string coll_fmt = "%s : [opCount : %lx, sendbuff : %p, recvbuff : %p, count : %zu, datatype : %d, op : %d, root : %d, comm : %p, nranks : %d, stream : %p, task : %d, globalrank : %d, context : [";
std::string coll_fmt = "%s : [opCount : %lx, sendbuff : [addr : %p, base : %p, size : %zu], recvbuff : [addr : %p, base : %p, size : %zu], count : %zu, datatype : %d, op : %d, root : %d, comm : %p, nranks : %d, stream : %p, task : %d, globalrank : %d, context : [";
Recorder::Recorder()
{
@@ -74,8 +81,8 @@ Recorder::Recorder()
}
logLevel = getenv("RCCL_LOG_LEVEL") ? std::stoi(getenv("RCCL_LOG_LEVEL")) : 1;
char hostname[1024];
getHostName(hostname, 1024, '.');
char hostname[256];
gethostname(hostname, 256);
pid = getpid();
output_json = 0;
@@ -94,13 +101,13 @@ Recorder::Recorder()
output_name = std::string(filename);
}
outputFile.open(output_name + std::to_string(pid) + output_extension,
outputFile.open(output_name + "." + std::to_string(pid) + "." + std::string(hostname) + output_extension,
output_json ? std::ofstream::out : std::ofstream::binary);
if (output_json)
{
outputFile << "{" << std::endl;
indent(2, outputFile);
outputFile << "hostname : " << hostname << ", version : 0,";
outputFile << "version : 1,";
}
}
@@ -110,6 +117,14 @@ Recorder& Recorder::instance()
return _instance;
}
void Recorder::skip(bool b)
{
if (filename.size())
{
skipped = b;
}
}
void Recorder::captureGpuContext(rcclApiCall& call) const
{
call.timestamp = duration_cast<duration<double>>(high_resolution_clock::now().time_since_epoch()).count() * 1000;
@@ -171,7 +186,7 @@ void Recorder::write(const rcclApiCall &call)
case rrCommRegister:
{
len = snprintf(buffer, 4096, ubr_fmt.c_str(),
rcclCallStr[call.type], call.comm, call.sendbuff, call.recvbuff, call.count);
rcclCallStr[call.type], call.comm, call.sendbuff, call.sendPtrBase, call.sendPtrExtent, call.recvbuff, call.count);
break;
}
case rrCommDeregister:
@@ -240,7 +255,8 @@ void Recorder::write(const rcclApiCall &call)
}
default: // collectives
len = snprintf(buffer, 4096, coll_fmt.c_str(),
rcclCallStr[call.type], call.opCount, call.sendbuff, call.recvbuff, call.count, call.datatype,
rcclCallStr[call.type], call.opCount, call.sendbuff, call.sendPtrBase, call.sendPtrExtent,
call.recvbuff, call.recvPtrBase, call.recvPtrExtent, call.count, call.datatype,
call.op, call.root, call.comm, call.nRanks, call.stream, call.nTasks, call.globalRank);
}
@@ -249,7 +265,6 @@ void Recorder::write(const rcclApiCall &call)
outputFile.write(buffer, len);
} else {
outputFile.write((char*)&call, sizeof(rcclApiCall));
outputFile << std::endl;
}
outputFile.flush();
return ;
@@ -314,7 +329,6 @@ ncclResult_t Recorder::record(rcclApiCall& call)
#if ROCM_VERSION >= 60100
hipStreamCaptureStatus status;
hipGraph_t graphCaptured;
unsigned long long graphID = 0;
CUDACHECK(hipStreamGetCaptureInfo_v2(call.stream, &status, &(call.graphID), &graphCaptured)); // shouldnt we need dependency?
if (status == hipStreamCaptureStatusActive) // when graph launched this should be disabled
@@ -338,7 +352,7 @@ ncclResult_t Recorder::record(rcclApiCall& call)
ncclResult_t Recorder::record(rcclCall_t type, const ncclInfo& info)
{
if (!filename.size())
if (!filename.size() || skipped)
{
return ncclSuccess;
}
@@ -366,9 +380,9 @@ ncclResult_t Recorder::record(rcclCall_t type, const void* sendbuff, void* recvb
ncclResult_t ret = record(call);
if (type == rrAllToAllv)
{
int size = call.nRanks - 1;
if (output_json)
{
int size = call.nRanks - 1;
outputFile << ", sendcounts : [";
for (int i = 0; i < size; i++) outputFile << sendcounts[i] << ", ";
outputFile << sendcounts[size] << "], sdispls : [";
@@ -378,9 +392,13 @@ ncclResult_t Recorder::record(rcclCall_t type, const void* sendbuff, void* recvb
outputFile << recvcounts[size] << "], rdispls : [";
for (int i = 0; i < size; i++) outputFile << rdispls[i] << ", ";
outputFile << rdispls[size] << "]";
outputFile.flush();
} else {
outputFile.write((char*)sendcounts, sizeof(size_t) * (size + 1));
outputFile.write((char*)sdispls, sizeof(size_t) * (size + 1));
outputFile.write((char*)recvcounts, sizeof(size_t) * (size + 1));
outputFile.write((char*)rdispls, sizeof(size_t) * (size + 1));
}
// else export to binary
outputFile.flush();
}
return ret;
}
@@ -406,7 +424,7 @@ ncclResult_t Recorder::record(rcclCall_t type, ncclRedOp_t op, ncclComm_t comm,
ncclResult_t Recorder::record(rcclCall_t type, int groupDepth)
{
if (!filename.size())
if (!filename.size() || skipped)
{
return ncclSuccess;
}
@@ -467,6 +485,7 @@ ncclResult_t Recorder::record(rcclCall_t type, ncclComm_t comm, void* handle, vo
call.recvbuff = handle;
if (type == rrCommRegister)
{
CUDACHECK(hipMemGetAddressRange(&call.sendPtrBase, &call.sendPtrExtent, userBuffer));
call.sendbuff = userBuffer;
call.count = size;
}
@@ -503,6 +522,7 @@ void Recorder::record(int groupDepth, ncclSimInfo_t *siminfo)
siminfo->size, siminfo->magic, siminfo->version, siminfo->estimatedTime, call.timestamp);
outputFile.write(buffer, len);
} // no tid for groupCall
// TODO: else flush siminfo in binary
outputFile.flush();
}
@@ -529,6 +549,7 @@ void Recorder::record(rcclCall_t type, int size, int rank, ncclUniqueId* commId,
outputFile.write(buffer, len);
outputFile.flush();
}
// TODO: else flush ncclConfig in binary
}
void Recorder::record(ncclComm_t* comms, int ndev, const int* devlist)
@@ -540,15 +561,20 @@ void Recorder::record(ncclComm_t* comms, int ndev, const int* devlist)
rcclApiCall call(rrCommInitAll);
call.root = ndev;
// call.sendbuff = comms; TODO: might log this too
call.sendbuff = devlist;
record(call);
if (output_json && devlist)
if (devlist)
{
outputFile << ", devlist : [";
for (int i = 0; i < call.root - 1; i++)
outputFile << devlist[i] << ", ";
outputFile << devlist[call.root - 1] << "]";
if (output_json)
{
outputFile << ", devlist : [";
for (int i = 0; i < call.root - 1; i++)
outputFile << devlist[i] << ", ";
outputFile << devlist[call.root - 1] << "]";
} else {
outputFile.write((char*)devlist, sizeof(int) * ndev);
}
outputFile.flush();
}
}
@@ -578,6 +604,7 @@ static rcclCall_t getFuncType(std::string func)
void parseJsonEntry(const char* entry, std::vector<rcclApiCall>& calls)
{
// TODO: parse comma too
rcclApiCall call;
std::string str(entry);
size_t begin = str.find_first_not_of(' ');
@@ -588,7 +615,7 @@ void parseJsonEntry(const char* entry, std::vector<rcclApiCall>& calls)
case rrCommRegister:
{
assert(sscanf(str.c_str() + end + 3, (ubr_fmt.substr(5) + ctxt_fmt).c_str(),
&call.comm, &call.sendbuff, &call.recvbuff, &call.count) == 4);
&call.comm, &call.sendbuff, &call.sendPtrBase, &call.sendPtrExtent, &call.recvbuff, &call.count) == 6);
break;
}
case rrCommDeregister:
@@ -658,9 +685,10 @@ void parseJsonEntry(const char* entry, std::vector<rcclApiCall>& calls)
}
default:
assert(sscanf(str.c_str() + end + 3, (coll_fmt.substr(5) + ctxt_fmt).c_str(),
&call.opCount, &call.sendbuff, &call.recvbuff, &call.count, &call.datatype, &call.op, &call.root,
&call.opCount, &call.sendbuff, &call.sendPtrBase, &call.sendPtrExtent, &call.recvbuff, &call.recvPtrBase, &call.recvPtrExtent,
&call.count, &call.datatype, &call.op, &call.root,
&call.comm, &call.nRanks, &call.stream, &call.nTasks, &call.globalRank, &call.timestamp, &call.tid,
&call.hipDev, &call.graphCaptured, &call.graphID) == 17);
&call.hipDev, &call.graphCaptured, &call.graphID) == 21);
}
calls.push_back(call);
}