# Applies to llama.cpp commit 9adc7f420c37641921b32e326b3d4a538256b878 # Tested with llama.cpp version 1892 # Tested on ROCm 10.0 diff --git a/common/common.h b/common/common.h index 9194d3dc..382913b4 100644 --- a/common/common.h +++ b/common/common.h @@ -1191,6 +1191,11 @@ struct common_prompt_checkpoint { std::vector data_tgt; std::vector data_dft; + // The actual storage mode used for each checkpoint. This can differ from + // the flags passed to update_* when an on-device save is not possible. + bool data_tgt_on_device = false; + bool data_dft_on_device = false; + // (optional) speculative-decoding implementation state stashed with the checkpoint // (e.g. eagle3's deferred-boundary g_embd row) std::vector data_spec; diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index e95fb63a..d04254d5 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -734,6 +734,196 @@ struct server_slot { } }; +// --------------------------------------------------------------------------- +// slot checkpoint sidecar persistence +// +// /slots save|restore persists the token list and the serialized KV/recurrent +// state, but not the server-side slot.prompt.checkpoints chain. Hybrid/recurrent +// prompt reuse requires that chain for rollback, so a saved slot would otherwise +// be re-prefilled from scratch on the next request. The sidecar file +// (.ckpt) persists the checkpoint chain next to the slot file. +// +// The header carries an FNV-1a 64 hash of the packed token list, so a stale +// sidecar (e.g. the server died after writing the slot file but before the +// sidecar) is rejected on restore instead of loading foreign recurrent state. +// +// Layout is native byte order; the sidecar is only valid on the host that wrote it. +// --------------------------------------------------------------------------- + +static const uint32_t SLOT_CKPT_SIDECAR_MAGIC = 0x4B434C53u; // 'SLCK' +static const uint32_t SLOT_CKPT_SIDECAR_VERSION = 2u; + +static uint64_t slot_ckpt_fnv1a(const uint8_t * p, size_t len) { + uint64_t h = 1469598103934665603ull; // FNV-1a 64 offset basis + for (size_t i = 0; i < len; ++i) { + h ^= p[i]; + h *= 1099511628211ull; // FNV-1a 64 prime + } + return h; +} + +static bool slot_ckpt_sidecar_save( + const std::string & path, + const void * tokens_data, size_t n_tokens_bytes, + const std::list & checkpoints) { + std::ofstream f(path, std::ios::binary | std::ios::trunc); + if (!f) { + return false; + } + + const uint32_t magic = SLOT_CKPT_SIDECAR_MAGIC; + const uint32_t version = SLOT_CKPT_SIDECAR_VERSION; + const uint32_t n = static_cast(checkpoints.size()); + const uint64_t hash = slot_ckpt_fnv1a(static_cast(tokens_data), n_tokens_bytes); + + auto w = [&f](const void * p, size_t len) { + f.write(reinterpret_cast(p), static_cast(len)); + }; + + w(&magic, sizeof(magic)); + w(&version, sizeof(version)); + w(&n, sizeof(n)); + w(&hash, sizeof(hash)); + + for (const auto & ckpt : checkpoints) { + const int64_t n_tokens = ckpt.n_tokens; + const int32_t id_task = ckpt.id_task; + const int32_t pos_min = ckpt.pos_min; + const int32_t pos_max = ckpt.pos_max; + const uint8_t flags = static_cast((ckpt.data_tgt_on_device ? 1u : 0u) | (ckpt.data_dft_on_device ? 2u : 0u)); + const uint64_t len_spec = ckpt.data_spec.size(); + const uint64_t len_tgt = ckpt.data_tgt.size(); + const uint64_t len_dft = ckpt.data_dft.size(); + + w(&n_tokens, sizeof(n_tokens)); + w(&id_task, sizeof(id_task)); + w(&pos_min, sizeof(pos_min)); + w(&pos_max, sizeof(pos_max)); + w(&flags, sizeof(flags)); + w(&len_spec, sizeof(len_spec)); + w(&len_tgt, sizeof(len_tgt)); + w(&len_dft, sizeof(len_dft)); + + if (len_spec > 0) w(ckpt.data_spec.data(), len_spec); + if (len_tgt > 0) w(ckpt.data_tgt.data(), len_tgt); + if (len_dft > 0) w(ckpt.data_dft.data(), len_dft); + } + + f.flush(); + f.close(); + return f.good(); +} + +static bool slot_ckpt_sidecar_restore( + const std::string & path, + const void * tokens_data, size_t n_tokens_bytes, + std::list & checkpoints, + size_t n_tokens_prompt, + size_t & n_restored) { + std::ifstream f(path, std::ios::binary); + if (!f) { + return false; // no sidecar + } + + // the total file size bounds every length field: a corrupt header can never + // drive an allocation beyond the actual file contents + f.seekg(0, std::ios::end); + const std::streamoff file_size = f.tellg(); + if (file_size <= 0) { + return false; + } + f.seekg(0, std::ios::beg); + + auto r = [&f](void * p, size_t len) { + f.read(reinterpret_cast(p), static_cast(len)); + return static_cast(f); + }; + + uint32_t magic = 0; + uint32_t version = 0; + uint32_t n = 0; + uint64_t hash = 0; + if (!r(&magic, sizeof(magic)) || !r(&version, sizeof(version)) || + !r(&n, sizeof(n)) || !r(&hash, sizeof(hash))) { + return false; + } + size_t off = sizeof(magic) + sizeof(version) + sizeof(n) + sizeof(hash); + if (magic != SLOT_CKPT_SIDECAR_MAGIC || version != SLOT_CKPT_SIDECAR_VERSION || n > 1024u) { + return false; + } + + // reject a sidecar that does not belong to this exact prompt (stale file) + if (hash != slot_ckpt_fnv1a(static_cast(tokens_data), n_tokens_bytes)) { + return false; + } + + std::list restored; + for (uint32_t i = 0; i < n; ++i) { + int64_t n_tokens = 0; + int32_t id_task = 0; + int32_t pos_min = 0; + int32_t pos_max = 0; + uint8_t flags = 0; + uint64_t len_spec = 0; + uint64_t len_tgt = 0; + uint64_t len_dft = 0; + + if (!r(&n_tokens, sizeof(n_tokens)) || !r(&id_task, sizeof(id_task)) || + !r(&pos_min, sizeof(pos_min)) || !r(&pos_max, sizeof(pos_max)) || + !r(&flags, sizeof(flags)) || + !r(&len_spec, sizeof(len_spec)) || !r(&len_tgt, sizeof(len_tgt)) || + !r(&len_dft, sizeof(len_dft))) { + return false; + } + off += sizeof(n_tokens) + sizeof(id_task) + sizeof(pos_min) + sizeof(pos_max) + + sizeof(flags) + sizeof(len_spec) + sizeof(len_tgt) + sizeof(len_dft); + if (off > static_cast(file_size)) { + return false; + } + const uint64_t remaining = static_cast(file_size) - off; + + // validate before installing: reject incompatible/stale data; + // the data lengths can never exceed the bytes actually left in the file + if (n_tokens <= 0 || + static_cast(n_tokens) > n_tokens_prompt || + pos_min < 0 || pos_min > pos_max || + static_cast(pos_max) >= n_tokens_prompt || + len_tgt == 0 || + len_spec > remaining || + len_tgt > remaining - len_spec || + len_dft > remaining - len_spec - len_tgt) { + return false; + } + + common_prompt_checkpoint ckpt; + ckpt.n_tokens = n_tokens; + ckpt.id_task = id_task; + ckpt.pos_min = pos_min; + ckpt.pos_max = pos_max; + ckpt.data_tgt_on_device = (flags & 1u) != 0; + ckpt.data_dft_on_device = (flags & 2u) != 0; + + ckpt.data_spec.resize(len_spec); + ckpt.data_tgt.resize(len_tgt); + ckpt.data_dft.resize(len_dft); + + if (len_spec > 0 && !r(ckpt.data_spec.data(), len_spec)) return false; + if (len_tgt > 0 && !r(ckpt.data_tgt.data(), len_tgt)) return false; + if (len_dft > 0 && !r(ckpt.data_dft.data(), len_dft)) return false; + off += static_cast(len_spec + len_tgt + len_dft); + + restored.push_back(std::move(ckpt)); + } + + if (off != static_cast(file_size)) { + return false; // truncated or trailing garbage + } + + n_restored = restored.size(); + checkpoints = std::move(restored); + return true; +} + // returns 0 on success // caller need to update prompt.tokens after a successful call to keep track of the processing progress // note: this is not a member of server_slot because we want to run it inside yield_to_queue @@ -1603,7 +1793,10 @@ private: } // if we are about to lose a large portion of the existing context - save it in the prompt cache - if (f_keep < 0.5f) { + // [TAG_PROMPT_CACHE_BEST_MATCH] also consult the cache when a cached prompt matches the + // incoming task better than the resident slot does (e.g. returning to a previous session + // while a short, prefix-similar session is still resident) + if (f_keep < 0.5f || (prompt_cache && prompt_cache->has_better(task.tokens, ret->prompt.tokens))) { update_cache = true; } } @@ -2579,6 +2772,20 @@ private: break; } + // persist the checkpoint chain in a sidecar file so /restore can rebuild prompt reuse + // (hybrid/recurrent models need the chain for rollback; without it the next + // request would re-prefill the whole prompt) + { + const std::string ckpt_path = filepath + ".ckpt"; + if (slot->prompt.checkpoints.empty()) { + std::remove(ckpt_path.c_str()); + } else if (slot_ckpt_sidecar_save(ckpt_path, packed.data(), packed.size(), slot->prompt.checkpoints)) { + SRV_INF("saved %zu context checkpoints to %s\n", slot->prompt.checkpoints.size(), ckpt_path.c_str()); + } else { + SRV_WRN("unable to save checkpoint sidecar to %s\n", ckpt_path.c_str()); + } + } + const int64_t t_end = ggml_time_us(); const double t_save_ms = (t_end - t_start) / 1000.0; @@ -2612,6 +2819,21 @@ private: std::string filename = task.slot_action.filename; std::string filepath = task.slot_action.filepath; + // [TAG_PROMPT_CACHE_RESTORE_PARK] a slot restore overwrites the slot's KV in + // place, bypassing get_available_slot() and therefore its prompt-cache save. + // With -np 1 that destroys the context of whichever session was resident, so a + // later request from that session has nothing to swap back to and re-prefills + // its whole unique tail. Park the displaced context in the RAM prompt cache + // first (same call get_available_slot() makes), so has_better() can bring it + // back on that session's next request. + if (prompt_cache && !slot->prompt.tokens.empty()) { + if (slot->prompt_save(*prompt_cache)) { + prompt_cache->update(); + SRV_INF("parked displaced context in prompt cache before slot restore (n_tokens = %zu)\n", + slot->prompt.tokens.size()); + } + } + size_t nread = 0; try { size_t n_packed = 0; @@ -2638,6 +2860,20 @@ private: slot->prompt.clear(); slot->prompt.tokens = std::move(restored); + + // rebuild the checkpoint chain from the sidecar file (if present) + { + const std::string ckpt_path = filepath + ".ckpt"; + size_t n_ckpt = 0; + // note: `packed` is a llama_tokens (vector), so size() is the token + // count; the sidecar hash is over the packed bytes (save hashes the + // vector whose size() is already bytes) + if (slot_ckpt_sidecar_restore(ckpt_path, packed.data(), packed.size() * sizeof(llama_token), slot->prompt.checkpoints, slot->prompt.tokens.size(), n_ckpt)) { + SRV_INF("restored %zu context checkpoints from %s\n", n_ckpt, ckpt_path.c_str()); + } else { + SRV_DBG("no valid checkpoint sidecar at %s (missing, stale, or corrupt)\n", ckpt_path.c_str()); + } + } } catch (const std::exception & err) { slot->prompt_clear(); send_error(task, std::string("Unable to restore slot: ") + err.what(), ERROR_TYPE_INVALID_REQUEST); diff --git a/tools/server/server-task.cpp b/tools/server/server-task.cpp index 0d3beb31..509e35eb 100644 --- a/tools/server/server-task.cpp +++ b/tools/server/server-task.cpp @@ -1867,6 +1867,33 @@ bool server_prompt_cache::load(server_prompt & prompt, const server_tokens & tok return true; } +// [TAG_PROMPT_CACHE_BEST_MATCH] probe the cache without committing: would load() pick a cached +// prompt over the currently resident one for the incoming tokens? +bool server_prompt_cache::has_better(const server_tokens & tokens_new, const server_tokens & tokens_cur) const { + const size_t lcp_base = tokens_cur.get_common_prefix(tokens_new); + + const float f_keep_best = tokens_cur.size() > 0 ? float(lcp_base) / tokens_cur.size() : -1.0f; // empty slot: any cache entry wins + const float f_sim_best = float(lcp_base) / tokens_new.size(); + + for (const auto & st : states) { + const size_t lcp_cur = st.prompt.tokens.get_common_prefix(tokens_new); + + const float f_keep_cur = st.prompt.tokens.size() > 0 ? float(lcp_cur) / st.prompt.tokens.size() : 0.0f; + const float f_sim_cur = float(lcp_cur) / tokens_new.size(); + + // mirrors load(): don't trash large prompts + if (f_keep_cur < 0.25f) { + continue; + } + + if (f_keep_best < f_keep_cur && f_sim_best < f_sim_cur) { + return true; + } + } + + return false; +} + void server_prompt_cache::update() { if (limit_size > 0) { while (!states.empty() && size() > limit_size) { diff --git a/tools/server/server-task.h b/tools/server/server-task.h index 9c99143f..eca2dd42 100644 --- a/tools/server/server-task.h +++ b/tools/server/server-task.h @@ -631,6 +631,10 @@ struct server_prompt_cache { bool load(server_prompt & prompt, const server_tokens & tokens_new, llama_context * ctx_tgt, llama_context * ctx_dft, int32_t id_slot); + // [TAG_PROMPT_CACHE_BEST_MATCH] true when some cached prompt would be loaded in preference to + // the currently resident prompt for the given incoming tokens (mirrors the selection criteria of load()) + bool has_better(const server_tokens & tokens_new, const server_tokens & tokens_cur) const; + void update(); };