From ac1e37b6593dd8f894d519c8dca956fce871e2b0 Mon Sep 17 00:00:00 2001 From: Aleksa Gordic Date: Mon, 24 Jun 2024 18:56:48 +0000 Subject: [PATCH] Remove linux socket funcs when windows --- llmc/zero.cuh | 225 +++++++++++++++++++++++++------------------------- 1 file changed, 113 insertions(+), 112 deletions(-) diff --git a/llmc/zero.cuh b/llmc/zero.cuh index df2df6a..a46ed75 100644 --- a/llmc/zero.cuh +++ b/llmc/zero.cuh @@ -83,15 +83,6 @@ typedef struct { } MultiGpuConfig; #ifdef MULTI_GPU -void send_nccl_id_to_clients(ncclUniqueId *nccl_id, int client_sockets[], int num_clients) { - for (int i = 0; i < num_clients; ++i) { - if (send(client_sockets[i], nccl_id, sizeof(*nccl_id), 0) == -1) { - printf("Failed to send nccl_id"); - exit(EXIT_FAILURE); - } - close(client_sockets[i]); - } -} #ifdef _WIN32 void send_nccl_id_to_clients_windows(ncclUniqueId *nccl_id, SOCKET client_sockets[], int num_clients) { @@ -104,113 +95,17 @@ void send_nccl_id_to_clients_windows(ncclUniqueId *nccl_id, SOCKET client_socket closesocket(client_sockets[i]); } } -#endif - -ncclUniqueId get_nccl_id_via_tcp(MultiGpuConfig* result, const char* server_ip) { - ncclUniqueId nccl_id; - - int SERVER_PORT = 12345; // hardcoded an arbitrary port number between 1024 and 49151 (registered ports) - if (result->process_rank == 0) { - ncclCheck(ncclGetUniqueId(&nccl_id)); - - int MAX_CLIENTS = result->num_processes - 1; - int client_sockets[MAX_CLIENTS]; - int num_clients = 0; - int server_socket, new_socket; - struct sockaddr_in address; - int addrlen = sizeof(address); - int opt = 1; - - // Step 1) create a server TCP socket - if ((server_socket = socket(AF_INET, SOCK_STREAM, 0)) < 0) { - printf("Socket failed"); +#else +void send_nccl_id_to_clients(ncclUniqueId *nccl_id, int client_sockets[], int num_clients) { + for (int i = 0; i < num_clients; ++i) { + if (send(client_sockets[i], nccl_id, sizeof(*nccl_id), 0) == -1) { + printf("Failed to send nccl_id"); exit(EXIT_FAILURE); } - - // Step 2) set socket options - // SOL_SOCKET - means that option is configured at socket level - // SO_REUSEADDR - allows to bind to an address which is in a TIME_WAIT state (already used by another socket) - useful when restarting the server - // SO_REUSEPORT - allows to bind to the same port multiple times - if (setsockopt(server_socket, SOL_SOCKET, SO_REUSEADDR | SO_REUSEPORT, &opt, sizeof(opt)) < 0) { - printf("Setsockopt failed"); - exit(EXIT_FAILURE); - } - - // Step 3) set the server address and port - address.sin_family = AF_INET; // IPv4 - address.sin_addr.s_addr = inet_addr(server_ip); // alternatively use INADDR_ANY to bind to all interfaces, currently we only allow ethernet - address.sin_port = htons(SERVER_PORT); - - // Step 4) bind the socket to the address and port - if (bind(server_socket, (struct sockaddr *)&address, sizeof(address)) < 0) { - printf("Bind failed"); - exit(EXIT_FAILURE); - } - - // Step 5) MAX_CLIENTS specifies the maximum number of clients that can be queued for this server - if (listen(server_socket, MAX_CLIENTS) < 0) { - printf("Listen failed"); - exit(EXIT_FAILURE); - } - - // Step 6) accept connections from clients - printf("Waiting for clients to connect...\n"); - while (num_clients < MAX_CLIENTS) { - if ((new_socket = accept(server_socket, (struct sockaddr *)&address, (socklen_t*)&addrlen)) < 0) { - printf("Accept failed"); - exit(EXIT_FAILURE); - } - client_sockets[num_clients++] = new_socket; - printf("Client %d connected\n", num_clients); - } - - // Step 7) send the NCCL ID to all clients - send_nccl_id_to_clients(&nccl_id, client_sockets, num_clients); - printf("NCCL ID sent to all clients\n"); - - close(server_socket); - } else { - int num_connection_attempts = 5; - int time_to_sleep = 2; - int client_socket; - struct sockaddr_in serv_addr; - - // Step 1) create a client TCP socket - if ((client_socket = socket(AF_INET, SOCK_STREAM, 0)) < 0) { - printf("Socket creation error"); - exit(EXIT_FAILURE); - } - - // Step 2) set the server address and port - serv_addr.sin_family = AF_INET; - serv_addr.sin_port = htons(SERVER_PORT); - if (inet_pton(AF_INET, server_ip, &serv_addr.sin_addr) <= 0) { - printf("Invalid address or address not supported"); - exit(EXIT_FAILURE); - } - - // Step 3) Try to connect to the server - retry up to `num_connection_attempts` times if the connection fails - while (connect(client_socket, (struct sockaddr *)&serv_addr, sizeof(serv_addr)) < 0) { - printf("%d Connection failed, retrying in %d seconds\n", result->process_rank, time_to_sleep); - if (--num_connection_attempts == 0) { - printf("Failed to connect to the server\n"); - exit(EXIT_FAILURE); - } - sleep(time_to_sleep); - } - - // Step 4) receive the NCCL ID from the server - if (recv(client_socket, &nccl_id, sizeof(nccl_id), 0) <= 0) { - printf("Failed to receive nccl_id"); - exit(EXIT_FAILURE); - } - - printf("Received NCCL ID\n"); - close(client_socket); + close(client_sockets[i]); } - - return nccl_id; } +#endif #ifdef _WIN32 // Same as get_nccl_id_via_tcp but for Windows @@ -330,6 +225,112 @@ ncclUniqueId get_nccl_id_via_tcp_windows(MultiGpuConfig* result, const char* ser WSACleanup(); return nccl_id; } +#else +ncclUniqueId get_nccl_id_via_tcp(MultiGpuConfig* result, const char* server_ip) { + ncclUniqueId nccl_id; + + int SERVER_PORT = 12345; // hardcoded an arbitrary port number between 1024 and 49151 (registered ports) + if (result->process_rank == 0) { + ncclCheck(ncclGetUniqueId(&nccl_id)); + + int MAX_CLIENTS = result->num_processes - 1; + int client_sockets[MAX_CLIENTS]; + int num_clients = 0; + int server_socket, new_socket; + struct sockaddr_in address; + int addrlen = sizeof(address); + int opt = 1; + + // Step 1) create a server TCP socket + if ((server_socket = socket(AF_INET, SOCK_STREAM, 0)) < 0) { + printf("Socket failed"); + exit(EXIT_FAILURE); + } + + // Step 2) set socket options + // SOL_SOCKET - means that option is configured at socket level + // SO_REUSEADDR - allows to bind to an address which is in a TIME_WAIT state (already used by another socket) - useful when restarting the server + // SO_REUSEPORT - allows to bind to the same port multiple times + if (setsockopt(server_socket, SOL_SOCKET, SO_REUSEADDR | SO_REUSEPORT, &opt, sizeof(opt)) < 0) { + printf("Setsockopt failed"); + exit(EXIT_FAILURE); + } + + // Step 3) set the server address and port + address.sin_family = AF_INET; // IPv4 + address.sin_addr.s_addr = inet_addr(server_ip); // alternatively use INADDR_ANY to bind to all interfaces, currently we only allow ethernet + address.sin_port = htons(SERVER_PORT); + + // Step 4) bind the socket to the address and port + if (bind(server_socket, (struct sockaddr *)&address, sizeof(address)) < 0) { + printf("Bind failed"); + exit(EXIT_FAILURE); + } + + // Step 5) MAX_CLIENTS specifies the maximum number of clients that can be queued for this server + if (listen(server_socket, MAX_CLIENTS) < 0) { + printf("Listen failed"); + exit(EXIT_FAILURE); + } + + // Step 6) accept connections from clients + printf("Waiting for clients to connect...\n"); + while (num_clients < MAX_CLIENTS) { + if ((new_socket = accept(server_socket, (struct sockaddr *)&address, (socklen_t*)&addrlen)) < 0) { + printf("Accept failed"); + exit(EXIT_FAILURE); + } + client_sockets[num_clients++] = new_socket; + printf("Client %d connected\n", num_clients); + } + + // Step 7) send the NCCL ID to all clients + send_nccl_id_to_clients(&nccl_id, client_sockets, num_clients); + printf("NCCL ID sent to all clients\n"); + + close(server_socket); + } else { + int num_connection_attempts = 5; + int time_to_sleep = 2; + int client_socket; + struct sockaddr_in serv_addr; + + // Step 1) create a client TCP socket + if ((client_socket = socket(AF_INET, SOCK_STREAM, 0)) < 0) { + printf("Socket creation error"); + exit(EXIT_FAILURE); + } + + // Step 2) set the server address and port + serv_addr.sin_family = AF_INET; + serv_addr.sin_port = htons(SERVER_PORT); + if (inet_pton(AF_INET, server_ip, &serv_addr.sin_addr) <= 0) { + printf("Invalid address or address not supported"); + exit(EXIT_FAILURE); + } + + // Step 3) Try to connect to the server - retry up to `num_connection_attempts` times if the connection fails + while (connect(client_socket, (struct sockaddr *)&serv_addr, sizeof(serv_addr)) < 0) { + printf("%d Connection failed, retrying in %d seconds\n", result->process_rank, time_to_sleep); + if (--num_connection_attempts == 0) { + printf("Failed to connect to the server\n"); + exit(EXIT_FAILURE); + } + sleep(time_to_sleep); + } + + // Step 4) receive the NCCL ID from the server + if (recv(client_socket, &nccl_id, sizeof(nccl_id), 0) <= 0) { + printf("Failed to receive nccl_id"); + exit(EXIT_FAILURE); + } + + printf("Received NCCL ID\n"); + close(client_socket); + } + + return nccl_id; +} #endif ncclUniqueId get_nccl_id_via_fs(MultiGpuConfig* result, char* fs_path) {