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 {