RPC: improve loading

This commit is contained in:
Aman Gupta
2026-09-25 08:16:34 +03:00
committed by Georgi Gerganov
parent ee2ddd0938
commit 7d8a964eb7
3 changed files with 136 additions and 42 deletions
+94 -30
View File
@@ -100,11 +100,16 @@ struct rpc_msg_hello_req {
uint8_t conn_caps[RPC_CONN_CAPS_SIZE];
};
// server feature flags in rpc_msg_hello_rsp::flags; old servers send 0
enum rpc_server_flags : uint8_t {
RPC_SERVER_FLAG_CACHE = 1 << 0, // server has a local tensor cache (rpc-server -c)
};
struct rpc_msg_hello_rsp {
uint8_t major;
uint8_t minor;
uint8_t patch;
uint8_t padding;
uint8_t flags;
uint8_t conn_caps[RPC_CONN_CAPS_SIZE];
};
@@ -345,21 +350,30 @@ static bool parse_endpoint(const std::string & endpoint, std::string & host, int
}
// RPC request : | rpc_cmd (1 byte) | request_size (8 bytes) | request_data (request_size bytes) |
// request_data is input followed by an optional payload, so large data does not need to be copied into input
// No response
static bool send_rpc_cmd(socket_ptr sock, enum rpc_cmd cmd, const void * input, size_t input_size) {
static bool send_rpc_cmd_payload(socket_ptr sock, enum rpc_cmd cmd, const void * input, size_t input_size, const void * payload, size_t payload_size) {
uint8_t cmd_byte = cmd;
if (!sock->send_data(&cmd_byte, sizeof(cmd_byte))) {
return false;
}
if (!sock->send_data(&input_size, sizeof(input_size))) {
const size_t request_size = input_size + payload_size;
if (!sock->send_data(&request_size, sizeof(request_size))) {
return false;
}
if (!sock->send_data(input, input_size)) {
return false;
}
if (payload_size > 0 && !sock->send_data(payload, payload_size)) {
return false;
}
return sock->flush();
}
static bool send_rpc_cmd(socket_ptr sock, enum rpc_cmd cmd, const void * input, size_t input_size) {
return send_rpc_cmd_payload(sock, cmd, input, input_size, nullptr, 0);
}
// RPC request : | rpc_cmd (1 byte) | request_size (8 bytes) | request_data (request_size bytes) |
// RPC response: | response_size (8 bytes) | response_data (response_size bytes) |
static bool send_rpc_cmd(socket_ptr sock, enum rpc_cmd cmd, const void * input, size_t input_size, void * output, size_t output_size) {
@@ -392,7 +406,7 @@ static inline void rpc_cpu_relax() {
// 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.
static bool negotiate_hello(const std::shared_ptr<socket_t> & sock) {
static bool negotiate_hello(const std::shared_ptr<socket_t> & sock, uint8_t & server_flags) {
rpc_msg_hello_req request = {};
rpc_msg_hello_rsp response = {};
@@ -408,6 +422,7 @@ static bool negotiate_hello(const std::shared_ptr<socket_t> & sock) {
}
sock->update_caps(response.conn_caps);
server_flags = response.flags;
return true;
}
@@ -470,6 +485,8 @@ public:
void send(enum rpc_cmd cmd, std::shared_ptr<const void> input, size_t input_size, void * output, size_t output_size);
void send_async(enum rpc_cmd cmd, std::shared_ptr<const void> input, size_t input_size);
void send_async(enum rpc_cmd cmd, std::shared_ptr<const void> input, size_t input_size, void * output, size_t output_size);
// payload is sent after input without a copy, the call returns once it is sent
void send_payload(enum rpc_cmd cmd, std::shared_ptr<const void> input, size_t input_size, const void * payload, size_t payload_size);
ggml_backend_event_t event_new(ggml_backend_dev_t dev);
void event_free(ggml_backend_event_t event);
@@ -483,6 +500,8 @@ public:
void start(const std::string & endpoint);
void work();
bool server_has_cache() const { return server_flags & RPC_SERVER_FLAG_CACHE; }
~rpc_dispatcher();
private:
@@ -492,6 +511,8 @@ private:
size_t input_size;
void * output;
size_t output_size;
const void * payload = nullptr;
size_t payload_size = 0;
std::promise<void> completion;
};
using rpc_msg_ptr = std::shared_ptr<rpc_msg>;
@@ -504,6 +525,7 @@ private:
std::unordered_map<uint32_t, std::unordered_set<uint64_t>> graph_uids;
rpc_msg_queue queue;
socket_ptr sock;
uint8_t server_flags = 0;
std::atomic_uint busy_spin_users = 0;
std::atomic_bool running;
std::thread thread;
@@ -536,6 +558,20 @@ void rpc_dispatcher::send_async(enum rpc_cmd cmd, std::shared_ptr<const void> in
GGML_ASSERT(queue.push(msg));
}
void rpc_dispatcher::send_payload(enum rpc_cmd cmd, std::shared_ptr<const void> input, size_t input_size, const void * payload, size_t payload_size) {
auto msg = std::make_shared<rpc_msg>();
msg->cmd = cmd;
msg->input = input;
msg->input_size = input_size;
msg->output = nullptr;
msg->output_size = 0;
msg->payload = payload;
msg->payload_size = payload_size;
GGML_ASSERT(queue.push(msg));
auto future = msg->completion.get_future();
future.wait();
}
void rpc_dispatcher::send(enum rpc_cmd cmd, std::shared_ptr<const void> input, size_t input_size, void * output, size_t output_size) {
auto msg = std::make_shared<rpc_msg>();
msg->cmd = cmd;
@@ -619,7 +655,7 @@ void rpc_dispatcher::start(const std::string & endpoint) {
if (sock == nullptr) {
GGML_ABORT("Failed to connect to %s\n", endpoint.c_str());
}
if (!negotiate_hello(sock)) {
if (!negotiate_hello(sock, server_flags)) {
GGML_ABORT("RPC handshake failed for %s\n", endpoint.c_str());
}
LOG_DBG("[%s] connected to %s\n", __func__, endpoint.c_str());
@@ -643,7 +679,7 @@ void rpc_dispatcher::work() {
bool status = send_rpc_cmd(sock, msg_ptr->cmd, msg_ptr->input.get(), msg_ptr->input_size, msg_ptr->output, msg_ptr->output_size);
RPC_STATUS_ASSERT(status);
} else {
bool status = send_rpc_cmd(sock, msg_ptr->cmd, msg_ptr->input.get(), msg_ptr->input_size);
bool status = send_rpc_cmd_payload(sock, msg_ptr->cmd, msg_ptr->input.get(), msg_ptr->input_size, msg_ptr->payload, msg_ptr->payload_size);
RPC_STATUS_ASSERT(status);
}
}
@@ -776,14 +812,17 @@ static void ggml_backend_rpc_buffer_memset_tensor(
}
// input serialization format: | rpc_tensor | cache_flag (1 byte) | offset (8 bytes) | data (size bytes)
// with size == 0 only the header is serialized and data is sent separately as payload
static std::shared_ptr<uint8_t> serialize_set_tensor(const rpc_tensor & rpc_tensor, uint8_t cache_flag, uint64_t offset, const void * data, size_t size, size_t & input_size) {
input_size = sizeof(rpc_tensor) + sizeof(cache_flag) + sizeof(offset) + size;
uint8_t * input = new uint8_t[input_size]();
uint8_t * input = new uint8_t[input_size];
uint8_t * p = input;
memcpy(p, &rpc_tensor, sizeof(rpc_tensor)); p += sizeof(rpc_tensor);
memcpy(p, &cache_flag, sizeof(cache_flag)); p += sizeof(cache_flag);
memcpy(p, &offset, sizeof(offset)); p += sizeof(offset);
memcpy(p, data, size);
if (size > 0) {
memcpy(p, data, size);
}
return std::shared_ptr<uint8_t>(input, std::default_delete<uint8_t[]>());
}
@@ -791,15 +830,16 @@ static std::shared_ptr<uint8_t> serialize_set_tensor(const rpc_tensor & rpc_tens
// compute-buffer inputs (the activations ggml_backend_sched copies between backends) must not
// take this path, otherwise with `rpc-server -c` every ubatch above the threshold is written
// to the cache directory and later served from there.
static bool rpc_use_hash_cache(const ggml_tensor * tensor, size_t size) {
return size > HASH_THRESHOLD && tensor->buffer->usage == GGML_BACKEND_BUFFER_USAGE_WEIGHTS;
// hashing is slow (~1 GB/s), so skip it when the server has no cache to look up.
static bool rpc_use_hash_cache(const rpc_dispatcher & dispatcher, const ggml_tensor * tensor, size_t size) {
return dispatcher.server_has_cache() && size > HASH_THRESHOLD && tensor->buffer->usage == GGML_BACKEND_BUFFER_USAGE_WEIGHTS;
}
static void ggml_backend_rpc_buffer_set_tensor(ggml_backend_buffer_t buffer, ggml_tensor * tensor, const void * data, size_t offset, size_t size) {
ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context;
rpc_tensor rpc_tensor = serialize_tensor(tensor);
uint8_t cache_flag = 0;
if (rpc_use_hash_cache(tensor, size)) {
if (rpc_use_hash_cache(*ctx->dispatcher, tensor, size)) {
auto request = std::make_shared<rpc_msg_set_tensor_hash_req>();
request->tensor = rpc_tensor;
request->offset = offset;
@@ -813,9 +853,9 @@ static void ggml_backend_rpc_buffer_set_tensor(ggml_backend_buffer_t buffer, ggm
// the server has no cache entry for this tensor - ask it to save one
cache_flag = 1;
}
size_t input_size;
auto input = serialize_set_tensor(rpc_tensor, cache_flag, offset, data, size, input_size);
ctx->dispatcher->send(RPC_CMD_SET_TENSOR, input, input_size);
size_t header_size;
auto header = serialize_set_tensor(rpc_tensor, cache_flag, offset, nullptr, 0, header_size);
ctx->dispatcher->send_payload(RPC_CMD_SET_TENSOR, header, header_size, data, size);
}
static void ggml_backend_rpc_buffer_set_tensor_2d(ggml_backend_buffer_t buffer, ggml_tensor * tensor, const void * data,
@@ -1064,7 +1104,7 @@ static void ggml_backend_rpc_set_tensor_async(ggml_backend_t backend, ggml_tenso
ggml_backend_rpc_context * ctx = (ggml_backend_rpc_context *)backend->context;
rpc_tensor rpc_tensor = serialize_tensor(tensor);
uint8_t cache_flag = 0;
if (rpc_use_hash_cache(tensor, size)) {
if (rpc_use_hash_cache(*ctx->dispatcher, tensor, size)) {
auto request = std::make_shared<rpc_msg_set_tensor_hash_req>();
request->tensor = rpc_tensor;
request->offset = offset;
@@ -1302,7 +1342,7 @@ public:
bool free_buffer(const rpc_msg_free_buffer_req & request);
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(socket_ptr sock, uint64_t input_size);
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);
@@ -1349,12 +1389,15 @@ private:
// 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;
// SET_TENSOR data is received in chunks through this buffer
std::vector<uint8_t> set_tensor_buf;
};
void rpc_server::hello(rpc_msg_hello_rsp & response) {
response.major = RPC_PROTO_MAJOR_VERSION;
response.minor = RPC_PROTO_MINOR_VERSION;
response.patch = RPC_PROTO_PATCH_VERSION;
response.flags = cache_dir ? RPC_SERVER_FLAG_CACHE : 0;
LOG_DBG("[%s] version: %d.%d.%d\n", __func__, response.major, response.minor, response.patch);
}
@@ -1591,19 +1634,21 @@ ggml_tensor * rpc_server::deserialize_tensor(struct ggml_context * ctx, const rp
}
bool rpc_server::set_tensor(const std::vector<uint8_t> & input) {
bool rpc_server::set_tensor(socket_ptr sock, uint64_t input_size) {
sync_all_backends();
// serialization format: | rpc_tensor | cache_flag (1 byte) | offset (8 bytes) | data (size bytes) |
// only the header is read here, the data is streamed from the socket below
uint8_t cache_flag;
uint64_t offset;
const size_t header_size = sizeof(rpc_tensor) + sizeof(cache_flag) + sizeof(offset);
if (input.size() < header_size) {
uint8_t header[sizeof(rpc_tensor) + sizeof(cache_flag) + sizeof(offset)];
const size_t header_size = sizeof(header);
if (input_size < header_size || !sock->recv_data(header, header_size)) {
return false;
}
const rpc_tensor * in_tensor = (const rpc_tensor *)input.data();
memcpy(&cache_flag, input.data() + sizeof(rpc_tensor), sizeof(cache_flag));
memcpy(&offset, input.data() + sizeof(rpc_tensor) + sizeof(cache_flag), sizeof(offset));
const size_t size = input.size() - header_size;
const rpc_tensor * in_tensor = (const rpc_tensor *)header;
memcpy(&cache_flag, header + sizeof(rpc_tensor), sizeof(cache_flag));
memcpy(&offset, header + sizeof(rpc_tensor) + sizeof(cache_flag), sizeof(offset));
const size_t size = input_size - header_size;
struct ggml_init_params params {
/*.mem_size =*/ ggml_tensor_overhead(),
@@ -1632,18 +1677,37 @@ bool rpc_server::set_tensor(const std::vector<uint8_t> & input) {
}
}
const void * data = input.data() + header_size;
if (cache_dir && cache_flag) {
uint64_t hash = fnv_hash((const uint8_t*)data, size);
// the whole data is needed to compute the hash
std::vector<uint8_t> data(size);
if (!sock->recv_data(data.data(), size)) {
return false;
}
uint64_t hash = fnv_hash(data.data(), size);
char hash_str[17];
snprintf(hash_str, sizeof(hash_str), "%016" PRIx64, hash);
// save to cache_dir/hash_str
fs::path cache_file = fs::path(cache_dir) / hash_str;
std::ofstream ofs(cache_file, std::ios::binary);
ofs.write((const char *)data, size);
ofs.write((const char *)data.data(), size);
GGML_LOG_INFO("[%s] saved to '%s'\n", __func__, cache_file.string().c_str());
ggml_backend_tensor_set(tensor, data.data(), offset, size);
return true;
}
// a buffer of the tensor size would be allocated and page faulted for every tensor, which is slow
const size_t chunk_size = 16*1024*1024;
if (set_tensor_buf.size() < std::min(size, chunk_size)) {
set_tensor_buf.resize(std::min(size, chunk_size));
}
for (size_t done = 0; done < size; ) {
const size_t n = std::min(size - done, chunk_size);
if (!sock->recv_data(set_tensor_buf.data(), n)) {
return false;
}
ggml_backend_tensor_set(tensor, set_tensor_buf.data(), offset + done, n);
done += n;
}
ggml_backend_tensor_set(tensor, data, offset, size);
return true;
}
@@ -2472,11 +2536,11 @@ static void rpc_serve_client(const std::vector<ggml_backend_t> & backends, const
break;
}
case RPC_CMD_SET_TENSOR: {
std::vector<uint8_t> input;
if (!recv_msg(sock, input)) {
uint64_t input_size;
if (!sock->recv_data(&input_size, sizeof(input_size))) {
return;
}
if (!server.set_tensor(input)) {
if (!server.set_tensor(sock, input_size)) {
return;
}
break;
+21 -6
View File
@@ -69,6 +69,11 @@ struct rdma_conn {
struct ibv_mr * rx_mr = nullptr;
int rx_head = 0;
// received slot that is only partly consumed, so recv sizes need not match send sizes
int rx_cur = -1;
size_t rx_off = 0;
size_t rx_len = 0;
uint32_t max_inline = 0;
uint8_t * rx_slot(int i) const {
@@ -458,14 +463,24 @@ bool socket_t::impl::rdma_recv(void * data, size_t size) {
uint8_t * dst = (uint8_t *)data;
size_t rem = size;
while (rem > 0) {
struct ibv_wc wc;
if (!rdma_poll(c->rcq, &wc)) return false;
if (c->rx_cur < 0) {
struct ibv_wc wc;
if (!rdma_poll(c->rcq, &wc)) return false;
int slot = (int)wc.wr_id;
size_t got = wc.byte_len;
memcpy(dst, c->rx_slot(slot), got);
c->rx_cur = (int)wc.wr_id;
c->rx_off = 0;
c->rx_len = wc.byte_len;
}
if (!c->post_rx(slot)) return false;
size_t got = std::min(rem, c->rx_len - c->rx_off);
memcpy(dst, c->rx_slot(c->rx_cur) + c->rx_off, got);
c->rx_off += got;
// repost the slot only once it is fully consumed
if (c->rx_off == c->rx_len) {
if (!c->post_rx(c->rx_cur)) return false;
c->rx_cur = -1;
}
dst += got;
rem -= got;
+21 -6
View File
@@ -1622,6 +1622,8 @@ bool llama_model_loader::load_all_data(
});
}
// staging buffer for non-host buffers without async uploads, reused because a fresh allocation per tensor page faults on every load
std::vector<no_init<uint8_t>> read_buf;
for (struct ggml_tensor * cur : tensors) {
const auto * weight = get_weight(ggml_get_name(cur));
if (weight == nullptr) {
@@ -1735,12 +1737,25 @@ bool llama_model_loader::load_all_data(
buffer_idx %= n_buffers;
}
} else {
// scoped to one tensor so only one staging buffer is alive at a time
std::vector<no_init<uint8_t>> read_buf(n_size);
file->seek(weight->offs, SEEK_SET);
file->read_raw(read_buf.data(), n_size);
ggml_backend_tensor_set(cur, read_buf.data(), 0, n_size);
if (check_tensors && !ggml_validate_row_data(cur->type, read_buf.data(), n_size)) {
// with direct IO read the aligned range around the tensor and use the data inside it, so no bounce buffer is needed
const size_t align = file->read_alignment();
const size_t aligned_offset = weight->offs & ~(align - 1);
const size_t pad = weight->offs - aligned_offset;
const size_t read_size = (pad + n_size + align - 1) & ~(align - 1);
const size_t buf_size = read_size + align - 1;
// tensors come biggest-first: shrink the buffer when it gets too large, so peak memory stays close to one tensor
if (read_buf.size() < buf_size || read_buf.size() > 2*buf_size) {
read_buf = {};
read_buf.resize(buf_size);
}
uint8_t * read_dst = (uint8_t *) GGML_PAD((uintptr_t) read_buf.data(), align);
file->seek(aligned_offset, SEEK_SET);
file->read_raw_unsafe(read_dst, read_size);
const uint8_t * data = read_dst + pad;
ggml_backend_tensor_set(cur, data, 0, n_size);
if (check_tensors && !ggml_validate_row_data(cur->type, data, n_size)) {
throw std::runtime_error(format("tensor '%s' has invalid data", ggml_get_name(cur)));
}
}