Enable multi-threading for MSCCL (#1203)
MSCCL can now run in a multi-threaded configuration. To test in the unit tests, added the ENABLE_OPENMP compile definition flag and the --openmp-test-enable flag to the unit test build script. To activate, set the environment variables UT_MULTITHREADED=1 and UT_PROCESS_MASK=1. Set Jenkins to use this mode.
[ROCm/rccl commit: 0c36d571ea]
Tá an tiomantas seo le fáil i:
tiomanta ag
GitHub
tuismitheoir
b5bc883f61
tiomantas
37bf54b8f8
@@ -25,8 +25,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 std::mutex mscclLifecycleMutex;
|
||||
|
||||
bool mscclEnabled() {
|
||||
#ifdef COMPILE_MSCCL_KERNEL
|
||||
@@ -56,23 +54,12 @@ bool mscclIsCaller() {
|
||||
return mscclGetThreadLocalStatus().mscclIsCallerFlag;
|
||||
}
|
||||
|
||||
bool mscclAvailable() {
|
||||
return mscclEnabled() && mscclInitialized.load(std::memory_order_acquire);
|
||||
bool mscclAvailable(int rank) {
|
||||
return mscclEnabled() && mscclInitialized(rank);
|
||||
}
|
||||
|
||||
static bool mscclCommCompatible(ncclComm_t comm) {
|
||||
std::map<uint64_t, std::set<uint64_t>> hostHashToPidHashes;
|
||||
for (int i = 0; i < comm->nRanks; i++) {
|
||||
uint64_t hostHash = comm->peerInfo[i].hostHash;
|
||||
uint64_t pidHash = comm->peerInfo[i].pidHash;
|
||||
if (hostHashToPidHashes.find(hostHash) != hostHashToPidHashes.end()) {
|
||||
auto& pidHashSet = hostHashToPidHashes[hostHash];
|
||||
if (pidHashSet.find(pidHash) != pidHashSet.end()) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
hostHashToPidHashes[hostHash].insert(pidHash);
|
||||
}
|
||||
// MSCCL is always compatible now. No need to guard against multi-thread.
|
||||
return true;
|
||||
}
|
||||
|
||||
@@ -86,8 +73,8 @@ static const char* mscclAlgoShareDirPath = "../share/rccl/msccl-algorithms";
|
||||
static const char* mscclUnitTestAlgoShareDirPath = "../share/rccl/msccl-unit-test-algorithms";
|
||||
|
||||
static ncclResult_t mscclInternalSchedulerInit(ncclComm_t comm, int* numChannelsRequired) {
|
||||
static bool mscclAlgoMetaLoaded = false;
|
||||
mscclStatus& status = mscclGetStatus();
|
||||
static thread_local bool mscclAlgoMetaLoaded = false;
|
||||
mscclStatus& status = mscclGetStatus(comm->rank);
|
||||
|
||||
*numChannelsRequired = 0;
|
||||
// Query numChannelsRequired from loaded algorithm metas
|
||||
@@ -166,9 +153,7 @@ ncclResult_t mscclSchedulerInit(ncclComm_t comm, int* numChannelsRequired) {
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
std::lock_guard<std::mutex> lock(mscclLifecycleMutex);
|
||||
|
||||
mscclStatus& status = mscclGetStatus();
|
||||
mscclStatus& status = mscclGetStatus(comm->rank);
|
||||
bool useInternalScheduler = false;
|
||||
|
||||
const char* mscclSchedulerPath = getenv(mscclSchedulerPathEnv);
|
||||
@@ -199,20 +184,11 @@ ncclResult_t mscclSchedulerInit(ncclComm_t comm, int* numChannelsRequired) {
|
||||
}
|
||||
|
||||
ncclResult_t mscclInit(ncclComm_t comm) {
|
||||
// Always initialize thread local status
|
||||
mscclThreadLocalStatus threadLocalStatus = mscclGetThreadLocalStatus();
|
||||
threadLocalStatus.groupStatus = mscclNoGroup;
|
||||
threadLocalStatus.groupDepth = 0;
|
||||
threadLocalStatus.captureId = ULLONG_MAX;
|
||||
threadLocalStatus.captureStatus = mscclNoCapture;
|
||||
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(mscclLifecycleMutex);
|
||||
|
||||
mscclStatus& status = mscclGetStatus();
|
||||
mscclStatus& status = mscclGetStatus(comm->rank);
|
||||
|
||||
// freeAlgoHandles and needsProxy are initialized globally once and before algorithm pre-processing and connection
|
||||
if (!mscclInitialized.load(std::memory_order_acquire)) {
|
||||
if (!mscclInitialized(comm->rank)) {
|
||||
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;
|
||||
@@ -229,6 +205,8 @@ ncclResult_t mscclInit(ncclComm_t comm) {
|
||||
if (m.nRanks == comm->nRanks) {
|
||||
// Load algorithms
|
||||
if (status.rankToAlgoHandles[i].find(comm->rank) == status.rankToAlgoHandles[i].end()) {
|
||||
static std::mutex loadAlgoMutex;
|
||||
std::lock_guard<std::mutex> lock(loadAlgoMutex);
|
||||
NCCLCHECK(mscclLoadAlgo(m.filePath.c_str(), &(status.rankToAlgoHandles[i][comm->rank]), comm->rank));
|
||||
}
|
||||
// Connect algorithms
|
||||
@@ -241,7 +219,7 @@ ncclResult_t mscclInit(ncclComm_t comm) {
|
||||
}
|
||||
}
|
||||
|
||||
if (mscclInitialized.load(std::memory_order_acquire)) {
|
||||
if (mscclInitialized(comm->rank)) {
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
@@ -250,7 +228,7 @@ ncclResult_t mscclInit(ncclComm_t comm) {
|
||||
status.lastStream = nullptr;
|
||||
NCCLCHECK(mscclInitWorkFifoStatus(&(status.defaultWorkFifoStatus)));
|
||||
|
||||
mscclInitialized.store(true, std::memory_order_release);
|
||||
mscclSetInitialized(comm->rank);
|
||||
}
|
||||
|
||||
INFO(NCCL_INIT, "MSCCL: Initialization finished, localSize %ld", mscclKernMaxLocalSize());
|
||||
@@ -266,8 +244,8 @@ ncclResult_t mscclGroupStart() {
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
static ncclResult_t mscclInternalSchedulerSelectAlgo(struct mscclSchedulerParam* param) {
|
||||
mscclStatus& status = mscclGetStatus();
|
||||
static ncclResult_t mscclInternalSchedulerSelectAlgo(int rank, struct mscclSchedulerParam* param) {
|
||||
mscclStatus& status = mscclGetStatus(rank);
|
||||
param->scheduled = false;
|
||||
|
||||
// Current MSCCL doesn't support pre/post op
|
||||
@@ -312,13 +290,13 @@ static ncclResult_t mscclInternalSchedulerSelectAlgo(struct mscclSchedulerParam*
|
||||
}
|
||||
|
||||
static ncclResult_t mscclSchedulerSelectAlgo(struct mscclSavedSchedulerParam* param) {
|
||||
mscclStatus& status = mscclGetStatus();
|
||||
mscclStatus& status = mscclGetStatus(param->comm->rank);
|
||||
if (status.mscclSchedulerPtr) {
|
||||
NCCLCHECK(status.mscclSchedulerPtr->selectAlgo(&(param->p)));
|
||||
} else {
|
||||
// Disable MSCCL algorithms if machine type is not matching
|
||||
if (param->comm->topo->mscclEnabled || mscclForceEnabled()) {
|
||||
NCCLCHECK(mscclInternalSchedulerSelectAlgo(&(param->p)));
|
||||
NCCLCHECK(mscclInternalSchedulerSelectAlgo(param->comm->rank, &(param->p)));
|
||||
} else {
|
||||
param->p.scheduled = false;
|
||||
}
|
||||
@@ -513,12 +491,30 @@ ncclResult_t mscclGroupEnd() {
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
static ncclResult_t mscclInternalSchedulerTeardown() {
|
||||
static ncclResult_t mscclInternalUnloadAlgo(int rank, mscclAlgoHandle_t mscclAlgoHandle) {
|
||||
mscclStatus& status = mscclGetStatus(rank);
|
||||
|
||||
free(status.hostAlgos[mscclAlgoHandle]);
|
||||
status.hostAlgos.erase(mscclAlgoHandle);
|
||||
|
||||
NCCLCHECK(ncclCudaFree(status.devAlgos[mscclAlgoHandle]));
|
||||
status.devAlgos.erase(mscclAlgoHandle);
|
||||
|
||||
status.freeAlgoHandles.push_back(mscclAlgoHandle);
|
||||
|
||||
for (auto &s : status.connectedAlgos) {
|
||||
s.second.erase(mscclAlgoHandle);
|
||||
}
|
||||
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
static ncclResult_t mscclInternalSchedulerTeardown(int rank) {
|
||||
ncclResult_t ret = ncclSuccess, tmpRet = ncclSuccess;
|
||||
mscclStatus& status = mscclGetStatus();
|
||||
mscclStatus& status = mscclGetStatus(rank);
|
||||
for (auto &m : status.rankToAlgoHandles) {
|
||||
for (auto &p : m) {
|
||||
tmpRet = mscclUnloadAlgo(p.second);
|
||||
tmpRet = mscclInternalUnloadAlgo(rank, p.second);
|
||||
if (ret == ncclSuccess) {
|
||||
ret = tmpRet;
|
||||
}
|
||||
@@ -529,18 +525,13 @@ static ncclResult_t mscclInternalSchedulerTeardown() {
|
||||
return ret;
|
||||
}
|
||||
|
||||
ncclResult_t mscclTeardown() {
|
||||
// Always teardown thread local status
|
||||
mscclThreadLocalStatus threadLocalStatus = mscclGetThreadLocalStatus();
|
||||
threadLocalStatus.savedSchedulerParams.clear();
|
||||
|
||||
ncclResult_t mscclTeardown(int rank) {
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(mscclLifecycleMutex);
|
||||
|
||||
if (!mscclInitialized.load(std::memory_order_acquire)) {
|
||||
if (!mscclInitialized(rank)) {
|
||||
mscclRemoveRank(rank);
|
||||
return ncclSuccess;
|
||||
}
|
||||
mscclStatus& status = mscclGetStatus();
|
||||
mscclStatus& status = mscclGetStatus(rank);
|
||||
for (auto &p : status.hostAlgos) {
|
||||
free(p.second);
|
||||
status.freeAlgoHandles.push_back(p.first);
|
||||
@@ -563,13 +554,14 @@ ncclResult_t mscclTeardown() {
|
||||
dlclose(status.mscclSchedulerLib);
|
||||
status.mscclSchedulerLib = nullptr;
|
||||
} else {
|
||||
NCCLCHECK(mscclInternalSchedulerTeardown());
|
||||
NCCLCHECK(mscclInternalSchedulerTeardown(rank));
|
||||
}
|
||||
NCCLCHECK(mscclDestroyWorkFifoStatus(&(status.defaultWorkFifoStatus)));
|
||||
for (auto &p : status.graphWorkFifoStatus) {
|
||||
NCCLCHECK(mscclDestroyWorkFifoStatus(&(p.second)));
|
||||
}
|
||||
mscclInitialized.store(false, std::memory_order_release);
|
||||
mscclSetInitialized(rank, false);
|
||||
mscclRemoveRank(rank);
|
||||
}
|
||||
|
||||
INFO(NCCL_INIT, "MSCCL: Teardown finished");
|
||||
|
||||
@@ -22,10 +22,10 @@ static inline size_t computeSizeNeeded(size_t nBytes, int nScratchChunks, int nC
|
||||
return (nBytes * (size_t)nScratchChunks) / (size_t)nChunksPerLoop;
|
||||
}
|
||||
|
||||
ncclResult_t mscclGetCaptureStatus(hipStream_t stream) {
|
||||
mscclStatus& status = mscclGetStatus();
|
||||
ncclResult_t mscclGetCaptureStatus(int rank, hipStream_t stream) {
|
||||
mscclStatus& status = mscclGetStatus(rank);
|
||||
mscclThreadLocalStatus& threadLocalStatus = mscclGetThreadLocalStatus();
|
||||
mscclSavedProxyArgs& savedProxyArgs = mscclGetSavedProxyArgs();
|
||||
mscclSavedProxyArgs& savedProxyArgs = mscclGetSavedProxyArgs(rank);
|
||||
cudaStreamCaptureStatus captureStatus;
|
||||
unsigned long long captureId;
|
||||
CUDACHECK(hipStreamGetCaptureInfo_v2(stream, &captureStatus, &captureId, &threadLocalStatus.graph, nullptr, nullptr));
|
||||
@@ -42,12 +42,12 @@ ncclResult_t mscclGetCaptureStatus(hipStream_t stream) {
|
||||
} else {
|
||||
threadLocalStatus.captureStatus = mscclNoCapture;
|
||||
}
|
||||
INFO(NCCL_NET,"mscclGetCaptureStatus: %d, captureId: %llu, size: %lu\n", threadLocalStatus.captureStatus, threadLocalStatus.captureId, mscclGetSavedProxyArgs()[captureId].size());
|
||||
INFO(NCCL_NET,"mscclGetCaptureStatus: %d, captureId: %llu, size: %lu\n", threadLocalStatus.captureStatus, threadLocalStatus.captureId, savedProxyArgs[captureId].size());
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
ncclResult_t mscclSetupCount(struct mscclAlgo* hostAlgo, ncclComm_t comm, size_t count, ncclDataType_t dataType) {
|
||||
mscclStatus& status = mscclGetStatus();
|
||||
mscclStatus& status = mscclGetStatus(comm->rank);
|
||||
status.stepSize = comm->buffSizes[hostAlgo->protocol] / NCCL_STEPS;
|
||||
status.chunkSteps = hostAlgo->protocol == NCCL_PROTO_SIMPLE ? hostAlgo->chunkSteps : 1;
|
||||
status.sliceSteps = hostAlgo->protocol == NCCL_PROTO_SIMPLE ? hostAlgo->sliceSteps : 1;
|
||||
@@ -69,12 +69,11 @@ ncclResult_t mscclSetupCount(struct mscclAlgo* hostAlgo, ncclComm_t comm, size_t
|
||||
}
|
||||
|
||||
ncclResult_t mscclSetupScratch(struct mscclAlgo* hostAlgo, hipStream_t stream) {
|
||||
mscclStatus& status = mscclGetStatus();
|
||||
return ncclSuccess;
|
||||
}
|
||||
|
||||
ncclResult_t mscclSetupSyncFlags(hipStream_t stream) {
|
||||
mscclStatus& status = mscclGetStatus();
|
||||
ncclResult_t mscclSetupSyncFlags(int rank, hipStream_t stream) {
|
||||
mscclStatus& status = mscclGetStatus(rank);
|
||||
mscclThreadLocalStatus& threadLocalStatus = mscclGetThreadLocalStatus();
|
||||
if (threadLocalStatus.captureStatus == mscclNewCapture ||
|
||||
status.workIndex > (1ULL << (8*sizeof(status.workIndex))) - 2 * NCCL_MAX_OPS - 1) {
|
||||
@@ -86,7 +85,7 @@ ncclResult_t mscclSetupSyncFlags(hipStream_t stream) {
|
||||
}
|
||||
|
||||
ncclResult_t mscclSetupConnections(struct mscclAlgo* hostAlgo, ncclComm_t comm) {
|
||||
mscclStatus& status = mscclGetStatus();
|
||||
mscclStatus& status = mscclGetStatus(comm->rank);
|
||||
|
||||
// Check whether there are enough channels
|
||||
if (hostAlgo->nChannels > comm->nChannels) {
|
||||
@@ -124,7 +123,7 @@ ncclResult_t mscclSetupConnections(struct mscclAlgo* hostAlgo, ncclComm_t comm)
|
||||
}
|
||||
|
||||
static ncclResult_t mscclSetupProxyImpl(struct mscclAlgo* hostAlgo, ncclComm_t comm) {
|
||||
mscclStatus& status = mscclGetStatus();
|
||||
mscclStatus& status = mscclGetStatus(comm->rank);
|
||||
mscclThreadLocalStatus& threadLocalStatus = mscclGetThreadLocalStatus();
|
||||
struct ncclProxyOp proxyOp = {};
|
||||
proxyOp.connIndex = 0;
|
||||
@@ -185,9 +184,9 @@ static void HIPRT_CB mscclSetupProxyCallback(void *args) {
|
||||
}
|
||||
|
||||
ncclResult_t mscclSetupProxy(struct mscclAlgo* hostAlgo, ncclComm_t comm, hipStream_t stream) {
|
||||
mscclStatus& status = mscclGetStatus();
|
||||
mscclStatus& status = mscclGetStatus(comm->rank);
|
||||
mscclThreadLocalStatus& threadLocalStatus = mscclGetThreadLocalStatus();
|
||||
mscclSavedProxyArgs& savedProxyArgs = mscclGetSavedProxyArgs();
|
||||
mscclSavedProxyArgs& savedProxyArgs = mscclGetSavedProxyArgs(comm->rank);
|
||||
if (threadLocalStatus.captureStatus == mscclNoCapture) {
|
||||
INFO(NCCL_NET,"mscclSetupProxy: no capture\n");
|
||||
NCCLCHECK(mscclSetupProxyImpl(hostAlgo, comm));
|
||||
@@ -203,7 +202,7 @@ ncclResult_t mscclSetupProxy(struct mscclAlgo* hostAlgo, ncclComm_t comm, hipStr
|
||||
p.userData = params;
|
||||
CUDACHECK(hipGraphAddHostNode(&callbackNode, threadLocalStatus.graph, nullptr, 0, &p));
|
||||
}
|
||||
mscclGetSavedProxyArgs()[threadLocalStatus.captureId].emplace_back(hostAlgo, comm);
|
||||
savedProxyArgs[threadLocalStatus.captureId].emplace_back(hostAlgo, comm);
|
||||
}
|
||||
return ncclSuccess;
|
||||
}
|
||||
@@ -410,7 +409,7 @@ RCCL_PARAM(MscclForceFullOps, "MSCCL_FORCE_FULLOPS", 0);
|
||||
ncclResult_t mscclSetupKernel(const void* sendBuff, void* recvBuff, size_t count,
|
||||
ncclDataType_t dataType, ncclRedOp_t op, struct mscclAlgo* hostAlgo, struct mscclAlgo* devAlgo,
|
||||
ncclComm_t comm, hipStream_t stream) {
|
||||
mscclStatus& status = mscclGetStatus();
|
||||
mscclStatus& status = mscclGetStatus(comm->rank);
|
||||
mscclThreadLocalStatus& threadLocalStatus = mscclGetThreadLocalStatus();
|
||||
|
||||
if (status.lastStream != stream && status.lastStream != nullptr) {
|
||||
|
||||
@@ -6,9 +6,60 @@
|
||||
#include "msccl/msccl_status.h"
|
||||
#include "msccl/msccl_struct.h"
|
||||
|
||||
mscclStatus& mscclGetStatus() {
|
||||
static mscclStatus status;
|
||||
return status;
|
||||
#include "debug.h"
|
||||
|
||||
#include <memory>
|
||||
#include <mutex>
|
||||
#include <unordered_map>
|
||||
using namespace std;
|
||||
|
||||
struct mscclRankState {
|
||||
int rank;
|
||||
bool initialized;
|
||||
mscclStatus status;
|
||||
mscclSavedProxyArgs savedProxyArgs;
|
||||
|
||||
mscclRankState() : rank(-1), initialized(false), status(), savedProxyArgs() {}
|
||||
explicit mscclRankState(const mscclRankState&) = default;
|
||||
};
|
||||
|
||||
static mutex rankStatesMutex;
|
||||
static unordered_map<int, shared_ptr<mscclRankState>> rankStates;
|
||||
|
||||
static inline mscclRankState& mscclGetRankState(int rank) {
|
||||
static thread_local shared_ptr<mscclRankState> threadRankState = make_shared<mscclRankState>();
|
||||
|
||||
if (rank < 0) {
|
||||
return *threadRankState;
|
||||
}
|
||||
|
||||
lock_guard<mutex> lock(rankStatesMutex);
|
||||
|
||||
auto rankStateIt = rankStates.find(rank);
|
||||
if (rankStateIt == rankStates.end()) {
|
||||
rankStateIt = rankStates.insert(make_pair(rank, make_shared<mscclRankState>(*threadRankState))).first;
|
||||
rankStateIt->second->rank = rank;
|
||||
}
|
||||
return *(rankStateIt->second);
|
||||
}
|
||||
|
||||
bool mscclInitialized(int rank) {
|
||||
return mscclGetRankState(rank).initialized;
|
||||
}
|
||||
|
||||
void mscclSetInitialized(int rank, bool initialized) {
|
||||
auto& state = mscclGetRankState(rank);
|
||||
assert(!initialized || !state.initialized);
|
||||
state.initialized = initialized;
|
||||
}
|
||||
|
||||
void mscclRemoveRank(int rank) {
|
||||
lock_guard<mutex> lock(rankStatesMutex);
|
||||
rankStates.erase(rank);
|
||||
}
|
||||
|
||||
mscclStatus& mscclGetStatus(int rank) {
|
||||
return mscclGetRankState(rank).status;
|
||||
}
|
||||
|
||||
mscclThreadLocalStatus& mscclGetThreadLocalStatus() {
|
||||
@@ -16,7 +67,6 @@ mscclThreadLocalStatus& mscclGetThreadLocalStatus() {
|
||||
return threadLocalStatus;
|
||||
}
|
||||
|
||||
mscclSavedProxyArgs& mscclGetSavedProxyArgs() {
|
||||
static mscclSavedProxyArgs savedProxyArgs;
|
||||
return savedProxyArgs;
|
||||
mscclSavedProxyArgs& mscclGetSavedProxyArgs(int rank) {
|
||||
return mscclGetRankState(rank).savedProxyArgs;
|
||||
}
|
||||
|
||||
Tagairt in Eagrán Nua
Cuir bac ar úsáideoir