파일
rocm-systems/src/bootstrap.cc
T

658 라인
25 KiB
C++
Raw 일반 보기 히스토리

2018-09-24 16:06:59 -07:00
/*************************************************************************
2022-01-07 06:39:55 -08:00
* Copyright (c) 2016-2022, NVIDIA CORPORATION. All rights reserved.
2018-09-24 16:06:59 -07:00
*
* See LICENSE.txt for license information
************************************************************************/
#include "nccl.h"
#include "core.h"
#include "utils.h"
#include "bootstrap.h"
#include "net.h"
#include <unistd.h>
#include <sys/types.h>
2022-01-07 06:39:55 -08:00
#include "proxy.h"
2023-09-26 05:47:28 -07:00
#include "param.h"
2018-09-24 16:06:59 -07:00
2022-11-29 04:27:46 -08:00
struct bootstrapRootArgs {
struct ncclSocket* listenSock;
uint64_t magic;
};
2019-06-25 13:22:47 -07:00
/* Init functions */
2020-09-04 14:35:05 -07:00
static char bootstrapNetIfName[MAX_IF_NAME_SIZE+1];
2022-01-07 06:39:55 -08:00
static union ncclSocketAddress bootstrapNetIfAddr;
2020-09-04 14:35:05 -07:00
static int bootstrapNetInitDone = 0;
2019-06-25 13:22:47 -07:00
pthread_mutex_t bootstrapNetLock = PTHREAD_MUTEX_INITIALIZER;
ncclResult_t bootstrapNetInit() {
2020-09-04 14:35:05 -07:00
if (bootstrapNetInitDone == 0) {
2019-06-25 13:22:47 -07:00
pthread_mutex_lock(&bootstrapNetLock);
2020-09-04 14:35:05 -07:00
if (bootstrapNetInitDone == 0) {
2023-09-26 05:47:28 -07:00
const char* env = ncclGetEnv("NCCL_COMM_ID");
2020-09-04 14:35:05 -07:00
if (env) {
2022-01-07 06:39:55 -08:00
union ncclSocketAddress remoteAddr;
2022-11-29 04:27:46 -08:00
if (ncclSocketGetAddrFromString(&remoteAddr, env) != ncclSuccess) {
2020-09-04 14:35:05 -07:00
WARN("Invalid NCCL_COMM_ID, please use format: <ipv4>:<port> or [<ipv6>]:<port> or <hostname>:<port>");
2023-09-26 05:47:28 -07:00
pthread_mutex_unlock(&bootstrapNetLock);
2020-09-04 14:35:05 -07:00
return ncclInvalidArgument;
}
2022-01-07 06:39:55 -08:00
if (ncclFindInterfaceMatchSubnet(bootstrapNetIfName, &bootstrapNetIfAddr, &remoteAddr, MAX_IF_NAME_SIZE, 1) <= 0) {
2020-09-04 14:35:05 -07:00
WARN("NET/Socket : No usable listening interface found");
2023-09-26 05:47:28 -07:00
pthread_mutex_unlock(&bootstrapNetLock);
2020-09-04 14:35:05 -07:00
return ncclSystemError;
}
2019-06-25 13:22:47 -07:00
} else {
2022-01-07 06:39:55 -08:00
int nIfs = ncclFindInterfaces(bootstrapNetIfName, &bootstrapNetIfAddr, MAX_IF_NAME_SIZE, 1);
2020-09-04 14:35:05 -07:00
if (nIfs <= 0) {
WARN("Bootstrap : no socket interface found");
2023-09-26 05:47:28 -07:00
pthread_mutex_unlock(&bootstrapNetLock);
2020-09-04 14:35:05 -07:00
return ncclInternalError;
2019-06-25 13:22:47 -07:00
}
}
2020-09-04 14:35:05 -07:00
char line[SOCKET_NAME_MAXLEN+MAX_IF_NAME_SIZE+2];
sprintf(line, " %s:", bootstrapNetIfName);
2022-01-07 06:39:55 -08:00
ncclSocketToString(&bootstrapNetIfAddr, line+strlen(line));
2020-09-04 14:35:05 -07:00
INFO(NCCL_INIT, "Bootstrap : Using%s", line);
bootstrapNetInitDone = 1;
2019-06-25 13:22:47 -07:00
}
pthread_mutex_unlock(&bootstrapNetLock);
}
return ncclSuccess;
}
2018-09-24 16:06:59 -07:00
2019-06-25 13:22:47 -07:00
/* Socket Interface Selection type */
enum bootstrapInterface_t { findSubnetIf = -1, dontCareIf = -2 };
// Additional sync functions
2022-01-07 06:39:55 -08:00
static ncclResult_t bootstrapNetSend(struct ncclSocket* sock, void* data, int size) {
NCCLCHECK(ncclSocketSend(sock, &size, sizeof(int)));
NCCLCHECK(ncclSocketSend(sock, data, size));
2018-09-24 16:06:59 -07:00
return ncclSuccess;
}
2022-01-07 06:39:55 -08:00
static ncclResult_t bootstrapNetRecv(struct ncclSocket* sock, void* data, int size) {
2019-06-25 13:22:47 -07:00
int recvSize;
2022-01-07 06:39:55 -08:00
NCCLCHECK(ncclSocketRecv(sock, &recvSize, sizeof(int)));
2019-06-25 13:22:47 -07:00
if (recvSize > size) {
2021-02-09 15:34:08 -08:00
WARN("Message truncated : received %d bytes instead of %d", recvSize, size);
2019-06-25 13:22:47 -07:00
return ncclInternalError;
}
2022-01-07 06:39:55 -08:00
NCCLCHECK(ncclSocketRecv(sock, data, std::min(recvSize, size)));
2018-09-24 16:06:59 -07:00
return ncclSuccess;
}
2024-03-26 06:08:55 -07:00
static ncclResult_t bootstrapNetSendRecv(struct ncclSocket* sendSock, void* sendData, int sendSize, struct ncclSocket* recvSock, void* recvData, int recvSize) {
int senderRecvSize;
NCCLCHECK(ncclSocketSendRecv(sendSock, &sendSize, sizeof(int), recvSock, &senderRecvSize, sizeof(int)));
if (senderRecvSize > recvSize) {
WARN("Message truncated : received %d bytes instead of %d", senderRecvSize, recvSize);
return ncclInternalError;
}
NCCLCHECK(ncclSocketSendRecv(sendSock, sendData, sendSize, recvSock, recvData, recvSize));
return ncclSuccess;
}
2018-09-24 16:06:59 -07:00
struct extInfo {
int rank;
int nranks;
2022-01-07 06:39:55 -08:00
union ncclSocketAddress extAddressListenRoot;
union ncclSocketAddress extAddressListen;
2018-09-24 16:06:59 -07:00
};
#include <sys/resource.h>
static ncclResult_t setFilesLimit() {
struct rlimit filesLimit;
SYSCHECK(getrlimit(RLIMIT_NOFILE, &filesLimit), "getrlimit");
filesLimit.rlim_cur = filesLimit.rlim_max;
SYSCHECK(setrlimit(RLIMIT_NOFILE, &filesLimit), "setrlimit");
return ncclSuccess;
}
2022-11-29 04:27:46 -08:00
static void *bootstrapRoot(void* rargs) {
struct bootstrapRootArgs* args = (struct bootstrapRootArgs*)rargs;
struct ncclSocket* listenSock = args->listenSock;
uint64_t magic = args->magic;
2020-09-04 14:35:05 -07:00
ncclResult_t res = ncclSuccess;
int nranks = 0, c = 0;
2018-09-24 16:06:59 -07:00
struct extInfo info;
2022-01-07 06:39:55 -08:00
union ncclSocketAddress *rankAddresses = NULL;
union ncclSocketAddress *rankAddressesRoot = NULL; // for initial rank <-> root information exchange
union ncclSocketAddress *zero = NULL;
2020-09-04 14:35:05 -07:00
NCCLCHECKGOTO(ncclCalloc(&zero, 1), res, out);
2018-09-24 16:06:59 -07:00
setFilesLimit();
2018-12-13 15:56:12 -08:00
TRACE(NCCL_INIT, "BEGIN");
2018-09-24 16:06:59 -07:00
/* Receive addresses from all ranks */
do {
2022-01-07 06:39:55 -08:00
struct ncclSocket sock;
2022-11-29 04:27:46 -08:00
NCCLCHECKGOTO(ncclSocketInit(&sock), res, out);
2022-01-07 06:39:55 -08:00
NCCLCHECKGOTO(ncclSocketAccept(&sock, listenSock), res, out);
NCCLCHECKGOTO(bootstrapNetRecv(&sock, &info, sizeof(info)), res, out);
2022-11-29 04:27:46 -08:00
NCCLCHECKGOTO(ncclSocketClose(&sock), res, out);
2018-10-24 14:44:59 -07:00
if (c == 0) {
2018-09-24 16:06:59 -07:00
nranks = info.nranks;
2020-09-04 14:35:05 -07:00
NCCLCHECKGOTO(ncclCalloc(&rankAddresses, nranks), res, out);
NCCLCHECKGOTO(ncclCalloc(&rankAddressesRoot, nranks), res, out);
2018-09-24 16:06:59 -07:00
}
if (nranks != info.nranks) {
WARN("Bootstrap Root : mismatch in rank count from procs %d : %d", nranks, info.nranks);
goto out;
}
2022-01-07 06:39:55 -08:00
if (memcmp(zero, &rankAddressesRoot[info.rank], sizeof(union ncclSocketAddress)) != 0) {
2018-10-24 14:44:59 -07:00
WARN("Bootstrap Root : rank %d of %d ranks has already checked in", info.rank, nranks);
goto out;
2018-09-24 16:06:59 -07:00
}
2018-12-13 15:56:12 -08:00
// Save the connection handle for that rank
2022-01-07 06:39:55 -08:00
memcpy(rankAddressesRoot+info.rank, &info.extAddressListenRoot, sizeof(union ncclSocketAddress));
memcpy(rankAddresses+info.rank, &info.extAddressListen, sizeof(union ncclSocketAddress));
2018-09-24 16:06:59 -07:00
2018-10-24 14:44:59 -07:00
++c;
2019-11-19 14:57:39 -08:00
TRACE(NCCL_INIT, "Received connect from rank %d total %d/%d", info.rank, c, nranks);
2018-10-24 14:44:59 -07:00
} while (c < nranks);
2019-11-19 14:57:39 -08:00
TRACE(NCCL_INIT, "COLLECTED ALL %d HANDLES", nranks);
2018-09-24 16:06:59 -07:00
2018-10-24 14:44:59 -07:00
// Send the connect handle for the next rank in the AllGather ring
for (int r=0; r<nranks; ++r) {
int next = (r+1) % nranks;
2022-01-07 06:39:55 -08:00
struct ncclSocket sock;
2022-11-29 04:27:46 -08:00
NCCLCHECKGOTO(ncclSocketInit(&sock, rankAddressesRoot+r, magic, ncclSocketTypeBootstrap), res, out);
2022-01-07 06:39:55 -08:00
NCCLCHECKGOTO(ncclSocketConnect(&sock), res, out);
NCCLCHECKGOTO(bootstrapNetSend(&sock, rankAddresses+next, sizeof(union ncclSocketAddress)), res, out);
2022-11-29 04:27:46 -08:00
NCCLCHECKGOTO(ncclSocketClose(&sock), res, out);
2018-10-24 14:44:59 -07:00
}
2019-11-19 14:57:39 -08:00
TRACE(NCCL_INIT, "SENT OUT ALL %d HANDLES", nranks);
2018-09-24 16:06:59 -07:00
out:
2022-11-29 04:27:46 -08:00
if (listenSock != NULL) {
ncclSocketClose(listenSock);
free(listenSock);
}
2020-09-04 14:35:05 -07:00
if (rankAddresses) free(rankAddresses);
if (rankAddressesRoot) free(rankAddressesRoot);
if (zero) free(zero);
2022-11-29 04:27:46 -08:00
free(rargs);
2018-12-13 15:56:12 -08:00
TRACE(NCCL_INIT, "DONE");
2018-09-24 16:06:59 -07:00
return NULL;
}
2022-11-29 04:27:46 -08:00
ncclResult_t bootstrapCreateRoot(struct ncclBootstrapHandle* handle, bool idFromEnv) {
2022-01-07 06:39:55 -08:00
struct ncclSocket* listenSock;
2022-11-29 04:27:46 -08:00
struct bootstrapRootArgs* args;
pthread_t thread;
2022-01-07 06:39:55 -08:00
NCCLCHECK(ncclCalloc(&listenSock, 1));
2022-11-29 04:27:46 -08:00
NCCLCHECK(ncclSocketInit(listenSock, &handle->addr, handle->magic, ncclSocketTypeBootstrap, NULL, 0));
2022-01-07 06:39:55 -08:00
NCCLCHECK(ncclSocketListen(listenSock));
2022-11-29 04:27:46 -08:00
NCCLCHECK(ncclSocketGetAddr(listenSock, &handle->addr));
NCCLCHECK(ncclCalloc(&args, 1));
args->listenSock = listenSock;
args->magic = handle->magic;
NEQCHECK(pthread_create(&thread, NULL, bootstrapRoot, (void*)args), 0);
2022-01-07 06:39:55 -08:00
ncclSetThreadName(thread, "NCCL BootstrapR");
2022-11-29 04:27:46 -08:00
NEQCHECK(pthread_detach(thread), 0); // will not be pthread_join()'d
2018-09-24 16:06:59 -07:00
return ncclSuccess;
}
2022-11-29 04:27:46 -08:00
ncclResult_t bootstrapGetUniqueId(struct ncclBootstrapHandle* handle) {
memset(handle, 0, sizeof(ncclBootstrapHandle));
2018-09-24 16:06:59 -07:00
2023-09-26 05:47:28 -07:00
const char* env = ncclGetEnv("NCCL_COMM_ID");
2018-09-24 16:06:59 -07:00
if (env) {
2020-05-12 14:40:18 -07:00
INFO(NCCL_ENV, "NCCL_COMM_ID set by environment to %s", env);
2022-11-29 04:27:46 -08:00
if (ncclSocketGetAddrFromString(&handle->addr, env) != ncclSuccess) {
2018-09-24 16:06:59 -07:00
WARN("Invalid NCCL_COMM_ID, please use format: <ipv4>:<port> or [<ipv6>]:<port> or <hostname>:<port>");
return ncclInvalidArgument;
}
2024-06-11 01:28:01 -07:00
handle->magic = NCCL_MAGIC;
2018-09-24 16:06:59 -07:00
} else {
2024-06-11 01:28:01 -07:00
NCCLCHECK(getRandomData(&handle->magic, sizeof(handle->magic)));
2022-11-29 04:27:46 -08:00
memcpy(&handle->addr, &bootstrapNetIfAddr, sizeof(union ncclSocketAddress));
NCCLCHECK(bootstrapCreateRoot(handle, false));
2018-09-24 16:06:59 -07:00
}
return ncclSuccess;
}
2018-12-13 15:56:12 -08:00
struct unexConn {
int peer;
2021-04-12 16:00:11 -07:00
int tag;
2022-01-07 06:39:55 -08:00
struct ncclSocket sock;
2018-12-13 15:56:12 -08:00
struct unexConn* next;
};
2022-01-07 06:39:55 -08:00
struct bootstrapState {
struct ncclSocket listenSock;
struct ncclSocket ringRecvSocket;
struct ncclSocket ringSendSocket;
union ncclSocketAddress* peerCommAddresses;
union ncclSocketAddress* peerProxyAddresses;
2024-02-05 05:06:02 -08:00
uint64_t* peerProxyAddressesUDS;
2018-12-13 15:56:12 -08:00
struct unexConn* unexpectedConnections;
2020-09-04 14:35:05 -07:00
int cudaDev;
2018-09-24 16:06:59 -07:00
int rank;
int nranks;
2022-11-29 04:27:46 -08:00
uint64_t magic;
2022-01-07 06:39:55 -08:00
volatile uint32_t *abortFlag;
2018-09-24 16:06:59 -07:00
};
2022-11-29 04:27:46 -08:00
ncclResult_t bootstrapInit(struct ncclBootstrapHandle* handle, struct ncclComm* comm) {
2022-01-07 06:39:55 -08:00
int rank = comm->rank;
int nranks = comm->nRanks;
struct bootstrapState* state;
2022-11-29 04:27:46 -08:00
struct ncclSocket* proxySocket;
ncclSocketAddress nextAddr;
struct ncclSocket sock, listenSockRoot;
struct extInfo info = { 0 };
2018-09-24 16:06:59 -07:00
NCCLCHECK(ncclCalloc(&state, 1));
state->rank = rank;
state->nranks = nranks;
2022-01-07 06:39:55 -08:00
state->abortFlag = comm->abortFlag;
comm->bootstrap = state;
2022-11-29 04:27:46 -08:00
comm->magic = state->magic = handle->magic;
2018-12-13 15:56:12 -08:00
TRACE(NCCL_INIT, "rank %d nranks %d", rank, nranks);
2018-09-24 16:06:59 -07:00
info.rank = rank;
info.nranks = nranks;
2022-01-07 06:39:55 -08:00
// Create socket for other ranks to contact me
2022-11-29 04:27:46 -08:00
NCCLCHECK(ncclSocketInit(&state->listenSock, &bootstrapNetIfAddr, comm->magic, ncclSocketTypeBootstrap, comm->abortFlag));
2022-01-07 06:39:55 -08:00
NCCLCHECK(ncclSocketListen(&state->listenSock));
2022-11-29 04:27:46 -08:00
NCCLCHECK(ncclSocketGetAddr(&state->listenSock, &info.extAddressListen));
2018-12-13 15:56:12 -08:00
2022-01-07 06:39:55 -08:00
// Create socket for root to contact me
2022-11-29 04:27:46 -08:00
NCCLCHECK(ncclSocketInit(&listenSockRoot, &bootstrapNetIfAddr, comm->magic, ncclSocketTypeBootstrap, comm->abortFlag));
2022-01-07 06:39:55 -08:00
NCCLCHECK(ncclSocketListen(&listenSockRoot));
2022-11-29 04:27:46 -08:00
NCCLCHECK(ncclSocketGetAddr(&listenSockRoot, &info.extAddressListenRoot));
2020-09-04 14:35:05 -07:00
// stagger connection times to avoid an overload of the root
2018-12-13 15:56:12 -08:00
if (nranks > 128) {
long msec = rank;
struct timespec tv;
tv.tv_sec = msec / 1000;
tv.tv_nsec = 1000000 * (msec % 1000);
TRACE(NCCL_INIT, "rank %d delaying connection to root by %ld msec", rank, msec);
(void) nanosleep(&tv, NULL);
}
2018-10-24 14:44:59 -07:00
2018-12-13 15:56:12 -08:00
// send info on my listening socket to root
2022-11-29 04:27:46 -08:00
NCCLCHECK(ncclSocketInit(&sock, &handle->addr, comm->magic, ncclSocketTypeBootstrap, comm->abortFlag));
2022-01-07 06:39:55 -08:00
NCCLCHECK(ncclSocketConnect(&sock));
NCCLCHECK(bootstrapNetSend(&sock, &info, sizeof(info)));
2022-11-29 04:27:46 -08:00
NCCLCHECK(ncclSocketClose(&sock));
2018-10-24 14:44:59 -07:00
// get info on my "next" rank in the bootstrap ring from root
2022-11-29 04:27:46 -08:00
NCCLCHECK(ncclSocketInit(&sock));
2022-01-07 06:39:55 -08:00
NCCLCHECK(ncclSocketAccept(&sock, &listenSockRoot));
2022-11-29 04:27:46 -08:00
NCCLCHECK(bootstrapNetRecv(&sock, &nextAddr, sizeof(union ncclSocketAddress)));
NCCLCHECK(ncclSocketClose(&sock));
NCCLCHECK(ncclSocketClose(&listenSockRoot));
2018-10-24 14:44:59 -07:00
2022-11-29 04:27:46 -08:00
NCCLCHECK(ncclSocketInit(&state->ringSendSocket, &nextAddr, comm->magic, ncclSocketTypeBootstrap, comm->abortFlag));
2022-01-07 06:39:55 -08:00
NCCLCHECK(ncclSocketConnect(&state->ringSendSocket));
2018-10-24 14:44:59 -07:00
// Accept the connect request from the previous rank in the AllGather ring
2022-11-29 04:27:46 -08:00
NCCLCHECK(ncclSocketInit(&state->ringRecvSocket));
2022-01-07 06:39:55 -08:00
NCCLCHECK(ncclSocketAccept(&state->ringRecvSocket, &state->listenSock));
2018-12-13 15:56:12 -08:00
// AllGather all listen handlers
2020-09-04 14:35:05 -07:00
NCCLCHECK(ncclCalloc(&state->peerCommAddresses, nranks));
2022-11-29 04:27:46 -08:00
NCCLCHECK(ncclSocketGetAddr(&state->listenSock, state->peerCommAddresses+rank));
2022-01-07 06:39:55 -08:00
NCCLCHECK(bootstrapAllGather(state, state->peerCommAddresses, sizeof(union ncclSocketAddress)));
// Create the service proxy
NCCLCHECK(ncclCalloc(&state->peerProxyAddresses, nranks));
2024-02-05 05:06:02 -08:00
NCCLCHECK(ncclCalloc(&state->peerProxyAddressesUDS, nranks));
2022-11-29 04:27:46 -08:00
// proxy is aborted through a message; don't set abortFlag
2022-01-07 06:39:55 -08:00
NCCLCHECK(ncclCalloc(&proxySocket, 1));
2022-11-29 04:27:46 -08:00
NCCLCHECK(ncclSocketInit(proxySocket, &bootstrapNetIfAddr, comm->magic, ncclSocketTypeProxy, comm->abortFlag));
2022-01-07 06:39:55 -08:00
NCCLCHECK(ncclSocketListen(proxySocket));
2022-11-29 04:27:46 -08:00
NCCLCHECK(ncclSocketGetAddr(proxySocket, state->peerProxyAddresses+rank));
2022-01-07 06:39:55 -08:00
NCCLCHECK(bootstrapAllGather(state, state->peerProxyAddresses, sizeof(union ncclSocketAddress)));
2024-02-05 05:06:02 -08:00
// cuMem UDS support
2024-02-26 02:52:26 -08:00
// Make sure we create a unique UDS socket name
uint64_t randId;
NCCLCHECK(getRandomData(&randId, sizeof(randId)));
state->peerProxyAddressesUDS[rank] = getPidHash()+randId;
2024-02-05 05:06:02 -08:00
NCCLCHECK(bootstrapAllGather(state, state->peerProxyAddressesUDS, sizeof(*state->peerProxyAddressesUDS)));
NCCLCHECK(ncclProxyInit(comm, proxySocket, state->peerProxyAddresses, state->peerProxyAddressesUDS));
2018-12-13 15:56:12 -08:00
TRACE(NCCL_INIT, "rank %d nranks %d - DONE", rank, nranks);
2018-09-24 16:06:59 -07:00
return ncclSuccess;
}
2023-04-03 05:32:07 -07:00
ncclResult_t bootstrapSplit(struct ncclBootstrapHandle* handle, struct ncclComm* comm, struct ncclComm* parent, int color, int key, int* parentRanks) {
ncclResult_t ret = ncclSuccess;
int rank = comm->rank;
int nranks = comm->nRanks;
int prev, next;
ncclSocketAddress listenAddr, tmpAddr;
struct ncclSocket* proxySocket;
struct bootstrapState* state;
NCCLCHECKGOTO(ncclCalloc(&state, 1), ret, fail);
state->rank = rank;
state->nranks = nranks;
state->abortFlag = comm->abortFlag;
comm->bootstrap = state;
comm->magic = state->magic = handle->magic;
prev = parentRanks[(rank-1+nranks)%nranks];
next = parentRanks[(rank+1)%nranks];
// Setup my sockets for the allgather ring and other p2p connections
NCCLCHECKGOTO(ncclSocketInit(&state->listenSock, &bootstrapNetIfAddr, comm->magic, ncclSocketTypeBootstrap, comm->abortFlag, 0), ret, fail);
NCCLCHECKGOTO(ncclSocketInit(&state->ringRecvSocket, NULL, comm->magic, ncclSocketTypeBootstrap, comm->abortFlag, 0), ret, fail);
// Create socket for other ranks to contact me
NCCLCHECKGOTO(ncclSocketListen(&state->listenSock), ret, fail);
// Get addr from next rank
NCCLCHECKGOTO(ncclSocketGetAddr(&state->listenSock, &listenAddr), ret, fail);
NCCLCHECKGOTO(bootstrapSend(parent->bootstrap, prev, -2, &listenAddr, sizeof(union ncclSocketAddress)), ret, fail);
NCCLCHECKGOTO(bootstrapRecv(parent->bootstrap, next, -2, &tmpAddr, sizeof(union ncclSocketAddress)), ret, fail);
NCCLCHECKGOTO(ncclSocketInit(&state->ringSendSocket, &tmpAddr, comm->magic, ncclSocketTypeBootstrap, comm->abortFlag, 0), ret, fail);
NCCLCHECKGOTO(ncclSocketConnect(&state->ringSendSocket), ret, fail);
// Accept the connect request from the previous rank in the AllGather ring
NCCLCHECKGOTO(ncclSocketAccept(&state->ringRecvSocket, &state->listenSock), ret, fail);
// AllGather all listen handlers
NCCLCHECKGOTO(ncclCalloc(&state->peerCommAddresses, nranks), ret, fail);
memcpy(state->peerCommAddresses+rank, &listenAddr, sizeof(union ncclSocketAddress));
NCCLCHECKGOTO(bootstrapAllGather(state, state->peerCommAddresses, sizeof(union ncclSocketAddress)), ret, fail);
if (parent->config.splitShare) {
/* map local rank to top parent local rank. */
for (int i = 0; i < nranks; ++i) {
comm->topParentRanks[i] = parent->topParentRanks[parentRanks[i]];
}
} else {
// Create the service proxy
NCCLCHECKGOTO(ncclCalloc(&state->peerProxyAddresses, nranks), ret, fail);
NCCLCHECKGOTO(ncclCalloc(&proxySocket, 1), ret, fail);
NCCLCHECKGOTO(ncclSocketInit(proxySocket, &bootstrapNetIfAddr, comm->magic, ncclSocketTypeProxy, comm->abortFlag, 0), ret, fail);
NCCLCHECKGOTO(ncclSocketListen(proxySocket), ret, fail);
NCCLCHECKGOTO(ncclSocketGetAddr(proxySocket, &tmpAddr), ret, fail);
memcpy(state->peerProxyAddresses + rank, &tmpAddr, sizeof(union ncclSocketAddress));
NCCLCHECKGOTO(bootstrapAllGather(state, state->peerProxyAddresses, sizeof(union ncclSocketAddress)), ret, fail);
2024-02-05 05:06:02 -08:00
// cuMem UDS support
NCCLCHECKGOTO(ncclCalloc(&state->peerProxyAddressesUDS, nranks), ret, fail);
2024-02-26 02:52:26 -08:00
// Make sure we create a unique UDS socket name
uint64_t randId;
NCCLCHECKGOTO(getRandomData(&randId, sizeof(randId)), ret, fail);
state->peerProxyAddressesUDS[rank] = getPidHash()+randId;
2024-02-05 05:06:02 -08:00
NCCLCHECKGOTO(bootstrapAllGather(state, state->peerProxyAddressesUDS, sizeof(*state->peerProxyAddressesUDS)), ret, fail);
NCCLCHECKGOTO(ncclProxyInit(comm, proxySocket, state->peerProxyAddresses, state->peerProxyAddressesUDS), ret, fail);
2023-04-03 05:32:07 -07:00
}
2024-02-05 05:06:02 -08:00
INFO(NCCL_INIT, "bootstrapSplit: comm %p parent %p rank %d nranks %d color %d key %d prev %d next %d - DONE", comm, parent, rank, nranks, color, key, prev, next);
2023-04-03 05:32:07 -07:00
exit:
return ret;
fail:
goto exit;
}
2024-03-26 06:08:55 -07:00
// Bootstrap send/receive functions
//
// We do not keep connections opened with all ranks at all times, and we have no guarantee
// that connections to our unique listen socket will arrive in the same order as we need
// them. Therefore, when establishing a connection, the sender sends a (peer, tag) tuple to
// allow the receiver to identify the flow, and keep it in an unexpected queue if needed.
2018-10-24 14:44:59 -07:00
2024-03-26 06:08:55 -07:00
ncclResult_t bootstrapConnect(void* commState, int peer, int tag, struct ncclSocket* sock) {
ncclResult_t ret = ncclSuccess;
struct bootstrapState* state = (struct bootstrapState*)commState;
2018-09-24 16:06:59 -07:00
2024-03-26 06:08:55 -07:00
NCCLCHECKGOTO(ncclSocketInit(sock, state->peerCommAddresses+peer, state->magic, ncclSocketTypeBootstrap), ret, fail);
NCCLCHECKGOTO(ncclSocketConnect(sock), ret, fail);
NCCLCHECKGOTO(bootstrapNetSend(sock, &state->rank, sizeof(int)), ret, fail);
NCCLCHECKGOTO(bootstrapNetSend(sock, &tag, sizeof(int)), ret, fail);
2018-09-24 16:06:59 -07:00
return ncclSuccess;
2024-03-26 06:08:55 -07:00
fail:
NCCLCHECK(ncclSocketClose(sock));
return ret;
2018-09-24 16:06:59 -07:00
}
2021-04-12 16:00:11 -07:00
ncclResult_t bootstrapSend(void* commState, int peer, int tag, void* data, int size) {
2022-11-29 04:27:46 -08:00
ncclResult_t ret = ncclSuccess;
2022-01-07 06:39:55 -08:00
struct ncclSocket sock;
2022-08-18 02:53:17 -07:00
2024-03-26 06:08:55 -07:00
TRACE(NCCL_BOOTSTRAP, "Sending to peer=%d tag=%d size=%d", peer, tag, size);
NCCLCHECK(bootstrapConnect(commState, peer, tag, &sock));
NCCLCHECKGOTO(bootstrapNetSend(&sock, data, size), ret, exit);
TRACE(NCCL_BOOTSTRAP, "Sent to peer=%d tag=%d size=%d", peer, tag, size);
2022-11-29 04:27:46 -08:00
exit:
NCCLCHECK(ncclSocketClose(&sock));
return ret;
2023-02-27 02:48:21 -08:00
}
2022-01-07 06:39:55 -08:00
ncclResult_t unexpectedEnqueue(struct bootstrapState* state, int peer, int tag, struct ncclSocket* sock) {
2018-12-13 15:56:12 -08:00
// New unex
struct unexConn* unex;
NCCLCHECK(ncclCalloc(&unex, 1));
unex->peer = peer;
2021-04-12 16:00:11 -07:00
unex->tag = tag;
2022-01-07 06:39:55 -08:00
memcpy(&unex->sock, sock, sizeof(struct ncclSocket));
2018-12-13 15:56:12 -08:00
// Enqueue
struct unexConn* list = state->unexpectedConnections;
if (list == NULL) {
state->unexpectedConnections = unex;
return ncclSuccess;
}
while (list->next) list = list->next;
list->next = unex;
return ncclSuccess;
}
2018-09-24 16:06:59 -07:00
2022-11-29 04:27:46 -08:00
ncclResult_t unexpectedDequeue(struct bootstrapState* state, int peer, int tag, struct ncclSocket* sock, int* found) {
2018-12-13 15:56:12 -08:00
struct unexConn* elem = state->unexpectedConnections;
struct unexConn* prev = NULL;
2022-11-29 04:27:46 -08:00
*found = 0;
2018-12-13 15:56:12 -08:00
while (elem) {
2021-04-12 16:00:11 -07:00
if (elem->peer == peer && elem->tag == tag) {
2018-12-13 15:56:12 -08:00
if (prev == NULL) {
state->unexpectedConnections = elem->next;
} else {
prev->next = elem->next;
}
2022-01-07 06:39:55 -08:00
memcpy(sock, &elem->sock, sizeof(struct ncclSocket));
2018-12-13 15:56:12 -08:00
free(elem);
2022-11-29 04:27:46 -08:00
*found = 1;
2022-01-07 06:39:55 -08:00
return ncclSuccess;
2018-12-13 15:56:12 -08:00
}
prev = elem;
elem = elem->next;
}
2022-01-07 06:39:55 -08:00
return ncclSuccess;
2018-12-13 15:56:12 -08:00
}
2022-11-29 04:27:46 -08:00
static void unexpectedFree(struct bootstrapState* state) {
struct unexConn* elem = state->unexpectedConnections;
struct unexConn* prev = NULL;
while (elem) {
prev = elem;
elem = elem->next;
free(prev);
}
return;
}
2018-12-13 15:56:12 -08:00
// We can't know who we'll receive from, so we need to receive everything at once
2024-03-26 06:08:55 -07:00
ncclResult_t bootstrapAccept(void* commState, int peer, int tag, struct ncclSocket* sock) {
2022-11-29 04:27:46 -08:00
ncclResult_t ret = ncclSuccess;
2022-01-07 06:39:55 -08:00
struct bootstrapState* state = (struct bootstrapState*)commState;
2022-11-29 04:27:46 -08:00
int newPeer, newTag;
2018-12-13 15:56:12 -08:00
// Search unexpected connections first
2022-11-29 04:27:46 -08:00
int found;
2024-03-26 06:08:55 -07:00
NCCLCHECK(unexpectedDequeue(state, peer, tag, sock, &found));
if (found) return ncclSuccess;
2018-12-13 15:56:12 -08:00
// Then look for new connections
while (1) {
2024-03-26 06:08:55 -07:00
NCCLCHECKGOTO(ncclSocketInit(sock), ret, fail);
NCCLCHECKGOTO(ncclSocketAccept(sock, &state->listenSock), ret, fail);
NCCLCHECKGOTO(bootstrapNetRecv(sock, &newPeer, sizeof(int)), ret, fail);
NCCLCHECKGOTO(bootstrapNetRecv(sock, &newTag, sizeof(int)), ret, fail);
if (newPeer == peer && newTag == tag) return ncclSuccess;
NCCLCHECKGOTO(unexpectedEnqueue(state, newPeer, newTag, sock), ret, fail);
2018-12-13 15:56:12 -08:00
}
2024-03-26 06:08:55 -07:00
return ncclSuccess;
fail:
NCCLCHECK(ncclSocketClose(sock));
return ret;
}
// We can't know who we'll receive from, so we need to receive everything at once
ncclResult_t bootstrapRecv(void* commState, int peer, int tag, void* data, int size) {
ncclResult_t ret;
struct ncclSocket sock;
NCCLCHECK(bootstrapAccept(commState, peer, tag, &sock));
TRACE(NCCL_BOOTSTRAP, "Receiving tag=%d peer=%d size=%d", tag, peer, size);
NCCLCHECKGOTO(bootstrapNetRecv(&sock, ((char*)data), size), ret, exit);
2022-11-29 04:27:46 -08:00
exit:
NCCLCHECK(ncclSocketClose(&sock));
return ret;
2024-03-26 06:08:55 -07:00
}
// Collective algorithms, based on bootstrapSend/Recv, and sometimes bootstrapConnect/Accept
ncclResult_t bootstrapRingAllGather(struct ncclSocket* prevSocket, struct ncclSocket* nextSocket, int rank, int nranks, char* data, int size) {
/* Simple ring based AllGather
* At each step i receive data from (rank-i-1) from prev
* and send previous step's data from (rank-i) to next
*/
for (int i=0; i<nranks-1; i++) {
size_t rslice = (rank - i - 1 + nranks) % nranks;
size_t sslice = (rank - i + nranks) % nranks;
// Send slice to the right, recv slice from the left
NCCLCHECK(bootstrapNetSendRecv(nextSocket, data+sslice*size, size, prevSocket, data+rslice*size, size));
}
return ncclSuccess;
}
ncclResult_t bootstrapAllGather(void* commState, void* allData, int size) {
struct bootstrapState* state = (struct bootstrapState*)commState;
int rank = state->rank;
int nranks = state->nranks;
TRACE(NCCL_INIT, "rank %d nranks %d size %d", rank, nranks, size);
NCCLCHECK(bootstrapRingAllGather(&state->ringRecvSocket, &state->ringSendSocket, rank, nranks, (char*)allData, size));
TRACE(NCCL_INIT, "rank %d nranks %d size %d - DONE", rank, nranks, size);
return ncclSuccess;
}
ncclResult_t bootstrapIntraNodeBarrier(void* commState, int *ranks, int rank, int nranks, int tag) {
if (nranks == 1) return ncclSuccess;
TRACE(NCCL_INIT, "rank %d nranks %d tag %x - ENTER", rank, nranks, tag);
/* Simple [intra] process barrier
*
* Based on the dissemination algorithm by Debra Hensgen, Raphael Finkel, and Udi Manbet,
* "Two Algorithms for Barrier Synchronization," International Journal of Parallel Programming, 17(1):1-17, 1988"
*/
int data[1];
for (int mask=1; mask<nranks; mask<<=1) {
int src = (rank - mask + nranks) % nranks;
int dst = (rank + mask) % nranks;
NCCLCHECK(bootstrapSend(commState, ranks ? ranks[dst] : dst, tag, data, sizeof(data)));
NCCLCHECK(bootstrapRecv(commState, ranks ? ranks[src] : src, tag, data, sizeof(data)));
}
TRACE(NCCL_INIT, "rank %d nranks %d tag %x - DONE", rank, nranks, tag);
return ncclSuccess;
}
ncclResult_t bootstrapBarrier(void* commState, int rank, int nranks, int tag) {
return bootstrapIntraNodeBarrier(commState, NULL, rank, nranks, tag);
}
ncclResult_t bootstrapIntraNodeAllGather(void* commState, int *ranks, int rank, int nranks, void* allData, int size) {
if (nranks == 1) return ncclSuccess;
TRACE(NCCL_INIT, "rank %d nranks %d size %d - ENTER", rank, nranks, size);
int prevRank = ranks[(rank - 1 + nranks)%nranks];
int nextRank = ranks[(rank + 1) % nranks];
struct ncclSocket prevSocket, nextSocket;
NCCLCHECK(bootstrapConnect(commState, nextRank, 0, &nextSocket));
NCCLCHECK(bootstrapAccept(commState, prevRank, 0, &prevSocket));
NCCLCHECK(bootstrapRingAllGather(&prevSocket, &nextSocket, rank, nranks, (char*)allData, size));
NCCLCHECK(ncclSocketClose(&nextSocket));
NCCLCHECK(ncclSocketClose(&prevSocket));
TRACE(NCCL_INIT, "rank %d nranks %d size %d - DONE", rank, nranks, size);
return ncclSuccess;
}
// [IntraNode] in-place Broadcast
ncclResult_t bootstrapIntraNodeBroadcast(void* commState, int *ranks, int rank, int nranks, int root, void* bcastData, int size) {
if (nranks == 1) return ncclSuccess;
TRACE(NCCL_INIT, "rank %d nranks %d root %d size %d - ENTER", rank, nranks, root, size);
if (rank == root) {
for (int i=0; i<nranks; i++) {
if (i != root) NCCLCHECK(bootstrapSend(commState, ranks ? ranks[i] : i, /*tag=*/ranks ? ranks[i] : i, bcastData, size));
}
}
else {
NCCLCHECK(bootstrapRecv(commState, ranks ? ranks[root] : root, /*tag=*/ranks ? ranks[rank] : rank, bcastData, size));
}
TRACE(NCCL_INIT, "rank %d nranks %d root %d size %d - DONE", rank, nranks, root, size);
return ncclSuccess;
}
ncclResult_t bootstrapBroadcast(void* commState, int rank, int nranks, int root, void* bcastData, int size) {
return bootstrapIntraNodeBroadcast(commState, NULL, rank, nranks, root, bcastData, size);
2018-12-13 15:56:12 -08:00
}
ncclResult_t bootstrapClose(void* commState) {
2022-01-07 06:39:55 -08:00
struct bootstrapState* state = (struct bootstrapState*)commState;
2018-12-13 15:56:12 -08:00
if (state->unexpectedConnections != NULL) {
2022-11-29 04:27:46 -08:00
unexpectedFree(state);
2024-06-11 01:28:01 -07:00
if (__atomic_load_n(state->abortFlag, __ATOMIC_ACQUIRE) == 0) {
2022-11-29 04:27:46 -08:00
WARN("Unexpected connections are not empty");
return ncclInternalError;
}
2018-12-13 15:56:12 -08:00
}
2022-11-29 04:27:46 -08:00
NCCLCHECK(ncclSocketClose(&state->listenSock));
NCCLCHECK(ncclSocketClose(&state->ringSendSocket));
NCCLCHECK(ncclSocketClose(&state->ringRecvSocket));
2018-09-24 16:06:59 -07:00
2020-09-04 14:35:05 -07:00
free(state->peerCommAddresses);
2018-09-24 16:06:59 -07:00
free(state);
return ncclSuccess;
}
2019-11-19 14:57:39 -08:00
ncclResult_t bootstrapAbort(void* commState) {
2022-01-07 06:39:55 -08:00
struct bootstrapState* state = (struct bootstrapState*)commState;
2021-04-26 14:24:50 -07:00
if (commState == NULL) return ncclSuccess;
2022-11-29 04:27:46 -08:00
NCCLCHECK(ncclSocketClose(&state->listenSock));
NCCLCHECK(ncclSocketClose(&state->ringSendSocket));
NCCLCHECK(ncclSocketClose(&state->ringRecvSocket));
2020-09-04 14:35:05 -07:00
free(state->peerCommAddresses);
2022-01-07 06:39:55 -08:00
free(state->peerProxyAddresses);
2024-02-05 05:06:02 -08:00
free(state->peerProxyAddressesUDS);
2019-11-19 14:57:39 -08:00
free(state);
return ncclSuccess;
}