Merge remote-tracking branch 'nccl/master' into develop

This commit is contained in:
Wenkai Du
2022-10-20 15:40:03 +00:00
51 changed files with 2139 additions and 1081 deletions
+71 -37
View File
@@ -407,17 +407,17 @@ ncclResult_t ncclProxySaveOp(struct ncclComm* comm, struct ncclProxyOp* op, bool
NCCLCHECK(SaveProxy(channel, proxyRecv, tree->up, op, 0, justInquire));
}
} break;
case ncclPatternCollTreeUpDown: {
// CollTree up
NCCLCHECK(SaveProxy(channel, proxySend, channel->collTree.out, op, 1, justInquire)); // For CollTree up, we are using push
// CollTree down
NCCLCHECK(SaveProxy(channel, proxyRecv, channel->collTree.out, op, 0, justInquire));
case ncclPatternCollnetChain: {
NCCLCHECK(SaveProxy(channel, proxySend, channel->collnetChain.up, op, 1, justInquire));
NCCLCHECK(SaveProxy(channel, proxyRecv, channel->collnetChain.up, op, 0, justInquire));
} break;
case ncclPatternCollnetDirect: {
NCCLCHECK(SaveProxy(channel, proxySend, channel->collnetDirect.out, op, 1, justInquire));
NCCLCHECK(SaveProxy(channel, proxyRecv, channel->collnetDirect.out, op, 0, justInquire));
} break;
case ncclPatternSend:
case ncclPatternRecv: {
if (op->root == comm->rank) return ncclSuccess;
op->nsteps = DIVUP(op->nbytes, op->chunkSize);
if (op->nsteps == 0) op->nsteps = 1;
NCCLCHECK(SaveProxy(channel, op->pattern == ncclPatternSend ? proxySend : proxyRecv, op->root, op, op->connIndex, justInquire));
} break;
}
@@ -433,16 +433,17 @@ ncclResult_t ncclProxyComputeP2p(struct ncclInfo* info, struct ncclProxyOp* op)
op->channelId = channelId;
op->sliceSteps = 1;
op->chunkSteps = 1;
op->protocol = NCCL_PROTO_SIMPLE;
op->dtype = info->datatype;
op->protocol = info->protocol;
int stepSize = info->comm->buffSizes[NCCL_PROTO_SIMPLE]/NCCL_STEPS;
if (info->comm->nNodes > 1) stepSize /= SENDRECV_SLICEFACTOR;
int stepSize = info->comm->buffSizes[op->protocol]/NCCL_STEPS;
// If nNodes > 1 and we're using Simple, reduce the stepSize to increase shared buffer utilization
if (info->comm->nNodes > 1 && op->protocol == NCCL_PROTO_SIMPLE) stepSize = info->comm->p2pNetChunkSize;
info->chunkSize = stepSize;
op->root = info->root;
op->nbytes = info->count;
struct ncclChannelPeer* peer = channel->peers + op->root;
struct ncclChannelPeer* peer = channel->peers + op->root;
if (info->coll == ncclFuncSend) {
op->pattern = ncclPatternSend;
if (op->root != info->comm->rank && peer->send[1].transportComm == &netTransport.send) {
@@ -465,6 +466,17 @@ ncclResult_t ncclProxyComputeP2p(struct ncclInfo* info, struct ncclProxyOp* op)
info->chunkSize = ncclParamChunkSize();
}
op->chunkSize = info->chunkSize;
// Compute nSteps for proxies
int chunkEffectiveSize = op->chunkSize;
if (op->protocol == NCCL_PROTO_LL) {
chunkEffectiveSize /= 2;
}
op->nbytes = stepSize;
op->nsteps = DIVUP(info->count, chunkEffectiveSize);
if (op->nsteps == 0) op->nsteps = 1;
return ncclSuccess;
}
@@ -617,7 +629,7 @@ ncclResult_t ncclSetThreadContext(struct ncclComm* comm) {
if (createThreadContext == -1) {
createThreadContext = ncclParamCreateThreadContext();
if (createThreadContext) {
if (CUPFN(cuCtxCreate_v3020) == nullptr || CUPFN(cuCtxDestroy) == nullptr || CUPFN(cuCtxSetCurrent) == nullptr) {
if (CUPFN(cuCtxCreate) == nullptr || CUPFN(cuCtxDestroy) == nullptr || CUPFN(cuCtxSetCurrent) == nullptr) {
WARN("Unable to create thread context due to old driver, disabling.");
createThreadContext = 0;
}
@@ -625,7 +637,7 @@ ncclResult_t ncclSetThreadContext(struct ncclComm* comm) {
}
if (createThreadContext) {
if (comm->proxyState.cudaCtx == NULL) {
if (CUPFN(cuCtxCreate_v3020(&comm->proxyState.cudaCtx,
if (CUPFN(cuCtxCreate(&comm->proxyState.cudaCtx,
CU_CTX_SCHED_SPIN|CU_CTX_MAP_HOST, comm->cudaDev)) != CUDA_SUCCESS) {
WARN("Failed to create CUDA context on device %d", comm->cudaDev);
createThreadContext = 0;
@@ -642,6 +654,9 @@ ncclResult_t ncclSetThreadContext(struct ncclComm* comm) {
return ncclSuccess;
}
// Set to SIGUSR1 or SIGUSR2 to help debug proxy state during hangs
NCCL_PARAM(ProxyDumpSignal, "PROXY_DUMP_SIGNAL", -1);
void* ncclProxyProgress(void *comm_) {
struct ncclComm* comm = (struct ncclComm*)comm_;
if (ncclSetThreadContext(comm) != ncclSuccess) {
@@ -653,7 +668,8 @@ void* ncclProxyProgress(void *comm_) {
struct ncclProxyProgressState* state = &comm->proxyState.progressState;
state->nextOps = -1;
signal(SIGUSR1, ncclDumpProxyState);
const int sig = ncclParamProxyDumpSignal();
if (sig != -1) signal(sig, ncclDumpProxyState);
ncclLastProxyState = state;
char threadName[NCCL_THREAD_NAMELEN];
snprintf(threadName, NCCL_THREAD_NAMELEN, "NCCL Progress%2d", comm->cudaDev);
@@ -665,7 +681,7 @@ void* ncclProxyProgress(void *comm_) {
int idle = 1;
ncclResult_t ret = progressOps(comm, state, state->active, &idle);
if (ret != ncclSuccess) {
comm->fatalError = ret;
(void) ncclCommSetAsyncError(comm, ret);
INFO(NCCL_ALL,"%s:%d -> %d [Proxy Thread]", __FILE__, __LINE__, ret);
return NULL;
}
@@ -677,7 +693,7 @@ void* ncclProxyProgress(void *comm_) {
ret = ncclProxyGetPostedOps(comm, &added);
if (added) { TIME_STOP(3); } else { TIME_CANCEL(3); }
if (ret != ncclSuccess) {
comm->fatalError = ret;
(void) ncclCommSetAsyncError(comm, ret);
INFO(NCCL_ALL,"%s:%d -> %d [Proxy Thread]", __FILE__, __LINE__, ret);
}
if (added == 0) {
@@ -783,9 +799,13 @@ static ncclResult_t ncclProxyGetConnection(struct ncclProxyConnectionPool* pool,
static ncclResult_t proxyFree(struct ncclProxyConnection* connection, struct ncclComm* comm) {
if (connection->send) {
NCCLCHECK(ncclTransports[connection->transport]->send.proxyFree(connection, comm));
if (ncclTransports[connection->transport]->send.proxyFree) {
NCCLCHECK(ncclTransports[connection->transport]->send.proxyFree(connection, comm));
}
} else {
NCCLCHECK(ncclTransports[connection->transport]->recv.proxyFree(connection, comm));
if (ncclTransports[connection->transport]->recv.proxyFree) {
NCCLCHECK(ncclTransports[connection->transport]->recv.proxyFree(connection, comm));
}
}
return ncclSuccess;
}
@@ -794,7 +814,10 @@ static ncclResult_t ncclProxyFreeConnections(struct ncclProxyConnectionPool* poo
for (int b=0; b<pool->banks; b++) {
int max = b == pool->banks-1 ? pool->offset : NCCL_PROXY_CONN_POOL_SIZE;
for (int i=0; i<max; i++) {
NCCLCHECK(proxyFree(pool->pools[b]+i, comm));
ncclProxyConnection *connection = pool->pools[b]+i;
if (connection->initFlag == true) {
NCCLCHECK(proxyFree(connection, comm));
}
}
free(pool->pools[b]);
}
@@ -813,8 +836,7 @@ ncclResult_t ncclProxyConnect(struct ncclComm* comm, int transport, int send, in
NCCLCHECK(ncclCalloc(&comm->proxyState.proxyOps, comm->localRanks));
NCCLCHECK(ncclCalloc(&comm->proxyState.sharedDevMems, comm->localRanks));
for (int r=0; r<comm->localRanks; r++) {
comm->proxyState.peerSocks[r].fd = -1;
comm->proxyState.peerSocks[r].abortFlag = comm->abortFlag;
NCCLCHECK(ncclSocketInit(&comm->proxyState.peerSocks[r], NULL, comm->abortFlag, 0));
}
}
NCCLCHECK(ncclTopoGetLocalRank(comm->topo, rank, &proxyConn->localRank));
@@ -944,6 +966,7 @@ static ncclResult_t proxyConnInit(struct ncclProxyLocalPeer* peer, struct ncclPr
NCCLCHECK(ncclSocketSend(sock, state->opsPoolShmSuffix, sizeof("XXXXXX")-1));
}
INFO(NCCL_NET, "New proxy %s connection %d from local rank %d, transport %d", connection->send ? "send":"recv", id, connection->localRank, connection->transport);
__atomic_store_n(&connection->initFlag, true, __ATOMIC_RELEASE);
return ncclSuccess;
}
@@ -1020,7 +1043,8 @@ void* ncclProxyService(void* _args) {
struct ncclProxyLocalPeer peers[NCCL_MAX_LOCAL_RANKS];
memset(&peers, 0, sizeof(struct ncclProxyLocalPeer)*NCCL_MAX_LOCAL_RANKS);
for (int s=0; s<NCCL_MAX_LOCAL_RANKS; s++) {
peers[s].sock.fd = pollfds[s].fd = -1;
ncclSocketInit(&peers[s].sock, NULL, comm->abortFlag, 0);
pollfds[s].fd = -1;
pollfds[s].events = POLLHUP|POLLIN;
}
pollfds[NCCL_MAX_LOCAL_RANKS].fd = comm->proxyState.listenSock->fd;
@@ -1030,8 +1054,9 @@ void* ncclProxyService(void* _args) {
int npeers = 0;
int stop = 0;
int asyncOpCount = 0;
while (stop == 0 || (stop == 1 && npeers > 0)) {
if (int error = poll(pollfds, NCCL_MAX_LOCAL_RANKS+1, asyncOpCount ? 0 : -1) < 0) {
while ((stop == 0 || (stop == 1 && npeers > 0)) && *comm->abortFlag == 0) {
/* never let proxy service thread blocks in poll, or it cannot receive abortFlag. */
if (int error = poll(pollfds, NCCL_MAX_LOCAL_RANKS+1, asyncOpCount ? 0 : 500) < 0) {
WARN("[Proxy Service] Poll failed with error %d", error);
return NULL;
}
@@ -1072,10 +1097,7 @@ void* ncclProxyService(void* _args) {
INFO(NCCL_INIT|NCCL_NET, "[Service thread] Connection closed by localRank %d", peer->localRank);
closeConn = 1;
} else {
if (type == ncclProxyMsgAbort) {
stop = 2;
closeConn = 1;
} else if (type == ncclProxyMsgStop) {
if (type == ncclProxyMsgStop) {
stop = 1;
closeConn = 1;
} else if (type == ncclProxyMsgClose) {
@@ -1105,6 +1127,10 @@ void* ncclProxyService(void* _args) {
}
}
}
/* wait until main thread flush all NCCL operations. */
while (*comm->abortFlag != 0 && __atomic_load_n(&comm->proxyState.safeAbortFlag, __ATOMIC_ACQUIRE) == 0)
usleep(1000);
// Wait for all operations to complete and stop progress thread before freeing any resource
if (ncclProxyProgressDestroy(comm) != ncclSuccess) {
WARN("[Proxy Service] proxyDestroy failed");
@@ -1134,15 +1160,23 @@ ncclResult_t ncclProxyCreate(struct ncclComm* comm) {
ncclResult_t ncclProxyDestroy(struct ncclComm* comm) {
struct ncclProxyState* state = &comm->proxyState;
if (state == NULL) return ncclSuccess;
if (state->peerAddresses) {
struct ncclSocket sock;
sock.abortFlag = NULL;
sock.asyncFlag = 0;
memcpy(&sock.addr, comm->proxyState.peerAddresses+comm->rank, sizeof(union ncclSocketAddress));
NCCLCHECK(ncclSocketConnect(&sock));
int type = (*comm->abortFlag) ? ncclProxyMsgAbort : ncclProxyMsgStop;
NCCLCHECK(ncclSocketSend(&sock, &type, sizeof(int)));
close(sock.fd);
if (*comm->abortFlag == 0) {
struct ncclSocket sock;
sock.abortFlag = NULL;
sock.asyncFlag = 0;
memcpy(&sock.addr, comm->proxyState.peerAddresses+comm->rank, sizeof(union ncclSocketAddress));
NCCLCHECK(ncclSocketConnect(&sock));
int type = ncclProxyMsgStop;
NCCLCHECK(ncclSocketSend(&sock, &type, sizeof(int)));
close(sock.fd);
} else {
/* when abortFlag is set, all socket related communications are no longer reliable. We need to
* set a flag to let proxy thread exit. */
__atomic_store_n(&state->safeAbortFlag, 1, __ATOMIC_RELEASE);
}
free(state->peerAddresses);
}
if (state->peerSocks) {