From f9b423302d246e529407cc075cd869a70bf1aa2b Mon Sep 17 00:00:00 2001 From: Aman Gupta Date: Sun, 13 Sep 2026 23:40:20 +0800 Subject: [PATCH] move graph_uids to rpc_dispatcher --- ggml/src/ggml-rpc/ggml-rpc.cpp | 44 ++++++++++++++++++---------------- 1 file changed, 23 insertions(+), 21 deletions(-) diff --git a/ggml/src/ggml-rpc/ggml-rpc.cpp b/ggml/src/ggml-rpc/ggml-rpc.cpp index f172903760..1093ba02be 100644 --- a/ggml/src/ggml-rpc/ggml-rpc.cpp +++ b/ggml/src/ggml-rpc/ggml-rpc.cpp @@ -258,8 +258,6 @@ struct ggml_backend_rpc_device_context { uint32_t device; std::string name; std::string description; - // uids of graphs cached by the server for this device - std::unordered_set graph_uids; }; struct ggml_backend_rpc_buffer_type_context { @@ -480,6 +478,7 @@ public: void synchronize(); void busy_spin_acquire(); void busy_spin_release(); + void graph_compute(uint32_t device, const ggml_cgraph * cgraph); void start(const std::string & endpoint); void work(); @@ -501,6 +500,8 @@ private: rpc_msg_ptr msg; std::shared_future sf; }; + std::mutex graph_mutex; + std::unordered_map> graph_uids; rpc_msg_queue queue; socket_ptr sock; std::atomic_uint busy_spin_users = 0; @@ -702,7 +703,7 @@ static bool ggml_backend_buffer_is_rpc(ggml_backend_buffer_t buffer) { return buffer->iface.free_buffer == ggml_backend_rpc_buffer_free_buffer; } -static rpc_tensor serialize_tensor(const ggml_tensor * tensor, const std::shared_ptr & dispatcher = nullptr) { +static rpc_tensor serialize_tensor(const ggml_tensor * tensor, const rpc_dispatcher * dispatcher = nullptr) { rpc_tensor result; if (!tensor) { memset(&result, 0, sizeof(result)); @@ -715,7 +716,7 @@ static rpc_tensor serialize_tensor(const ggml_tensor * tensor, const std::shared ggml_backend_buffer_t buffer = tensor->buffer; ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context; // ref: https://github.com/ggml-org/llama.cpp/pull/26500 - if (ctx != nullptr && (dispatcher == nullptr || ctx->dispatcher == dispatcher)) { + if (ctx != nullptr && (dispatcher == nullptr || ctx->dispatcher.get() == dispatcher)) { result.buffer = ctx->remote_ptr; result.data = reinterpret_cast(tensor->data); } else { @@ -1081,7 +1082,7 @@ static void ggml_backend_rpc_synchronize(ggml_backend_t backend) { rpc_ctx->dispatcher->synchronize(); } -static void add_tensor(ggml_tensor * tensor, const ggml_cgraph * cgraph, const std::shared_ptr & dispatcher, std::vector & tensors, std::unordered_set & visited) { +static void add_tensor(ggml_tensor * tensor, const ggml_cgraph * cgraph, const rpc_dispatcher * dispatcher, std::vector & tensors, std::unordered_set & visited) { if (tensor == nullptr) { return; } @@ -1101,7 +1102,7 @@ static void add_tensor(ggml_tensor * tensor, const ggml_cgraph * cgraph, const s tensors.push_back(result); } -static uint8_t * serialize_graph(uint32_t device, const ggml_cgraph * cgraph, const std::shared_ptr & dispatcher, size_t * output_size) { +static uint8_t * serialize_graph(uint32_t device, const ggml_cgraph * cgraph, const rpc_dispatcher * dispatcher, size_t * output_size) { uint32_t n_nodes = cgraph->n_nodes; std::vector tensors; std::unordered_set visited; @@ -1131,31 +1132,33 @@ static uint8_t * serialize_graph(uint32_t device, const ggml_cgraph * cgraph, co return output; } -static enum ggml_status ggml_backend_rpc_graph_compute(ggml_backend_t backend, ggml_cgraph * cgraph) { - ggml_backend_rpc_context * rpc_ctx = (ggml_backend_rpc_context *)backend->context; - ggml_backend_dev_t rpc_dev = ggml_backend_get_device(backend); - ggml_backend_rpc_device_context * rpc_dev_ctx = (ggml_backend_rpc_device_context *)rpc_dev->context; - +void rpc_dispatcher::graph_compute(uint32_t device, const ggml_cgraph * cgraph) { + std::lock_guard lock(graph_mutex); GGML_ASSERT(cgraph->n_nodes > 0); - auto & graph_uids = rpc_dev_ctx->graph_uids; - bool reuse = cgraph->uid != 0 && graph_uids.count(cgraph->uid) > 0; + auto & device_graph_uids = graph_uids[device]; + bool reuse = cgraph->uid != 0 && device_graph_uids.count(cgraph->uid) > 0; if (reuse) { auto request = std::make_shared(); - request->device = rpc_ctx->device; + request->device = device; request->uid = cgraph->uid; - rpc_ctx->dispatcher->send_async(RPC_CMD_GRAPH_RECOMPUTE, request, sizeof(*request)); + send_async(RPC_CMD_GRAPH_RECOMPUTE, request, sizeof(*request)); } else { if (cgraph->uid != 0) { - if (graph_uids.size() >= GRAPH_CACHE_MAX) { - graph_uids.clear(); + if (device_graph_uids.size() >= GRAPH_CACHE_MAX) { + device_graph_uids.clear(); } - graph_uids.insert(cgraph->uid); + device_graph_uids.insert(cgraph->uid); } size_t input_size = 0; - uint8_t * input = serialize_graph(rpc_ctx->device, cgraph, rpc_ctx->dispatcher, &input_size); + uint8_t * input = serialize_graph(device, cgraph, this, &input_size); std::shared_ptr input_ptr(input, std::default_delete()); - rpc_ctx->dispatcher->send_async(RPC_CMD_GRAPH_COMPUTE, input_ptr, input_size); + send_async(RPC_CMD_GRAPH_COMPUTE, input_ptr, input_size); } +} + +static enum ggml_status ggml_backend_rpc_graph_compute(ggml_backend_t backend, ggml_cgraph * cgraph) { + ggml_backend_rpc_context * rpc_ctx = (ggml_backend_rpc_context *)backend->context; + rpc_ctx->dispatcher->graph_compute(rpc_ctx->device, cgraph); return GGML_STATUS_SUCCESS; } @@ -3016,7 +3019,6 @@ ggml_backend_reg_t ggml_backend_rpc_add_server(const char * endpoint) { /* .device = */ ind, /* .name = */ dev_name, /* .description = */ dev_desc, - /* .graph_uids = */ {}, }; ggml_backend_dev_t dev = new ggml_backend_device {