feat: Add configurable network port options
This commit is contained in:
+20
-9
@@ -3436,6 +3436,9 @@ struct llama_context {
|
||||
// sockets
|
||||
std::string master_ip = "localhost";
|
||||
std::string next_node_ip = "localhost";
|
||||
uint32_t next_node_data_port = 9000;
|
||||
uint32_t next_node_signal_port = 10000;
|
||||
uint32_t master_data_port = 9000;
|
||||
uint32_t data_port = 9000;
|
||||
uint32_t signal_port = 10000;
|
||||
zmq::context_t * sock_context = nullptr;
|
||||
@@ -20237,6 +20240,11 @@ struct llama_context_params llama_context_default_params() {
|
||||
/*.keep_out_in_metal =*/ true,
|
||||
/*.master_ip =*/ nullptr,
|
||||
/*.next_node_ip =*/ nullptr,
|
||||
/*.data_port =*/ 9000,
|
||||
/*.signal_port =*/ 10000,
|
||||
/*.master_data_port =*/ 9000,
|
||||
/*.next_node_data_port =*/ 9000,
|
||||
/*.next_node_signal_port =*/ 10000,
|
||||
/*.n_ctx =*/ 512,
|
||||
/*.n_predict =*/ 512,
|
||||
/*.n_batch =*/ 2048,
|
||||
@@ -20439,11 +20447,10 @@ void llama_init_sockets(struct llama_context * ctx, uint32_t n_world, uint32_t m
|
||||
ctx->master_socket = ctx->send_socket; // No need for master_socket in the last rank, reuse send_socket for communication
|
||||
}
|
||||
|
||||
const uint32_t next_rank = (my_rank + 1) % n_world;
|
||||
std::string recv_endp = "tcp://*:" + std::to_string(map_rank_to_port(my_rank, ctx->data_port));
|
||||
std::string send_endp = "tcp://" + ctx->next_node_ip + ":" + std::to_string(map_rank_to_port(next_rank, ctx->data_port));
|
||||
std::string master_endp = "tcp://" + ctx->master_ip + ":" + std::to_string(map_rank_to_port(0, ctx->data_port));
|
||||
std::string signal_endp = "tcp://*:" + std::to_string(map_rank_to_port(my_rank, ctx->signal_port));
|
||||
std::string recv_endp = "tcp://*:" + std::to_string(ctx->data_port);
|
||||
std::string send_endp = "tcp://" + ctx->next_node_ip + ":" + std::to_string(ctx->next_node_data_port);
|
||||
std::string master_endp = "tcp://" + ctx->master_ip + ":" + std::to_string(ctx->master_data_port);
|
||||
std::string signal_endp = "tcp://*:" + std::to_string(ctx->signal_port);
|
||||
|
||||
try {
|
||||
ctx->recv_socket->bind(recv_endp);
|
||||
@@ -20653,7 +20660,7 @@ int llama_rebuild_topo(llama_context * ctx, uint32_t * n_layer_window, device_in
|
||||
if ((my_rank + 1) % n_world != next_rank) {
|
||||
socket_to_close = ctx->send_socket;
|
||||
ctx->send_socket = new zmq::socket_t(*ctx->sock_context, zmq::socket_type::push);
|
||||
std::string send_endp = "tcp://" + next_ip + ":" + std::to_string(map_rank_to_port(next_rank, ctx->data_port));
|
||||
std::string send_endp = "tcp://" + next_ip + ":" + std::to_string(ctx->next_node_data_port);
|
||||
ctx->send_socket->connect(send_endp);
|
||||
ctx->next_node_ip = next_ip;
|
||||
ctx->cparams.original_next_rank = next_rank;
|
||||
@@ -20723,15 +20730,13 @@ int llama_recv_layer_setup(struct llama_context * ctx, uint32_t * n_layer_window
|
||||
void llama_free_sockets(struct llama_context * ctx, char ** msg) {
|
||||
const uint32_t n_world = ctx->cparams.n_world;
|
||||
const uint32_t my_rank = ctx->cparams.rank;
|
||||
// to adapt to the new topology, use old next_rank
|
||||
const uint32_t next_rank = ctx->cparams.original_next_rank;
|
||||
|
||||
if (n_world == 1) {
|
||||
return;
|
||||
}
|
||||
|
||||
zmq::socket_t signal_sender(*ctx->sock_context, zmq::socket_type::push);
|
||||
std::string endp = "tcp://" + ctx->next_node_ip + ":" + std::to_string(map_rank_to_port(next_rank, ctx->signal_port));
|
||||
std::string endp = "tcp://" + ctx->next_node_ip + ":" + std::to_string(ctx->next_node_signal_port);
|
||||
signal_sender.connect(endp);
|
||||
|
||||
if (my_rank == 0) {
|
||||
@@ -20770,6 +20775,12 @@ struct llama_context * llama_new_context_with_model(
|
||||
|
||||
ctx->master_ip = params.master_ip;
|
||||
ctx->next_node_ip = params.next_node_ip;
|
||||
ctx->next_node_data_port = params.next_node_data_port;
|
||||
ctx->next_node_signal_port = params.next_node_signal_port;
|
||||
ctx->master_data_port = params.master_data_port;
|
||||
ctx->data_port = params.data_port;
|
||||
ctx->signal_port = params.signal_port;
|
||||
|
||||
ctx->cparams.n_world = params.n_world;
|
||||
ctx->cparams.rank = params.rank;
|
||||
ctx->cparams.force = params.force;
|
||||
|
||||
Reference in New Issue
Block a user