server : avoid checkpoint data host copies (#22558)

* server : avoid checkpoint data host copies

* llama : refactor llama_io_read_i
This commit is contained in:
Georgi Gerganov
2026-05-02 18:03:25 +03:00
committed by GitHub
parent 09294365a9
commit 0754b7b6fe
6 changed files with 132 additions and 72 deletions
+66 -10
View File
@@ -2253,6 +2253,28 @@ public:
llama_io_write_buffer(
uint8_t * p, size_t len) : ptr(p), buf_size(len) {}
~llama_io_write_buffer() {
#if 1
// TODO: add backend support to batch tensor_get? or some other way to speed this up
for (const auto & info : winfos) {
ggml_backend_tensor_get(info.tensor, info.ptr, info.offset, info.size);
}
#else
// flush the writes asynchronously
// this helps on Macs, but on other devices - it does not. just an example
std::vector<std::future<void>> futures;
futures.reserve(winfos.size());
for (const auto & info : winfos) {
futures.push_back(std::async(std::launch::async, [info]() {
ggml_backend_tensor_get(info.tensor, info.ptr, info.offset, info.size);
}));
}
for (auto & f : futures) {
f.wait();
}
#endif
}
void write(const void * src, size_t size) override {
if (size > buf_size) {
throw std::runtime_error("unexpectedly reached end of buffer");
@@ -2267,7 +2289,10 @@ public:
if (size > buf_size) {
throw std::runtime_error("unexpectedly reached end of buffer");
}
ggml_backend_tensor_get(tensor, ptr, offset, size);
// save the write for later during destruction
winfos.push_back({tensor, ptr, size, offset});
ptr += size;
size_written += size;
buf_size -= size;
@@ -2281,25 +2306,48 @@ private:
uint8_t * ptr;
size_t buf_size = 0;
size_t size_written = 0;
struct write_info {
const ggml_tensor * tensor;
uint8_t * ptr;
size_t size;
size_t offset;
};
std::vector<write_info> winfos;
};
class llama_io_read_buffer : public llama_io_read_i {
public:
llama_io_read_buffer(const uint8_t * p, size_t len) : ptr(p), buf_size(len) {}
const uint8_t * read(size_t size) override {
const uint8_t * base_ptr = ptr;
~llama_io_read_buffer() {
// flush the reads
for (const auto & info : rinfos) {
ggml_backend_tensor_set(info.tensor, info.ptr, info.offset, info.size);
}
}
void read(void * dst, size_t size) override {
if (size > buf_size) {
throw std::runtime_error("unexpectedly reached end of buffer");
}
memcpy(dst, ptr, size);
ptr += size;
size_read += size;
buf_size -= size;
return base_ptr;
}
void read_to(void * dst, size_t size) override {
memcpy(dst, read(size), size);
void read_tensor(ggml_tensor * tensor, size_t offset, size_t size) override {
if (size > buf_size) {
throw std::runtime_error("unexpectedly reached end of buffer");
}
// save for later during destruction
rinfos.push_back({tensor, ptr, size, offset});
ptr += size;
size_read += size;
buf_size -= size;
}
size_t n_bytes() override {
@@ -2310,6 +2358,14 @@ private:
const uint8_t * ptr;
size_t buf_size = 0;
size_t size_read = 0;
struct read_info {
ggml_tensor * tensor;
const uint8_t * ptr;
size_t size;
size_t offset;
};
std::vector<read_info> rinfos;
};
class llama_io_write_file : public llama_io_write_i {
@@ -2341,15 +2397,15 @@ class llama_io_read_file : public llama_io_read_i {
public:
llama_io_read_file(llama_file * f) : file(f) {}
void read_to(void * dst, size_t size) override {
void read(void * dst, size_t size) override {
file->read_raw(dst, size);
size_read += size;
}
const uint8_t * read(size_t size) override {
void read_tensor(ggml_tensor * tensor, size_t offset, size_t size) override {
temp_buffer.resize(size);
read_to(temp_buffer.data(), size);
return temp_buffer.data();
read(temp_buffer.data(), size);
ggml_backend_tensor_set(tensor, temp_buffer.data(), offset, size);
}
size_t n_bytes() override {