diff --git a/ggml/include/ggml-rpc.h b/ggml/include/ggml-rpc.h index cbfe400139cf..1f8cb7906cb2 100644 --- a/ggml/include/ggml-rpc.h +++ b/ggml/include/ggml-rpc.h @@ -6,7 +6,7 @@ extern "C" { #endif -#define RPC_PROTO_MAJOR_VERSION 6 +#define RPC_PROTO_MAJOR_VERSION 7 #define RPC_PROTO_MINOR_VERSION 0 #define RPC_PROTO_PATCH_VERSION 0 diff --git a/ggml/src/ggml-backend-meta.cpp b/ggml/src/ggml-backend-meta.cpp index 3ec40fb1af7f..4a7189bdf47a 100644 --- a/ggml/src/ggml-backend-meta.cpp +++ b/ggml/src/ggml-backend-meta.cpp @@ -865,7 +865,12 @@ static struct ggml_backend_meta_split_state ggml_backend_meta_get_split_state( ggml_backend_meta_split_state split_state; switch (tensor->op) { case GGML_OP_NONE: { - split_state = {GGML_BACKEND_SPLIT_AXIS_MIRRORED, {0}, {1}, 1}; + if (tensor->view_src != nullptr) { + // full-tensor view created with ggml_view_tensor, transparent for the split state + split_state = ggml_backend_meta_get_split_state(stc, tensor->view_src, assume_sync); + } else { + split_state = {GGML_BACKEND_SPLIT_AXIS_MIRRORED, {0}, {1}, 1}; + } } break; case GGML_OP_DUP: { split_state = handle_generic(src_ss, /*scalar_only =*/ true); @@ -2283,6 +2288,14 @@ static enum ggml_status ggml_backend_meta_graph_compute(ggml_backend_t backend, cgraph_ij->uid = ggml_graph_next_uid(); } } + + // Aux graph contents are rewritten on every compute but are identical across calls while the subgraphs are reused, + // so they can get stable uids on rebuild. Only safe without a comm backend, where the fallback usage is deterministic. + if (backend_ctx->comm_ctx == nullptr) { + for (ggml_cgraph * cgraph_aux : backend_ctx->cgraphs_aux) { + cgraph_aux->uid = ggml_graph_next_uid(); + } + } } size_t iga = 0; // i graph aux diff --git a/ggml/src/ggml-rpc/ggml-rpc.cpp b/ggml/src/ggml-rpc/ggml-rpc.cpp index 9aa5883d80de..e54fdadb62d8 100644 --- a/ggml/src/ggml-rpc/ggml-rpc.cpp +++ b/ggml/src/ggml-rpc/ggml-rpc.cpp @@ -5,9 +5,11 @@ #include "transport.h" #include +#include #include #include #include +#include #include #include #include @@ -21,7 +23,6 @@ #include #include #include -#include static const char * RPC_DEBUG = std::getenv("GGML_RPC_DEBUG"); @@ -77,6 +78,11 @@ enum rpc_cmd { RPC_CMD_DEVICE_COUNT, RPC_CMD_GRAPH_RECOMPUTE, RPC_CMD_MEMSET_TENSOR, + RPC_CMD_SET_TENSOR_2D, + RPC_CMD_GET_TENSOR_2D, + RPC_CMD_COMM_INIT, + RPC_CMD_COMM_ALLREDUCE, + RPC_CMD_COMM_FREE, RPC_CMD_NONE, RPC_CMD_COUNT, }; @@ -86,6 +92,10 @@ static_assert(RPC_CMD_HELLO == 14, "RPC_CMD_HELLO must be always 14"); // Try RPC_CMD_SET_TENSOR_HASH first when data size is larger than this threshold const size_t HASH_THRESHOLD = 10 * 1024 * 1024; +// Maximum number of graphs cached per device; client and server must use the same value +// so that both sides clear their caches at the same point in the message stream +const size_t GRAPH_CACHE_MAX = 1024; + struct rpc_msg_hello_req { uint8_t conn_caps[RPC_CONN_CAPS_SIZE]; }; @@ -202,6 +212,36 @@ struct rpc_msg_get_device_memory_rsp { struct rpc_msg_graph_recompute_req { uint32_t device; + uint64_t uid; +}; + +struct rpc_msg_get_tensor_2d_req { + rpc_tensor tensor; + uint64_t offset; + uint64_t size; + uint64_t n_copies; + uint64_t stride; +}; + +struct rpc_msg_comm_init_req { + uint32_t device; + uint32_t rank; + uint32_t world; + uint32_t port; // rank 0: port to listen on; rank > 0: rank 0's comm port + char host[64]; // rank > 0: rank 0's host +}; + +struct rpc_msg_comm_init_rsp { + uint8_t ok; +}; + +struct rpc_msg_comm_allreduce_req { + uint32_t device; + rpc_tensor tensor; +}; + +struct rpc_msg_comm_free_req { + uint32_t device; }; #pragma pack(pop) @@ -218,7 +258,8 @@ struct ggml_backend_rpc_device_context { uint32_t device; std::string name; std::string description; - uint64_t last_graph_uid; + // uids of graphs cached by the server for this device + std::unordered_set graph_uids; }; struct ggml_backend_rpc_buffer_type_context { @@ -232,6 +273,7 @@ struct ggml_backend_rpc_buffer_type_context { class rpc_dispatcher; struct ggml_backend_rpc_context { std::shared_ptr dispatcher; + std::string endpoint; uint32_t device; std::string name; }; @@ -341,6 +383,14 @@ static bool send_rpc_cmd(socket_ptr sock, enum rpc_cmd cmd, const void * input, // RPC client-side implementation +static inline void rpc_cpu_relax() { +#if defined(__aarch64__) && (defined(__clang__) || defined(__GNUC__)) + __asm__ volatile("yield" ::: "memory"); +#else + std::this_thread::yield(); +#endif +} + // Performs HELLO handshake with transport auto-negotiation. // Advertises local capabilities via conn_caps; if the server responds with // matching capabilities, the socket is upgraded transparently. @@ -389,6 +439,16 @@ class message_queue { return true; } + bool try_pop(T* out) { + std::unique_lock lock(mutex); + if (interrupted || queue.empty()) { + return false; + } + *out = queue.front(); + queue.pop(); + return true; + } + void interrupt() { std::unique_lock lock(mutex); interrupted = true; @@ -418,6 +478,8 @@ class rpc_dispatcher { void event_synchronize(ggml_backend_event_t event); void event_record(ggml_backend_event_t event); void synchronize(); + void busy_spin_acquire(); + void busy_spin_release(); void start(const std::string & endpoint); void work(); @@ -441,6 +503,7 @@ class rpc_dispatcher { }; rpc_msg_queue queue; socket_ptr sock; + std::atomic_uint busy_spin_users = 0; std::atomic_bool running; std::thread thread; }; @@ -532,6 +595,15 @@ void rpc_dispatcher::synchronize() { msg->completion.get_future().wait(); } +void rpc_dispatcher::busy_spin_acquire() { + busy_spin_users.fetch_add(1, std::memory_order_relaxed); +} + +void rpc_dispatcher::busy_spin_release() { + const unsigned previous = busy_spin_users.fetch_sub(1, std::memory_order_relaxed); + GGML_ASSERT(previous > 0); +} + void rpc_dispatcher::start(const std::string & endpoint) { std::string host; int port; @@ -557,7 +629,12 @@ void rpc_dispatcher::start(const std::string & endpoint) { void rpc_dispatcher::work() { while (running) { rpc_msg_ptr msg_ptr; - if (!queue.pop(&msg_ptr)) { + if (busy_spin_users.load(std::memory_order_relaxed) != 0) { + if (!queue.try_pop(&msg_ptr)) { + rpc_cpu_relax(); + continue; + } + } else if (!queue.pop(&msg_ptr)) { break; } if (msg_ptr->cmd != RPC_CMD_NONE) { @@ -716,6 +793,46 @@ static void ggml_backend_rpc_buffer_set_tensor(ggml_backend_buffer_t buffer, ggm ctx->dispatcher->send(RPC_CMD_SET_TENSOR, input_ptr, input_size); } +static void ggml_backend_rpc_buffer_set_tensor_2d(ggml_backend_buffer_t buffer, ggml_tensor * tensor, const void * data, + size_t offset, size_t size, size_t n_copies, size_t stride_tensor, size_t stride_data) { + ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context; + rpc_tensor rpc_tensor = serialize_tensor(tensor); + // input serialization format: | rpc_tensor | offset (8 bytes) | size (8 bytes) | n_copies (8 bytes) | stride (8 bytes) | data (size * n_copies bytes) | + size_t input_size = sizeof(rpc_tensor) + 4*sizeof(uint64_t) + size*n_copies; + uint8_t * input = new uint8_t[input_size](); + uint8_t * dest = input; + memcpy(dest, &rpc_tensor, sizeof(rpc_tensor)); + dest += sizeof(rpc_tensor); + uint64_t header[4] = { offset, size, n_copies, stride_tensor }; + memcpy(dest, header, sizeof(header)); + dest += sizeof(header); + for (size_t i = 0; i < n_copies; i++) { + memcpy(dest + i*size, (const char *)data + i*stride_data, size); + } + std::shared_ptr input_ptr(input, std::default_delete()); + ctx->dispatcher->send(RPC_CMD_SET_TENSOR_2D, input_ptr, input_size); +} + +static void ggml_backend_rpc_buffer_get_tensor_2d(ggml_backend_buffer_t buffer, const ggml_tensor * tensor, void * data, + size_t offset, size_t size, size_t n_copies, size_t stride_tensor, size_t stride_data) { + ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context; + auto request = std::make_shared(); + request->tensor = serialize_tensor(tensor); + request->offset = offset; + request->size = size; + request->n_copies = n_copies; + request->stride = stride_tensor; + if (stride_data == size) { + ctx->dispatcher->send(RPC_CMD_GET_TENSOR_2D, request, sizeof(*request), data, size*n_copies); + } else { + std::vector packed(size*n_copies); + ctx->dispatcher->send(RPC_CMD_GET_TENSOR_2D, request, sizeof(*request), packed.data(), packed.size()); + for (size_t i = 0; i < n_copies; i++) { + memcpy((char *)data + i*stride_data, packed.data() + i*size, size); + } + } +} + static void ggml_backend_rpc_buffer_get_tensor(ggml_backend_buffer_t buffer, const ggml_tensor * tensor, void * data, size_t offset, size_t size) { ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context; auto request = std::make_shared(); @@ -761,8 +878,8 @@ static ggml_backend_buffer_i ggml_backend_rpc_buffer_interface = { /* .memset_tensor = */ ggml_backend_rpc_buffer_memset_tensor, /* .set_tensor = */ ggml_backend_rpc_buffer_set_tensor, /* .get_tensor = */ ggml_backend_rpc_buffer_get_tensor, - /* .set_tensor_2d = */ NULL, - /* .get_tensor_2d = */ NULL, + /* .set_tensor_2d = */ ggml_backend_rpc_buffer_set_tensor_2d, + /* .get_tensor_2d = */ ggml_backend_rpc_buffer_get_tensor_2d, /* .cpy_tensor = */ ggml_backend_rpc_buffer_cpy_tensor, /* .clear = */ ggml_backend_rpc_buffer_clear, /* .reset = */ NULL, @@ -986,13 +1103,15 @@ static uint8_t * serialize_graph(uint32_t device, const ggml_cgraph * cgraph, si add_tensor(cgraph->nodes[i], cgraph, tensors, visited); } // serialization format: - // | device (4 bytes) | n_nodes (4 bytes) | nodes (n_nodes * sizeof(uint64_t) | n_tensors (4 bytes) | tensors (n_tensors * sizeof(rpc_tensor)) | + // | device (4 bytes) | uid (8 bytes) | n_nodes (4 bytes) | nodes (n_nodes * sizeof(uint64_t) | n_tensors (4 bytes) | tensors (n_tensors * sizeof(rpc_tensor)) | uint32_t n_tensors = tensors.size(); - *output_size = 2*sizeof(uint32_t) + n_nodes * sizeof(uint64_t) + sizeof(uint32_t) + n_tensors * sizeof(rpc_tensor); + *output_size = 2*sizeof(uint32_t) + sizeof(uint64_t) + n_nodes * sizeof(uint64_t) + sizeof(uint32_t) + n_tensors * sizeof(rpc_tensor); uint8_t * output = new uint8_t[*output_size](); uint8_t * dest = output; memcpy(dest, &device, sizeof(device)); dest += sizeof(device); + memcpy(dest, &cgraph->uid, sizeof(cgraph->uid)); + dest += sizeof(cgraph->uid); memcpy(dest, &n_nodes, sizeof(n_nodes)); dest += sizeof(n_nodes); for (uint32_t i = 0; i < n_nodes; i++) { @@ -1012,13 +1131,20 @@ static enum ggml_status ggml_backend_rpc_graph_compute(ggml_backend_t backend, g ggml_backend_rpc_device_context * rpc_dev_ctx = (ggml_backend_rpc_device_context *)rpc_dev->context; GGML_ASSERT(cgraph->n_nodes > 0); - bool reuse = cgraph->uid != 0 && rpc_dev_ctx->last_graph_uid == cgraph->uid; + auto & graph_uids = rpc_dev_ctx->graph_uids; + bool reuse = cgraph->uid != 0 && graph_uids.count(cgraph->uid) > 0; if (reuse) { auto request = std::make_shared(); request->device = rpc_ctx->device; + request->uid = cgraph->uid; rpc_ctx->dispatcher->send_async(RPC_CMD_GRAPH_RECOMPUTE, request, sizeof(*request)); } else { - rpc_dev_ctx->last_graph_uid = cgraph->uid; + if (cgraph->uid != 0) { + if (graph_uids.size() >= GRAPH_CACHE_MAX) { + graph_uids.clear(); + } + graph_uids.insert(cgraph->uid); + } size_t input_size = 0; uint8_t * input = serialize_graph(rpc_ctx->device, cgraph, &input_size); std::shared_ptr input_ptr(input, std::default_delete()); @@ -1038,13 +1164,25 @@ static void ggml_backend_rpc_event_wait(ggml_backend_t backend, ggml_backend_eve GGML_UNUSED(event); } +static void ggml_backend_rpc_set_tensor_2d_async(ggml_backend_t backend, ggml_tensor * tensor, const void * data, + size_t offset, size_t size, size_t n_copies, size_t stride_tensor, size_t stride_data) { + ggml_backend_tensor_set_2d(tensor, data, offset, size, n_copies, stride_tensor, stride_data); + GGML_UNUSED(backend); +} + +static void ggml_backend_rpc_get_tensor_2d_async(ggml_backend_t backend, const ggml_tensor * tensor, void * data, + size_t offset, size_t size, size_t n_copies, size_t stride_tensor, size_t stride_data) { + ggml_backend_tensor_get_2d(tensor, data, offset, size, n_copies, stride_tensor, stride_data); + GGML_UNUSED(backend); +} + static ggml_backend_i ggml_backend_rpc_interface = { /* .get_name = */ ggml_backend_rpc_name, /* .free = */ ggml_backend_rpc_free, /* .set_tensor_async = */ ggml_backend_rpc_set_tensor_async, /* .get_tensor_async = */ ggml_backend_rpc_get_tensor_async, - /* .set_tensor_2d_async = */ NULL, - /* .get_tensor_2d_async = */ NULL, + /* .set_tensor_2d_async = */ ggml_backend_rpc_set_tensor_2d_async, + /* .get_tensor_2d_async = */ ggml_backend_rpc_get_tensor_2d_async, /* .cpy_tensor_async = */ NULL, /* .synchronize = */ ggml_backend_rpc_synchronize, /* .graph_plan_create = */ NULL, @@ -1092,6 +1230,7 @@ ggml_backend_t ggml_backend_rpc_init(const char * endpoint, uint32_t device) { auto dispatcher = get_dispatcher(endpoint); ggml_backend_rpc_context * ctx = new ggml_backend_rpc_context { /* .dispatcher = */ dispatcher, + /* .endpoint = */ endpoint, /* .device = */ device, /* .name = */ dev_name, }; @@ -1126,6 +1265,7 @@ class rpc_server { rpc_server(std::vector all_backends, const char * cache_dir) : backends(std::move(all_backends)), cache_dir(cache_dir) { stored_graphs.resize(backends.size()); + comm_states.resize(backends.size()); } ~rpc_server(); @@ -1138,11 +1278,16 @@ class rpc_server { bool buffer_clear(const rpc_msg_buffer_clear_req & request); bool memset_tensor(const rpc_msg_memset_tensor_req & request); bool set_tensor(const std::vector & input); + bool set_tensor_2d(const std::vector & input); bool set_tensor_hash(const rpc_msg_set_tensor_hash_req & request, rpc_msg_set_tensor_hash_rsp & response); bool get_tensor(const rpc_msg_get_tensor_req & request, std::vector & response); + bool get_tensor_2d(const rpc_msg_get_tensor_2d_req & request, std::vector & response); bool copy_tensor(const rpc_msg_copy_tensor_req & request, rpc_msg_copy_tensor_rsp & response); bool graph_compute(const std::vector & input); bool graph_recompute(const rpc_msg_graph_recompute_req & request); + bool comm_init(const rpc_msg_comm_init_req & request, rpc_msg_comm_init_rsp & response); + bool comm_allreduce(const rpc_msg_comm_allreduce_req & request); + bool comm_free(const rpc_msg_comm_free_req & request); bool init_tensor(const rpc_msg_init_tensor_req & request); bool get_alloc_size(const rpc_msg_get_alloc_size_req & request, rpc_msg_get_alloc_size_rsp & response); bool get_device_memory(const rpc_msg_get_device_memory_req & request, rpc_msg_get_device_memory_rsp & response); @@ -1153,6 +1298,7 @@ class rpc_server { }; private: + void sync_all_backends(); bool get_cached_file(uint64_t hash, std::vector & data); ggml_tensor * deserialize_tensor(struct ggml_context * ctx, const rpc_tensor * tensor); ggml_tensor * create_node(uint64_t id, @@ -1161,11 +1307,23 @@ class rpc_server { std::unordered_map & tensor_map); + // pairwise allreduce over a direct connection to the peer server + struct comm_state { + socket_ptr peer; + uint32_t rank = 0; + uint32_t world = 0; + ggml_backend_buffer_ptr scratch; + size_t scratch_size = 0; + std::vector send_buf; + std::vector recv_buf; + }; + std::vector backends; const char * cache_dir; std::unordered_set buffers; - // store the last computed graph for each backend - std::vector stored_graphs; + // computed graphs cached per backend, keyed by uid + std::vector> stored_graphs; + std::vector comm_states; }; void rpc_server::hello(rpc_msg_hello_rsp & response) { @@ -1273,6 +1431,7 @@ bool rpc_server::buffer_get_base(const rpc_msg_buffer_get_base_req & request, rp } bool rpc_server::free_buffer(const rpc_msg_free_buffer_req & request) { + sync_all_backends(); LOG_DBG("[%s] remote_ptr: %" PRIx64 "\n", __func__, request.remote_ptr); ggml_backend_buffer_t buffer = reinterpret_cast(request.remote_ptr); if (buffers.find(buffer) == buffers.end()) { @@ -1285,6 +1444,7 @@ bool rpc_server::free_buffer(const rpc_msg_free_buffer_req & request) { } bool rpc_server::buffer_clear(const rpc_msg_buffer_clear_req & request) { + sync_all_backends(); LOG_DBG("[%s] remote_ptr: %" PRIx64 ", value: %u\n", __func__, request.remote_ptr, request.value); ggml_backend_buffer_t buffer = reinterpret_cast(request.remote_ptr); if (buffers.find(buffer) == buffers.end()) { @@ -1296,6 +1456,7 @@ bool rpc_server::buffer_clear(const rpc_msg_buffer_clear_req & request) { } bool rpc_server::memset_tensor(const rpc_msg_memset_tensor_req & request) { + sync_all_backends(); struct ggml_init_params params { /*.mem_size =*/ ggml_tensor_overhead(), /*.mem_buffer =*/ NULL, @@ -1371,13 +1532,20 @@ ggml_tensor * rpc_server::deserialize_tensor(struct ggml_context * ctx, const rp result->buffer = nullptr; } - if (result->buffer) { + if (result->buffer && ggml_nelements(result) > 0) { // require that the tensor data does not go beyond the buffer end uint64_t tensor_size = (uint64_t) ggml_nbytes(result); uint64_t buffer_start = (uint64_t) ggml_backend_buffer_get_base(result->buffer); uint64_t buffer_size = (uint64_t) ggml_backend_buffer_get_size(result->buffer); - GGML_ASSERT(tensor->data + tensor_size >= tensor->data); // check for overflow - GGML_ASSERT(tensor->data >= buffer_start && tensor->data + tensor_size <= buffer_start + buffer_size); + if (tensor->data + tensor_size < tensor->data || + tensor->data < buffer_start || tensor->data + tensor_size > buffer_start + buffer_size) { + GGML_LOG_ERROR("[%s] tensor '%s' (op %s, type %s, ne [%" PRId64 ", %" PRId64 ", %" PRId64 ", %" PRId64 "]) " + "data [0x%" PRIx64 ", 0x%" PRIx64 ") out of buffer bounds [0x%" PRIx64 ", 0x%" PRIx64 ")\n", + __func__, tensor->name, ggml_op_name((ggml_op) tensor->op), ggml_type_name(result->type), + result->ne[0], result->ne[1], result->ne[2], result->ne[3], + tensor->data, tensor->data + tensor_size, buffer_start, buffer_start + buffer_size); + return nullptr; + } } result->op = (ggml_op) tensor->op; @@ -1392,6 +1560,7 @@ ggml_tensor * rpc_server::deserialize_tensor(struct ggml_context * ctx, const rp bool rpc_server::set_tensor(const std::vector & input) { + sync_all_backends(); // serialization format: | rpc_tensor | offset (8 bytes) | data (size bytes) | if (input.size() < sizeof(rpc_tensor) + sizeof(uint64_t)) { return false; @@ -1443,6 +1612,67 @@ bool rpc_server::set_tensor(const std::vector & input) { return true; } +bool rpc_server::set_tensor_2d(const std::vector & input) { + sync_all_backends(); + // serialization format: | rpc_tensor | offset (8 bytes) | size (8 bytes) | n_copies (8 bytes) | stride (8 bytes) | data (size * n_copies bytes) | + if (input.size() < sizeof(rpc_tensor) + 4*sizeof(uint64_t)) { + return false; + } + const rpc_tensor * in_tensor = (const rpc_tensor *)input.data(); + uint64_t header[4]; + memcpy(header, input.data() + sizeof(rpc_tensor), sizeof(header)); + const uint64_t offset = header[0]; + const uint64_t size = header[1]; + const uint64_t n_copies = header[2]; + const uint64_t stride = header[3]; + + const uint64_t data_size = input.size() - sizeof(rpc_tensor) - 4*sizeof(uint64_t); + if (n_copies == 0 || size == 0 || size > data_size / n_copies || size * n_copies != data_size) { + return false; + } + + struct ggml_init_params params { + /*.mem_size =*/ ggml_tensor_overhead(), + /*.mem_buffer =*/ NULL, + /*.no_alloc =*/ true, + }; + ggml_context_ptr ctx_ptr { ggml_init(params) }; + GGML_ASSERT(ctx_ptr != nullptr); + ggml_context * ctx = ctx_ptr.get(); + ggml_tensor * tensor = deserialize_tensor(ctx, in_tensor); + if (tensor == nullptr || tensor->buffer == nullptr) { + GGML_LOG_ERROR("[%s] error deserializing tensor\n", __func__); + return false; + } + LOG_DBG("[%s] buffer: %p, data: %p, offset: %" PRIu64 ", size: %" PRIu64 ", n_copies: %" PRIu64 ", stride: %" PRIu64 "\n", + __func__, (void*)tensor->buffer, tensor->data, offset, size, n_copies, stride); + + // sanitize tensor->data + { + if (stride != 0 && n_copies - 1 > (UINT64_MAX - size) / stride) { + return false; + } + const uint64_t span = (n_copies - 1)*stride + size; + const uint64_t p0 = (uint64_t) ggml_backend_buffer_get_base(tensor->buffer); + const uint64_t p1 = p0 + ggml_backend_buffer_get_size(tensor->buffer); + + if (in_tensor->data < p0 || in_tensor->data > p1 || offset > p1 - in_tensor->data || span > p1 - in_tensor->data - offset) { + GGML_LOG_ERROR("[%s] tensor data region (data=0x%" PRIx64 ", offset=%" PRIu64 ", span=%" PRIu64 ") out of buffer bounds [0x%" PRIx64 ", 0x%" PRIx64 ")\n", + __func__, in_tensor->data, offset, span, p0, p1); + return false; + } + if (offset > ggml_nbytes(tensor) || span > ggml_nbytes(tensor) - offset) { + GGML_LOG_ERROR("[%s] tensor write region (offset=%" PRIu64 ", span=%" PRIu64 ") out of tensor bounds (%zu)\n", + __func__, offset, span, ggml_nbytes(tensor)); + return false; + } + } + + const void * data = input.data() + sizeof(rpc_tensor) + 4*sizeof(uint64_t); + ggml_backend_tensor_set_2d(tensor, data, offset, size, n_copies, stride, size); + return true; +} + bool rpc_server::get_cached_file(uint64_t hash, std::vector & data) { if (!cache_dir) { return false; @@ -1465,6 +1695,7 @@ bool rpc_server::get_cached_file(uint64_t hash, std::vector & data) { bool rpc_server::set_tensor_hash(const rpc_msg_set_tensor_hash_req & request, rpc_msg_set_tensor_hash_rsp & response) { + sync_all_backends(); std::vector cached_file; if (!get_cached_file(request.hash, cached_file)) { response.result = 0; @@ -1506,6 +1737,7 @@ bool rpc_server::set_tensor_hash(const rpc_msg_set_tensor_hash_req & request, rp } bool rpc_server::init_tensor(const rpc_msg_init_tensor_req & request) { + sync_all_backends(); struct ggml_init_params params { /*.mem_size =*/ ggml_tensor_overhead(), /*.mem_buffer =*/ NULL, @@ -1541,6 +1773,7 @@ bool rpc_server::init_tensor(const rpc_msg_init_tensor_req & request) { } bool rpc_server::get_tensor(const rpc_msg_get_tensor_req & request, std::vector & response) { + sync_all_backends(); struct ggml_init_params params { /*.mem_size =*/ ggml_tensor_overhead(), /*.mem_buffer =*/ NULL, @@ -1575,7 +1808,56 @@ bool rpc_server::get_tensor(const rpc_msg_get_tensor_req & request, std::vector< return true; } +bool rpc_server::get_tensor_2d(const rpc_msg_get_tensor_2d_req & request, std::vector & response) { + sync_all_backends(); + struct ggml_init_params params { + /*.mem_size =*/ ggml_tensor_overhead(), + /*.mem_buffer =*/ NULL, + /*.no_alloc =*/ true, + }; + ggml_context_ptr ctx_ptr { ggml_init(params) }; + GGML_ASSERT(ctx_ptr != nullptr); + ggml_context * ctx = ctx_ptr.get(); + ggml_tensor * tensor = deserialize_tensor(ctx, &request.tensor); + if (tensor == nullptr || tensor->buffer == nullptr) { + GGML_LOG_ERROR("[%s] error deserializing tensor\n", __func__); + return false; + } + LOG_DBG("[%s] buffer: %p, data: %p, offset: %" PRIu64 ", size: %" PRIu64 ", n_copies: %" PRIu64 ", stride: %" PRIu64 "\n", + __func__, (void*)tensor->buffer, tensor->data, request.offset, request.size, request.n_copies, request.stride); + + // sanitize tensor->data + { + if (request.n_copies == 0 || request.size == 0 || request.size > UINT64_MAX / request.n_copies) { + return false; + } + if (request.stride != 0 && request.n_copies - 1 > (UINT64_MAX - request.size) / request.stride) { + return false; + } + const uint64_t span = (request.n_copies - 1)*request.stride + request.size; + const uint64_t p0 = (uint64_t) ggml_backend_buffer_get_base(tensor->buffer); + const uint64_t p1 = p0 + ggml_backend_buffer_get_size(tensor->buffer); + + if (request.tensor.data < p0 || request.tensor.data > p1 || request.offset > p1 - request.tensor.data || + span > p1 - request.tensor.data - request.offset) { + GGML_LOG_ERROR("[%s] tensor data region (data=0x%" PRIx64 ", offset=%" PRIu64 ", span=%" PRIu64 ") out of buffer bounds [0x%" PRIx64 ", 0x%" PRIx64 ")\n", + __func__, request.tensor.data, request.offset, span, p0, p1); + return false; + } + if (request.offset > ggml_nbytes(tensor) || span > ggml_nbytes(tensor) - request.offset) { + GGML_LOG_ERROR("[%s] tensor read region (offset=%" PRIu64 ", span=%" PRIu64 ") out of tensor bounds (%zu)\n", + __func__, request.offset, span, ggml_nbytes(tensor)); + return false; + } + } + + response.resize(request.size * request.n_copies, 0); + ggml_backend_tensor_get_2d(tensor, response.data(), request.offset, request.size, request.n_copies, request.stride, request.size); + return true; +} + bool rpc_server::copy_tensor(const rpc_msg_copy_tensor_req & request, rpc_msg_copy_tensor_rsp & response) { + sync_all_backends(); struct ggml_init_params params { /*.mem_size =*/ 2*ggml_tensor_overhead(), /*.mem_buffer =*/ NULL, @@ -1674,8 +1956,8 @@ ggml_tensor * rpc_server::create_node(uint64_t id, bool rpc_server::graph_compute(const std::vector & input) { // serialization format: - // | device (4 bytes) | n_nodes (4 bytes) | nodes (n_nodes * sizeof(uint64_t) | n_tensors (4 bytes) | tensors (n_tensors * sizeof(rpc_tensor)) | - if (input.size() < 2*sizeof(uint32_t)) { + // | device (4 bytes) | uid (8 bytes) | n_nodes (4 bytes) | nodes (n_nodes * sizeof(uint64_t) | n_tensors (4 bytes) | tensors (n_tensors * sizeof(rpc_tensor)) | + if (input.size() < 2*sizeof(uint32_t) + sizeof(uint64_t)) { return false; } const uint8_t * src = input.data(); @@ -1685,10 +1967,13 @@ bool rpc_server::graph_compute(const std::vector & input) { if (device >= backends.size()) { return false; } + uint64_t uid; + memcpy(&uid, src, sizeof(uid)); + src += sizeof(uid); uint32_t n_nodes; memcpy(&n_nodes, src, sizeof(n_nodes)); src += sizeof(n_nodes); - if (input.size() < 2*sizeof(uint32_t) + n_nodes*sizeof(uint64_t) + sizeof(uint32_t)) { + if (input.size() < 2*sizeof(uint32_t) + sizeof(uint64_t) + n_nodes*sizeof(uint64_t) + sizeof(uint32_t)) { return false; } const uint64_t * nodes = (const uint64_t *)src; @@ -1696,19 +1981,26 @@ bool rpc_server::graph_compute(const std::vector & input) { uint32_t n_tensors; memcpy(&n_tensors, src, sizeof(n_tensors)); src += sizeof(n_tensors); - if (input.size() < 2*sizeof(uint32_t) + n_nodes*sizeof(uint64_t) + sizeof(uint32_t) + n_tensors*sizeof(rpc_tensor)) { + if (input.size() < 2*sizeof(uint32_t) + sizeof(uint64_t) + n_nodes*sizeof(uint64_t) + sizeof(uint32_t) + n_tensors*sizeof(rpc_tensor)) { return false; } const rpc_tensor * tensors = (const rpc_tensor *)src; - LOG_DBG("[%s] device: %u, n_nodes: %u, n_tensors: %u\n", __func__, device, n_nodes, n_tensors); + LOG_DBG("[%s] device: %u, uid: %" PRIu64 ", n_nodes: %u, n_tensors: %u\n", __func__, device, uid, n_nodes, n_tensors); + + // graphs with uid == 0 are not cached, see GRAPH_CACHE_MAX for the eviction policy + if (uid != 0 && stored_graphs[device].size() >= GRAPH_CACHE_MAX) { + stored_graphs[device].clear(); + } + stored_graph sg_tmp; + stored_graph & sg = uid != 0 ? stored_graphs[device][uid] : sg_tmp; size_t buf_size = ggml_tensor_overhead()*(n_nodes + n_tensors) + ggml_graph_overhead_custom(n_nodes, false); - if (stored_graphs[device].buffer.size() < buf_size) { - stored_graphs[device].buffer.resize(buf_size); + if (sg.buffer.size() < buf_size) { + sg.buffer.resize(buf_size); } struct ggml_init_params params = { /*.mem_size =*/ buf_size, - /*.mem_buffer =*/ stored_graphs[device].buffer.data(), + /*.mem_buffer =*/ sg.buffer.data(), /*.no_alloc =*/ true, }; ggml_context_ptr ctx_ptr { ggml_init(params) }; @@ -1740,9 +2032,9 @@ bool rpc_server::graph_compute(const std::vector & input) { graph->use_counts[hash_pos] = tensor_ptrs.at(id)->use_count; } } - ggml_status status = ggml_backend_graph_compute(backends[device], graph); + ggml_status status = ggml_backend_graph_compute_async(backends[device], graph); GGML_ASSERT(status == GGML_STATUS_SUCCESS && "Unsuccessful graph computations are not supported with RPC"); - stored_graphs[device].graph = graph; + sg.graph = graph; return true; } @@ -1751,16 +2043,216 @@ bool rpc_server::graph_recompute(const rpc_msg_graph_recompute_req & request) { if (device >= backends.size()) { return false; } - if (stored_graphs[device].graph == nullptr) { + auto it = stored_graphs[device].find(request.uid); + if (it == stored_graphs[device].end() || it->second.graph == nullptr) { + GGML_LOG_ERROR("[%s] device: %u, graph with uid %" PRIu64 " not found\n", __func__, device, request.uid); return false; } - ggml_cgraph * graph = stored_graphs[device].graph; - LOG_DBG("[%s] device: %u\n", __func__, device); - ggml_status status = ggml_backend_graph_compute(backends[device], graph); + ggml_cgraph * graph = it->second.graph; + LOG_DBG("[%s] device: %u, uid: %" PRIu64 "\n", __func__, device, request.uid); + ggml_status status = ggml_backend_graph_compute_async(backends[device], graph); GGML_ASSERT(status == GGML_STATUS_SUCCESS && "Unsuccessful graph computations are not supported with RPC"); return true; } +// graph compute is asynchronous; commands that read or write buffer data synchronize first +void rpc_server::sync_all_backends() { + for (ggml_backend_t backend : backends) { + ggml_backend_synchronize(backend); + } +} + +// The comm link between two servers uses the same caps negotiation as the client HELLO, +// so it gets the same transport upgrades (e.g. RDMA). +bool rpc_server::comm_init(const rpc_msg_comm_init_req & request, rpc_msg_comm_init_rsp & response) { + response.ok = 0; + if (request.device >= backends.size() || request.world != 2 || request.rank >= request.world) { + return true; + } + comm_state & state = comm_states[request.device]; + if (state.peer != nullptr) { + response.ok = 1; + return true; + } + uint8_t local_caps[RPC_CONN_CAPS_SIZE] = {}; + uint8_t remote_caps[RPC_CONN_CAPS_SIZE] = {}; + if (request.rank == 0) { + socket_ptr srv = socket_t::create_server("0.0.0.0", request.port); + if (srv == nullptr) { + GGML_LOG_ERROR("[%s] failed to listen on comm port %u\n", __func__, request.port); + return true; + } + state.peer = srv->accept(); + if (state.peer == nullptr) { + return true; + } + if (!state.peer->recv_data(remote_caps, sizeof(remote_caps))) { + state.peer = nullptr; + return true; + } + state.peer->get_caps(local_caps); + if (!state.peer->send_data(local_caps, sizeof(local_caps))) { + state.peer = nullptr; + return true; + } + state.peer->update_caps(remote_caps); + } else { + const std::string host(request.host, strnlen(request.host, sizeof(request.host))); + // rank 0 may not be listening yet, retry for a few seconds + for (int i = 0; i < 100 && state.peer == nullptr; i++) { + state.peer = socket_t::connect(host.c_str(), request.port); + if (state.peer == nullptr) { + std::this_thread::sleep_for(std::chrono::milliseconds(50)); + } + } + if (state.peer == nullptr) { + GGML_LOG_ERROR("[%s] failed to connect to peer %s:%u\n", __func__, host.c_str(), request.port); + return true; + } + state.peer->get_caps(local_caps); + if (!state.peer->send_data(local_caps, sizeof(local_caps)) || + !state.peer->recv_data(remote_caps, sizeof(remote_caps))) { + state.peer = nullptr; + return true; + } + state.peer->update_caps(remote_caps); + } + state.rank = request.rank; + state.world = request.world; + GGML_LOG_INFO("[%s] device %u joined pairwise comm as rank %u\n", __func__, request.device, request.rank); + response.ok = 1; + return true; +} + +bool rpc_server::comm_allreduce(const rpc_msg_comm_allreduce_req & request) { + if (request.device >= backends.size()) { + return false; + } + comm_state & state = comm_states[request.device]; + if (state.peer == nullptr) { + GGML_LOG_ERROR("[%s] no communicator for device %u\n", __func__, request.device); + return false; + } + ggml_backend_t backend = backends[request.device]; + + size_t ctx_size = 16*ggml_tensor_overhead() + 2*ggml_graph_overhead_custom(8, false); + struct ggml_init_params params = { + /*.mem_size =*/ ctx_size, + /*.mem_buffer =*/ NULL, + /*.no_alloc =*/ true, + }; + ggml_context_ptr ctx_ptr { ggml_init(params) }; + GGML_ASSERT(ctx_ptr != nullptr); + ggml_context * ctx = ctx_ptr.get(); + ggml_tensor * t_dst = deserialize_tensor(ctx, &request.tensor); + if (t_dst == nullptr || t_dst->buffer == nullptr) { + GGML_LOG_ERROR("[%s] error deserializing tensor\n", __func__); + return false; + } + const size_t nbytes = ggml_nbytes(t_dst); + const int64_t ne = ggml_nelements(t_dst); + if (nbytes == 0) { + return true; + } + // reduce large partials in bf16 to halve the wire bytes; small (decode-sized) ones + // stay f32 since the extra casts and sync cost more than the bytes saved + const bool wire_bf16 = t_dst->type == GGML_TYPE_F32 && ne >= 32768; + const size_t wire_bytes = wire_bf16 ? (size_t) ne*2 : nbytes; + const size_t need = wire_bf16 ? 2*nbytes : nbytes; + if (state.scratch_size < need) { + state.scratch.reset(ggml_backend_alloc_buffer(backend, need)); + state.scratch_size = need; + } + char * scratch_base = (char *) ggml_backend_buffer_get_base(state.scratch.get()); + state.send_buf.resize(wire_bytes); + state.recv_buf.resize(wire_bytes); + + auto new_scratch_tensor = [&](ggml_type type, size_t offset) { + ggml_tensor * t = ggml_new_tensor_4d(ctx, type, t_dst->ne[0], t_dst->ne[1], t_dst->ne[2], t_dst->ne[3]); + t->buffer = state.scratch.get(); + t->data = scratch_base + offset; + return t; + }; + auto new_cpy_node = [&](ggml_tensor * src, ggml_tensor * dst) { + ggml_tensor * t = ggml_new_tensor_4d(ctx, dst->type, dst->ne[0], dst->ne[1], dst->ne[2], dst->ne[3]); + t->op = GGML_OP_CPY; + t->src[0] = src; + t->src[1] = dst; + t->buffer = dst->buffer; + t->data = dst->data; + t->flags |= GGML_TENSOR_FLAG_COMPUTE; + return t; + }; + auto compute_nodes = [&](ggml_tensor * n0, ggml_tensor * n1) { + ggml_cgraph * graph = ggml_new_graph_custom(ctx, 2, false); + graph->nodes[0] = n0; + graph->nodes[1] = n1; + graph->n_nodes = n1 != nullptr ? 2 : 1; + ggml_status status = ggml_backend_graph_compute_async(backend, graph); + GGML_ASSERT(status == GGML_STATUS_SUCCESS && "Unsuccessful graph computations are not supported with RPC"); + }; + + // wait for the pending subgraph that produced this partial + ggml_backend_synchronize(backend); + + ggml_tensor * t_wire_send = nullptr; + ggml_tensor * t_wire_recv = nullptr; + if (wire_bf16) { + t_wire_send = new_scratch_tensor(GGML_TYPE_BF16, 0); + t_wire_recv = new_scratch_tensor(GGML_TYPE_BF16, ne*2); + compute_nodes(new_cpy_node(t_dst, t_wire_send), nullptr); + ggml_backend_synchronize(backend); + ggml_backend_tensor_get(t_wire_send, state.send_buf.data(), 0, wire_bytes); + } else { + ggml_backend_tensor_get(t_dst, state.send_buf.data(), 0, wire_bytes); + } + + // rank 0 sends first, rank 1 receives first, so large payloads cannot deadlock + if (state.rank == 0) { + if (!state.peer->send_data(state.send_buf.data(), wire_bytes) || !state.peer->flush() || + !state.peer->recv_data(state.recv_buf.data(), wire_bytes)) { + return false; + } + } else { + if (!state.peer->recv_data(state.recv_buf.data(), wire_bytes) || + !state.peer->send_data(state.send_buf.data(), wire_bytes) || !state.peer->flush()) { + return false; + } + } + + ggml_tensor * t_peer = new_scratch_tensor(t_dst->type, wire_bf16 ? (size_t) ne*4 : 0); + ggml_tensor * t_cast = nullptr; + if (wire_bf16) { + ggml_backend_tensor_set(t_wire_recv, state.recv_buf.data(), 0, wire_bytes); + t_cast = new_cpy_node(t_wire_recv, t_peer); + } else { + ggml_backend_tensor_set(t_peer, state.recv_buf.data(), 0, wire_bytes); + } + + ggml_tensor * t_red = ggml_new_tensor_4d(ctx, t_dst->type, t_dst->ne[0], t_dst->ne[1], t_dst->ne[2], t_dst->ne[3]); + t_red->op = GGML_OP_ADD; + t_red->src[0] = t_dst; + t_red->src[1] = t_peer; + t_red->buffer = t_dst->buffer; + t_red->data = t_dst->data; + t_red->flags |= GGML_TENSOR_FLAG_COMPUTE; + + if (t_cast != nullptr) { + compute_nodes(t_cast, t_red); + } else { + compute_nodes(t_red, nullptr); + } + return true; +} + +bool rpc_server::comm_free(const rpc_msg_comm_free_req & request) { + if (request.device >= backends.size()) { + return false; + } + comm_states[request.device] = comm_state(); + return true; +} + bool rpc_server::get_device_memory(const rpc_msg_get_device_memory_req & request, rpc_msg_get_device_memory_rsp & response) { uint32_t dev_id = request.device; if (dev_id >= backends.size()) { @@ -1955,6 +2447,30 @@ static void rpc_serve_client(const std::vector & backends, const } break; } + case RPC_CMD_SET_TENSOR_2D: { + std::vector input; + if (!recv_msg(sock, input)) { + return; + } + if (!server.set_tensor_2d(input)) { + return; + } + break; + } + case RPC_CMD_GET_TENSOR_2D: { + rpc_msg_get_tensor_2d_req request; + if (!recv_msg(sock, &request, sizeof(request))) { + return; + } + std::vector response; + if (!server.get_tensor_2d(request, response)) { + return; + } + if (!send_msg(sock, response.data(), response.size())) { + return; + } + break; + } case RPC_CMD_SET_TENSOR_HASH: { rpc_msg_set_tensor_hash_req request; if (!recv_msg(sock, &request, sizeof(request))) { @@ -2027,6 +2543,40 @@ static void rpc_serve_client(const std::vector & backends, const } break; } + case RPC_CMD_COMM_INIT: { + rpc_msg_comm_init_req request; + if (!recv_msg(sock, &request, sizeof(request))) { + return; + } + rpc_msg_comm_init_rsp response; + if (!server.comm_init(request, response)) { + return; + } + if (!send_msg(sock, &response, sizeof(response))) { + return; + } + break; + } + case RPC_CMD_COMM_ALLREDUCE: { + rpc_msg_comm_allreduce_req request; + if (!recv_msg(sock, &request, sizeof(request))) { + return; + } + if (!server.comm_allreduce(request)) { + return; + } + break; + } + case RPC_CMD_COMM_FREE: { + rpc_msg_comm_free_req request; + if (!recv_msg(sock, &request, sizeof(request))) { + return; + } + if (!server.comm_free(request)) { + return; + } + break; + } case RPC_CMD_GET_DEVICE_MEMORY: { rpc_msg_get_device_memory_req request; if (!recv_msg(sock, &request, sizeof(request))) { @@ -2256,6 +2806,136 @@ static ggml_backend_dev_t ggml_backend_rpc_reg_get_device(ggml_backend_reg_t reg } } +// Pairwise allreduce between two RPC servers over a direct server-to-server connection. +// The client only sends fire-and-forget COMM_ALLREDUCE commands; the tensor data is +// exchanged between the servers and never passes through the client. +struct ggml_backend_rpc_comm_context { + struct rank_info { + std::string endpoint; + uint32_t device; + std::shared_ptr dispatcher; + }; + std::vector ranks; +}; + +static void ggml_backend_rpc_comm_free(void * comm_ctx_v) { + ggml_backend_rpc_comm_context * comm_ctx = (ggml_backend_rpc_comm_context *) comm_ctx_v; + if (comm_ctx == nullptr) { + return; + } + for (const auto & rank : comm_ctx->ranks) { + auto request = std::make_shared(); + request->device = rank.device; + rank.dispatcher->send(RPC_CMD_COMM_FREE, request, sizeof(*request)); + rank.dispatcher->busy_spin_release(); + } + delete comm_ctx; +} + +static void * ggml_backend_rpc_comm_init(ggml_backend_t * backends, size_t n_backends) { + if (n_backends != 2 || std::getenv("GGML_RPC_NO_COMM") != nullptr) { + return nullptr; + } + std::vector ranks; + ranks.reserve(n_backends); + for (size_t i = 0; i < n_backends; i++) { + if (!ggml_backend_is_rpc(backends[i])) { + return nullptr; + } + ggml_backend_rpc_context * rpc_ctx = (ggml_backend_rpc_context *) backends[i]->context; + // one rank per endpoint: a server processes its socket sequentially, so a second + // COMM_INIT on the same connection would deadlock behind the first + for (const auto & rank : ranks) { + if (rank.endpoint == rpc_ctx->endpoint) { + GGML_LOG_WARN("%s: multiple ranks on endpoint %s are not supported\n", __func__, rpc_ctx->endpoint.c_str()); + return nullptr; + } + } + ranks.push_back({rpc_ctx->endpoint, rpc_ctx->device, rpc_ctx->dispatcher}); + } + + // rank 1 connects to rank 0 on its serving host; endpoints must be mutually reachable + // (e.g. do not bind the servers to 127.0.0.1 when they run on different machines) + std::string host0; + int port0; + if (!parse_endpoint(ranks[0].endpoint, host0, port0)) { + return nullptr; + } + const uint32_t comm_port = (uint32_t) port0 + 1000; + if (host0.size() >= 64) { + return nullptr; + } + + for (const auto & rank : ranks) { + rank.dispatcher->busy_spin_acquire(); + } + + // Send all init requests before reading any response: rank 0 blocks in accept + // until rank 1 has connected. + std::vector responses(n_backends); + for (size_t i = 0; i < n_backends; i++) { + auto request = std::make_shared(); + request->device = ranks[i].device; + request->rank = (uint32_t) i; + request->world = (uint32_t) n_backends; + request->port = comm_port; + if (i > 0) { + memcpy(request->host, host0.c_str(), host0.size()); + } + ranks[i].dispatcher->send_async(RPC_CMD_COMM_INIT, request, sizeof(*request), &responses[i], sizeof(responses[i])); + } + for (size_t i = 0; i < n_backends; i++) { + ranks[i].dispatcher->synchronize(); + } + bool ok = true; + for (size_t i = 0; i < n_backends; i++) { + if (!responses[i].ok) { + GGML_LOG_WARN("%s: rank %zu (%s) failed to initialize\n", __func__, i, ranks[i].endpoint.c_str()); + ok = false; + } + } + if (!ok) { + for (const auto & rank : ranks) { + rank.dispatcher->busy_spin_release(); + } + return nullptr; + } + GGML_LOG_INFO("%s: pairwise communicator initialized (%s <-> %s)\n", __func__, + ranks[0].endpoint.c_str(), ranks[1].endpoint.c_str()); + return new ggml_backend_rpc_comm_context{std::move(ranks)}; +} + +static bool ggml_backend_rpc_comm_allreduce_tensor(void * comm_ctx_v, ggml_tensor ** tensors) { + ggml_backend_rpc_comm_context * comm_ctx = (ggml_backend_rpc_comm_context *) comm_ctx_v; + if (comm_ctx == nullptr) { + return false; + } + const size_t n_ranks = comm_ctx->ranks.size(); + const int64_t ne = ggml_nelements(tensors[0]); + if (ne == 0) { + return true; + } + for (size_t i = 0; i < n_ranks; i++) { + if (tensors[i] == nullptr || tensors[i]->type != GGML_TYPE_F32 || ggml_nelements(tensors[i]) != ne || + !ggml_is_contiguously_allocated(tensors[i]) || + tensors[i]->buffer == nullptr || !ggml_backend_buffer_is_rpc(tensors[i]->buffer)) { + return false; + } + // a rank with a disabled node has garbage in its partial and must contribute zeros, + // which only the fallback path handles + if ((tensors[i]->flags & GGML_TENSOR_FLAG_COMPUTE) == 0) { + return false; + } + } + for (size_t i = 0; i < n_ranks; i++) { + auto request = std::make_shared(); + request->device = comm_ctx->ranks[i].device; + request->tensor = serialize_tensor(tensors[i]); + comm_ctx->ranks[i].dispatcher->send_async(RPC_CMD_COMM_ALLREDUCE, request, sizeof(*request)); + } + return true; +} + static void * ggml_backend_rpc_get_proc_address(ggml_backend_reg_t reg, const char * name) { if (std::strcmp(name, "ggml_backend_rpc_add_server") == 0) { return (void *)ggml_backend_rpc_add_server; @@ -2263,6 +2943,15 @@ static void * ggml_backend_rpc_get_proc_address(ggml_backend_reg_t reg, const ch if (std::strcmp(name, "ggml_backend_rpc_start_server") == 0) { return (void *)ggml_backend_rpc_start_server; } + if (std::strcmp(name, "ggml_backend_comm_init") == 0) { + return (void *)ggml_backend_rpc_comm_init; + } + if (std::strcmp(name, "ggml_backend_comm_free") == 0) { + return (void *)ggml_backend_rpc_comm_free; + } + if (std::strcmp(name, "ggml_backend_comm_allreduce_tensor") == 0) { + return (void *)ggml_backend_rpc_comm_allreduce_tensor; + } return NULL; GGML_UNUSED(reg); @@ -2321,7 +3010,7 @@ ggml_backend_reg_t ggml_backend_rpc_add_server(const char * endpoint) { /* .device = */ ind, /* .name = */ dev_name, /* .description = */ dev_desc, - /* .last_graph_uid = */ 0, + /* .graph_uids = */ {}, }; ggml_backend_dev_t dev = new ggml_backend_device {