Commit a46709b68 for llama.cpp
commit a46709b683aba9274d8ab29f5b42f7e551d1dff6
Author: Aman Gupta <amangupta052@gmail.com>
Date: Tue Oct 6 20:55:39 2026 +0530
RPC: add `-sm tensor` (#26610)
* rpc: allow -sm tensor
* fix flush for apple rdma
* move graph_uids to rpc_dispatcher
* cont : fix conflict
* cont: stop spinning dispatcher thread
* remove meta backend change
* add TODO to simplify logic
* rpc: bump major version
---------
Co-authored-by: Georgi Gerganov <ggerganov@gmail.com>
diff --git a/ggml/include/ggml-rpc.h b/ggml/include/ggml-rpc.h
index 1f8cb7906..482bd3666 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 7
+#define RPC_PROTO_MAJOR_VERSION 8
#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 0394433c0..8ed5f4ebb 100644
--- a/ggml/src/ggml-backend-meta.cpp
+++ b/ggml/src/ggml-backend-meta.cpp
@@ -869,7 +869,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);
diff --git a/ggml/src/ggml-rpc/ggml-rpc.cpp b/ggml/src/ggml-rpc/ggml-rpc.cpp
index 158a15bb8..e9af71228 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 <array>
+#include <chrono>
#include <cinttypes>
#include <optional>
#include <string>
+#include <thread>
#include <vector>
#include <queue>
#include <condition_variable>
@@ -21,7 +23,6 @@
#include <filesystem>
#include <algorithm>
#include <atomic>
-#include <thread>
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,6 @@ struct ggml_backend_rpc_device_context {
uint32_t device;
std::string name;
std::string description;
- uint64_t last_graph_uid;
};
struct ggml_backend_rpc_buffer_type_context {
@@ -232,6 +271,7 @@ struct ggml_backend_rpc_buffer_type_context {
class rpc_dispatcher;
struct ggml_backend_rpc_context {
std::shared_ptr<rpc_dispatcher> dispatcher;
+ std::string endpoint;
uint32_t device;
std::string name;
};
@@ -341,6 +381,17 @@ static bool send_rpc_cmd(socket_ptr sock, enum rpc_cmd cmd, const void * input,
// RPC client-side implementation
+// with busy spinning on, the dispatcher still blocks on its queue after this long without commands
+static constexpr auto RPC_BUSY_SPIN_IDLE_TIME = std::chrono::milliseconds(100);
+
+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 +440,16 @@ public:
return true;
}
+ bool try_pop(T* out) {
+ std::unique_lock<std::mutex> lock(mutex);
+ if (interrupted || queue.empty()) {
+ return false;
+ }
+ *out = queue.front();
+ queue.pop();
+ return true;
+ }
+
void interrupt() {
std::unique_lock<std::mutex> lock(mutex);
interrupted = true;
@@ -418,6 +479,9 @@ public:
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 graph_compute(uint32_t device, const ggml_cgraph * cgraph);
void start(const std::string & endpoint);
void work();
@@ -439,8 +503,11 @@ private:
rpc_msg_ptr msg;
std::shared_future<void> sf;
};
+ std::mutex graph_mutex;
+ std::unordered_map<uint32_t, std::unordered_set<uint64_t>> graph_uids;
rpc_msg_queue queue;
socket_ptr sock;
+ std::atomic_uint busy_spin_users = 0;
std::atomic_bool running;
std::thread thread;
};
@@ -532,6 +599,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;
@@ -555,9 +631,18 @@ void rpc_dispatcher::start(const std::string & endpoint) {
}
void rpc_dispatcher::work() {
+ auto last_cmd = std::chrono::steady_clock::now();
while (running) {
rpc_msg_ptr msg_ptr;
- if (!queue.pop(&msg_ptr)) {
+ // spin only while commands keep coming, so an idle dispatcher does not keep a core busy
+ const bool spin = busy_spin_users.load(std::memory_order_relaxed) != 0 &&
+ std::chrono::steady_clock::now() - last_cmd < RPC_BUSY_SPIN_IDLE_TIME;
+ if (spin) {
+ if (!queue.try_pop(&msg_ptr)) {
+ rpc_cpu_relax();
+ continue;
+ }
+ } else if (!queue.pop(&msg_ptr)) {
break;
}
if (msg_ptr->cmd != RPC_CMD_NONE) {
@@ -570,6 +655,7 @@ void rpc_dispatcher::work() {
}
}
msg_ptr->completion.set_value();
+ last_cmd = std::chrono::steady_clock::now();
}
}
@@ -625,7 +711,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<rpc_dispatcher> & 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));
@@ -638,7 +724,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<uint64_t>(tensor->data);
} else {
@@ -740,6 +826,46 @@ static void ggml_backend_rpc_buffer_set_tensor(ggml_backend_buffer_t buffer, ggm
ctx->dispatcher->send(RPC_CMD_SET_TENSOR, input, 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<uint8_t> input_ptr(input, std::default_delete<uint8_t[]>());
+ 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<rpc_msg_get_tensor_2d_req>();
+ 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<uint8_t> 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<rpc_msg_get_tensor_req>();
@@ -785,8 +911,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,
@@ -989,7 +1115,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<rpc_dispatcher> & dispatcher, std::vector<rpc_tensor> & tensors, std::unordered_set<ggml_tensor*> & visited) {
+static void add_tensor(ggml_tensor * tensor, const ggml_cgraph * cgraph, const rpc_dispatcher * dispatcher, std::vector<rpc_tensor> & tensors, std::unordered_set<ggml_tensor*> & visited) {
if (tensor == nullptr) {
return;
}
@@ -1009,7 +1135,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<rpc_dispatcher> & 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<rpc_tensor> tensors;
std::unordered_set<ggml_tensor*> visited;
@@ -1017,13 +1143,15 @@ static uint8_t * serialize_graph(uint32_t device, const ggml_cgraph * cgraph, co
add_tensor(cgraph->nodes[i], cgraph, dispatcher, 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++) {
@@ -1037,24 +1165,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<std::mutex> lock(graph_mutex);
GGML_ASSERT(cgraph->n_nodes > 0);
- bool reuse = cgraph->uid != 0 && rpc_dev_ctx->last_graph_uid == cgraph->uid;
+ 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<rpc_msg_graph_recompute_req>();
- request->device = rpc_ctx->device;
- rpc_ctx->dispatcher->send_async(RPC_CMD_GRAPH_RECOMPUTE, request, sizeof(*request));
+ request->device = device;
+ request->uid = cgraph->uid;
+ send_async(RPC_CMD_GRAPH_RECOMPUTE, request, sizeof(*request));
} else {
- rpc_dev_ctx->last_graph_uid = cgraph->uid;
+ if (cgraph->uid != 0) {
+ if (device_graph_uids.size() >= GRAPH_CACHE_MAX) {
+ device_graph_uids.clear();
+ }
+ 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<uint8_t> input_ptr(input, std::default_delete<uint8_t[]>());
- 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;
}
@@ -1069,13 +1206,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,
@@ -1123,6 +1272,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,
};
@@ -1157,6 +1307,7 @@ public:
rpc_server(std::vector<ggml_backend_t> 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();
@@ -1169,11 +1320,16 @@ public:
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<uint8_t> & input);
+ bool set_tensor_2d(const std::vector<uint8_t> & 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<uint8_t> & response);
+ bool get_tensor_2d(const rpc_msg_get_tensor_2d_req & request, std::vector<uint8_t> & response);
bool copy_tensor(const rpc_msg_copy_tensor_req & request, rpc_msg_copy_tensor_rsp & response);
bool graph_compute(const std::vector<uint8_t> & 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);
@@ -1184,6 +1340,7 @@ public:
};
private:
+ void sync_all_backends();
bool get_cached_file(uint64_t hash, std::vector<uint8_t> & data);
ggml_tensor * deserialize_tensor(struct ggml_context * ctx, const rpc_tensor * tensor);
ggml_tensor * create_node(uint64_t id,
@@ -1192,11 +1349,23 @@ private:
std::unordered_map<uint64_t, struct ggml_tensor*> & 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<uint8_t> send_buf;
+ std::vector<uint8_t> recv_buf;
+ };
+
std::vector<ggml_backend_t> backends;
const char * cache_dir;
std::unordered_set<ggml_backend_buffer_t> buffers;
- // store the last computed graph for each backend
- std::vector<stored_graph> stored_graphs;
+ // computed graphs cached per backend, keyed by uid
+ std::vector<std::unordered_map<uint64_t, stored_graph>> stored_graphs;
+ std::vector<comm_state> comm_states;
};
void rpc_server::hello(rpc_msg_hello_rsp & response) {
@@ -1304,6 +1473,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<ggml_backend_buffer_t>(request.remote_ptr);
if (buffers.find(buffer) == buffers.end()) {
@@ -1312,8 +1482,10 @@ bool rpc_server::free_buffer(const rpc_msg_free_buffer_req & request) {
}
// Discard all cached graphs to avoid use-after-free in graph_recompute,
// since their nodes may hold pointers to the buffer being freed.
- for (auto & sg : stored_graphs) {
- sg.graph = nullptr;
+ for (auto & sgs : stored_graphs) {
+ for (auto & sg : sgs) {
+ sg.second.graph = nullptr;
+ }
}
ggml_backend_buffer_free(buffer);
buffers.erase(buffer);
@@ -1321,6 +1493,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<ggml_backend_buffer_t>(request.remote_ptr);
if (buffers.find(buffer) == buffers.end()) {
@@ -1332,6 +1505,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,
@@ -1407,13 +1581,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;
@@ -1428,6 +1609,7 @@ ggml_tensor * rpc_server::deserialize_tensor(struct ggml_context * ctx, const rp
bool rpc_server::set_tensor(const std::vector<uint8_t> & input) {
+ sync_all_backends();
// serialization format: | rpc_tensor | cache_flag (1 byte) | offset (8 bytes) | data (size bytes) |
uint8_t cache_flag;
uint64_t offset;
@@ -1482,6 +1664,67 @@ bool rpc_server::set_tensor(const std::vector<uint8_t> & input) {
return true;
}
+bool rpc_server::set_tensor_2d(const std::vector<uint8_t> & 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<uint8_t> & data) {
if (!cache_dir) {
return false;
@@ -1504,6 +1747,7 @@ bool rpc_server::get_cached_file(uint64_t hash, std::vector<uint8_t> & 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<uint8_t> cached_file;
if (!get_cached_file(request.hash, cached_file)) {
response.result = 0;
@@ -1545,6 +1789,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,
@@ -1580,6 +1825,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<uint8_t> & response) {
+ sync_all_backends();
struct ggml_init_params params {
/*.mem_size =*/ ggml_tensor_overhead(),
/*.mem_buffer =*/ NULL,
@@ -1614,7 +1860,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<uint8_t> & 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,
@@ -1713,8 +2008,8 @@ ggml_tensor * rpc_server::create_node(uint64_t id,
bool rpc_server::graph_compute(const std::vector<uint8_t> & 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();
@@ -1724,10 +2019,13 @@ bool rpc_server::graph_compute(const std::vector<uint8_t> & 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;
@@ -1735,19 +2033,26 @@ bool rpc_server::graph_compute(const std::vector<uint8_t> & 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) };
@@ -1778,9 +2083,9 @@ bool rpc_server::graph_compute(const std::vector<uint8_t> & 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;
}
@@ -1789,16 +2094,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()) {
@@ -1993,6 +2498,30 @@ static void rpc_serve_client(const std::vector<ggml_backend_t> & backends, const
}
break;
}
+ case RPC_CMD_SET_TENSOR_2D: {
+ std::vector<uint8_t> 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<uint8_t> 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))) {
@@ -2065,6 +2594,40 @@ static void rpc_serve_client(const std::vector<ggml_backend_t> & 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))) {
@@ -2294,6 +2857,137 @@ 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<rpc_dispatcher> dispatcher;
+ };
+ std::vector<rank_info> 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<rpc_msg_comm_free_req>();
+ 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<ggml_backend_rpc_comm_context::rank_info> 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});
+ }
+
+ // TODO: simplify this logic
+ // 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<rpc_msg_comm_init_rsp> responses(n_backends);
+ for (size_t i = 0; i < n_backends; i++) {
+ auto request = std::make_shared<rpc_msg_comm_init_req>();
+ 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<rpc_msg_comm_allreduce_req>();
+ 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;
@@ -2301,6 +2995,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);
@@ -2359,7 +3062,6 @@ ggml_backend_reg_t ggml_backend_rpc_add_server(const char * endpoint) {
/* .device = */ ind,
/* .name = */ dev_name,
/* .description = */ dev_desc,
- /* .last_graph_uid = */ 0,
};
ggml_backend_dev_t dev = new ggml_backend_device {