Move MSCCL algorithm loading to initialization to workaround HIP graph conflict (#982)

* MSCCL: pre-specify channels and pre-load algorithms

* add mutex

* fix bug

* clean include

* disable all-gathers temporarily
Tento commit je obsažen v:
Ziyue Yang
2023-12-01 01:47:20 +08:00
odevzdal GitHub
rodič 20b02af19b
revize 4bb0b4a380
12 změnil soubory, kde provedl 85 přidání a 4504 odebrání
+64 -19
Zobrazit soubor
@@ -26,7 +26,6 @@ RCCL_PARAM(MscclEnabled, "MSCCL_ENABLE", 1);
RCCL_PARAM(MscclForceEnabled, "MSCCL_FORCE_ENABLE", 0);
static const char* mscclAlgoFilePathEnv = "MSCCL_ALGO_FILE_PATH";
static std::atomic<bool> mscclInitialized;
static bool mscclSchedulerTriedLoadAlgo = false;
static std::mutex mscclLifecycleMutex;
bool mscclEnabled() {
@@ -78,8 +77,21 @@ static const char* mscclUnitTestAlgoDefaultDir = "msccl-unit-test-algorithms";
static const char* mscclAlgoShareDirPath = "../share/rccl/msccl-algorithms";
static const char* mscclUnitTestAlgoShareDirPath = "../share/rccl/msccl-unit-test-algorithms";
static ncclResult_t mscclInternalSchedulerInit() {
static ncclResult_t mscclInternalSchedulerInit(ncclComm_t comm, int* numChannelsRequired) {
static bool mscclAlgoMetaLoaded = false;
mscclStatus& status = mscclGetStatus();
*numChannelsRequired = 0;
// Query numChannelsRequired from loaded algorithm metas
if (mscclAlgoMetaLoaded) {
for (auto& m : status.algoMetas) {
if (comm->nRanks == m.nRanks) {
*numChannelsRequired = std::max(*numChannelsRequired, m.nChannels);
}
}
return ncclSuccess;
}
const char* mscclAlgoDir = getenv(mscclAlgoDirEnv);
const char* mscclAlgoShareDir = nullptr;
std::string mscclAlgoDirStr;
@@ -117,6 +129,7 @@ static ncclResult_t mscclInternalSchedulerInit() {
fullDirPath = mscclAlgoDir;
}
INFO(NCCL_INIT, "Using MSCCL files from %s", fullDirPath);
while ((entry = readdir(dp))) {
if (entry->d_type != DT_LNK && entry->d_type != DT_REG) {
continue;
@@ -126,16 +139,28 @@ static ncclResult_t mscclInternalSchedulerInit() {
fullPath += "/";
fullPath += entry->d_name;
NCCLCHECK(mscclGetAlgoMetaFromXmlFile(fullPath.c_str(), &(status.algoMetas.back())));
if (status.algoMetas.back().nRanks == comm->nRanks) {
*numChannelsRequired = std::max(*numChannelsRequired, status.algoMetas.back().nChannels);
}
}
if (closedir(dp)) {
WARN("MSCCL Internal Scheduler: closedir failed, error %d", errno);
return ncclInvalidUsage;
}
status.rankToAlgoHandles.resize(status.algoMetas.size());
mscclAlgoMetaLoaded = true;
return ncclSuccess;
}
static ncclResult_t mscclSchedulerInit() {
ncclResult_t mscclSchedulerInit(ncclComm_t comm, int* numChannelsRequired) {
*numChannelsRequired = 0;
comm->mscclCompatible = mscclCommCompatible(comm);
if (!comm->mscclCompatible) {
return ncclSuccess;
}
std::lock_guard<std::mutex> lock(mscclLifecycleMutex);
mscclStatus& status = mscclGetStatus();
bool useInternalScheduler = false;
@@ -155,11 +180,14 @@ static ncclResult_t mscclSchedulerInit() {
useInternalScheduler = true;
}
}
if (useInternalScheduler) {
NCCLCHECK(mscclInternalSchedulerInit());
NCCLCHECK(mscclInternalSchedulerInit(comm, numChannelsRequired));
} else {
NCCLCHECK(status.mscclSchedulerPtr->init());
*numChannelsRequired = MAXCHANNELS;
}
return ncclSuccess;
}
@@ -170,30 +198,53 @@ ncclResult_t mscclInit(ncclComm_t comm) {
threadLocalStatus.groupDepth = 0;
threadLocalStatus.captureId = ULLONG_MAX;
threadLocalStatus.captureStatus = mscclNoCapture;
comm->mscclCompatible = mscclCommCompatible(comm);
{
std::lock_guard<std::mutex> lock(mscclLifecycleMutex);
mscclStatus& status = mscclGetStatus();
// Free algorithm handles are initialized globally once and before algorithm pre-processing
if (!mscclInitialized.load(std::memory_order_acquire)) {
status.freeAlgoHandles.resize(MSCCL_MAX_NUM_ALGOS);
for (int i = 0; i < MSCCL_MAX_NUM_ALGOS; i++) {
status.freeAlgoHandles[i] = MSCCL_MAX_NUM_ALGOS - i - 1;
}
}
// Pre-process all algorithms for internal scheduler and for different comms.
// This is a temp fix to bypass the issue that stream cannot be synchronized during HIP graph capturing,
// should use dynamic loading approach after the issue is fixed.
if (comm->mscclCompatible && !status.mscclSchedulerPtr) {
for (size_t i = 0; i < status.algoMetas.size(); i++) {
auto &m = status.algoMetas[i];
mscclAlgoHandle_t mscclAlgoHandle;
if (m.nRanks == comm->nRanks) {
// Load algorithms
if (status.rankToAlgoHandles[i].find(comm->rank) == status.rankToAlgoHandles[i].end()) {
NCCLCHECK(mscclLoadAlgo(m.filePath.c_str(), &mscclAlgoHandle, comm->rank));
status.rankToAlgoHandles[i][comm->rank] = mscclAlgoHandle;
}
// Connect algorithms
if (status.connectedAlgos[comm].find(mscclAlgoHandle) == status.connectedAlgos[comm].end()) {
NCCLCHECK(mscclSetupConnections(status.hostAlgos[mscclAlgoHandle], comm));
status.connectedAlgos[comm].insert(mscclAlgoHandle);
}
}
}
}
if (mscclInitialized.load(std::memory_order_acquire)) {
return ncclSuccess;
}
mscclStatus& status = mscclGetStatus();
status.scratchBuffer = nullptr;
status.scratchBufferSize = 0;
status.workIndex = 1;
status.freeAlgoHandles.resize(MSCCL_MAX_NUM_ALGOS);
for (int i = 0; i < MSCCL_MAX_NUM_ALGOS; i++) {
status.freeAlgoHandles[i] = MSCCL_MAX_NUM_ALGOS - i - 1;
}
NCCLCHECK(ncclCudaCalloc(&status.syncFlags, MSCCL_MAX_NUM_THREAD_BLOCKS));
status.lastStream = nullptr;
status.needsProxy = false;
NCCLCHECK(mscclInitWorkFifoStatus(&(status.defaultWorkFifoStatus)));
mscclSchedulerTriedLoadAlgo = false;
NCCLCHECK(mscclSchedulerInit());
mscclInitialized.store(true, std::memory_order_release);
}
@@ -248,12 +299,6 @@ static ncclResult_t mscclInternalSchedulerSelectAlgo(struct mscclSchedulerParam*
m.nRanks == param->nRanks &&
m.func == param->func &&
(isInPlace ? m.inPlace : m.outOfPlace)) {
// If not loaded for current rank, load it
if (status.rankToAlgoHandles[i].find(param->rank) == status.rankToAlgoHandles[i].end()) {
mscclAlgoHandle_t algoHandle;
NCCLCHECK(mscclLoadAlgo(m.filePath.c_str(), &algoHandle, param->rank));
status.rankToAlgoHandles[i][param->rank] = algoHandle;
}
param->handle = status.rankToAlgoHandles[i][param->rank];
param->scheduled = true;
return ncclSuccess;
+4
Zobrazit soubor
@@ -731,6 +731,10 @@ ncclResult_t mscclGetAlgoMetaFromXmlFile(const char* str, struct mscclAlgoMeta*
NCCLCHECK(mscclXmlGetAttrInt(node, "nchunksperloop", &nChunksPerLoop));
algoMeta->nChunksPerLoop = nChunksPerLoop;
int nChannels;
NCCLCHECK(mscclXmlGetAttrInt(node, "nchannels", &nChannels));
algoMeta->nChannels = nChannels;
int nGpus;
NCCLCHECK(mscclXmlGetAttrInt(node, "ngpus", &nGpus));
algoMeta->nRanks = nGpus;
+2 -7
Zobrazit soubor
@@ -95,14 +95,9 @@ ncclResult_t mscclSetupConnections(struct mscclAlgo* hostAlgo, ncclComm_t comm)
mscclStatus& status = mscclGetStatus();
// Check whether there are enough channels
if (hostAlgo->nChannels > MAXCHANNELS) {
WARN("MSCCL: max number of channels available (%d) less than required (%d)", MAXCHANNELS, hostAlgo->nChannels);
return ncclInvalidUsage;
}
if (hostAlgo->nChannels > comm->nChannels) {
for (int channelId = comm->nChannels; channelId < hostAlgo->nChannels; channelId++) {
NCCLCHECK(initChannel(comm, channelId));
}
WARN("MSCCL: number of channels available (%d) less than required (%d)", comm->nChannels, hostAlgo->nChannels);
return ncclInvalidUsage;
}
// Flag MSCCL connections