diff --git a/common/common.cpp b/common/common.cpp index 4045903a..860545f4 100644 --- a/common/common.cpp +++ b/common/common.cpp @@ -22,6 +22,7 @@ #include #include #include +#include #include #include #include @@ -2265,6 +2266,44 @@ bool common_prompt_batch_decode( return true; } +// State buffers of dropped checkpoints. A fresh buffer is page-faulted and zeroed before the state is +// copied in: ~65 ms of an ~85 ms checkpoint of a 150 MB recurrent state. +static std::mutex ckpt_pool_mutex; +static std::vector> ckpt_pool; + +static void ckpt_buffer_take(std::vector & buf, size_t size) { + if (buf.capacity() >= size) { + return; + } + { + std::lock_guard lock(ckpt_pool_mutex); + for (auto it = ckpt_pool.begin(); it != ckpt_pool.end(); ++it) { + if (it->capacity() >= size) { + buf = std::move(*it); + ckpt_pool.erase(it); + return; + } + } + } + // the state grows a little with the sequence; slack keeps a recycled buffer big enough for later checkpoints + buf.reserve(size + size/16); +} + +static void ckpt_buffer_give(std::vector & buf) { + if (buf.capacity() < (1u << 20)) { + return; + } + std::lock_guard lock(ckpt_pool_mutex); + if (ckpt_pool.size() < 4) { + ckpt_pool.push_back(std::move(buf)); + } +} + +common_prompt_checkpoint::~common_prompt_checkpoint() { + ckpt_buffer_give(data_tgt); + ckpt_buffer_give(data_dft); +} + size_t common_prompt_checkpoint::size() const { return data_tgt.size() + data_dft.size() + data_spec.size(); } @@ -2303,6 +2342,7 @@ void common_prompt_checkpoint::update_tgt( const size_t ckpt_size = llama_state_seq_get_size_ext(ctx, seq_id, flags); + ckpt_buffer_take(data_tgt, ckpt_size); data_tgt.resize(ckpt_size); const size_t n = llama_state_seq_get_data_ext(ctx, data_tgt.data(), ckpt_size, seq_id, flags); @@ -2321,6 +2361,7 @@ void common_prompt_checkpoint::update_dft( const size_t ckpt_size = llama_state_seq_get_size_ext(ctx, seq_id, flags); + ckpt_buffer_take(data_dft, ckpt_size); data_dft.resize(ckpt_size); const size_t n = llama_state_seq_get_data_ext(ctx, data_dft.data(), ckpt_size, seq_id, flags); diff --git a/common/common.h b/common/common.h index eb50a6cb..d796448d 100644 --- a/common/common.h +++ b/common/common.h @@ -1155,6 +1155,15 @@ struct common_prompt_checkpoint { // (e.g. eagle3's deferred-boundary g_embd row) std::vector data_spec; + common_prompt_checkpoint() = default; + common_prompt_checkpoint(const common_prompt_checkpoint &) = default; + common_prompt_checkpoint(common_prompt_checkpoint &&) = default; + common_prompt_checkpoint & operator=(const common_prompt_checkpoint &) = default; + common_prompt_checkpoint & operator=(common_prompt_checkpoint &&) = default; + + // returns the state buffers to a small pool that update_tgt/update_dft draw from + ~common_prompt_checkpoint(); + size_t size() const; bool empty() const; diff --git a/src/llama-context.cpp b/src/llama-context.cpp index cb6cf1f2..1d186472 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -696,6 +696,16 @@ void llama_context::resolve_fused_ops(const llama_memory_context_i * mctx, uint3 } } +// names of the samplers in a chain: two chains with the same names add the same ops to the graph +static std::string llama_sampler_chain_sig(llama_sampler * chain) { + std::string sig; + for (int i = 0; i < llama_sampler_chain_n(chain); ++i) { + sig += llama_sampler_name(llama_sampler_chain_get(chain, i)); + sig += ';'; + } + return sig; +} + void llama_context::sched_reserve() { if (!sched_need_reserve) { return; @@ -703,6 +713,10 @@ void llama_context::sched_reserve() { sched_need_reserve = false; + for (const auto & [seq_id, sampler] : sampling.samplers) { + sampler_sigs_reserved[seq_id] = llama_sampler_chain_sig(sampler); + } + LLAMA_LOG_INFO("%s: reserving ...\n", __func__); synchronize(); @@ -1461,7 +1475,11 @@ bool llama_context::set_sampler(llama_seq_id seq_id, llama_sampler * sampler) { sampling.samplers[seq_id] = sampler; - sched_need_reserve = true; + // a server sets a fresh chain for every request; reserving again for the same chain costs ~160 ms + const auto it = sampler_sigs_reserved.find(seq_id); + if (it == sampler_sigs_reserved.end() || it->second != llama_sampler_chain_sig(sampler)) { + sched_need_reserve = true; + } return true; } @@ -1478,10 +1496,9 @@ bool llama_context::set_sampler(llama_seq_id seq_id, llama_sampler * sampler) { return false; } + // removing a sampler only shrinks the graph, the reserved buffers still fit it sampling.samplers.erase(seq_id); - sched_need_reserve = true; - return true; } diff --git a/src/llama-context.h b/src/llama-context.h index 6057da7f..bced1881 100644 --- a/src/llama-context.h +++ b/src/llama-context.h @@ -372,6 +372,9 @@ private: bool sched_need_reserve = true; + // sampler chains (by sequence) whose backend ops the last reserves were sized for + std::map sampler_sigs_reserved; + ggml_backend_t backend_cpu = nullptr; std::vector backends; diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index 102da8f8..49d787c2 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -2239,6 +2239,16 @@ private: void create_checkpoint(server_slot & slot, const int64_t n_tokens_cur, llama_pos pos_min, llama_pos pos_max) { const int id_task = slot.task->id; + // a prompt that matches the cache up to its last checkpoint restores that checkpoint first, and + // writing an identical copy of it costs ~100 ms for a hybrid model's recurrent state + if (!slot.prompt.checkpoints.empty()) { + auto & last = slot.prompt.checkpoints.back(); + if (last.n_tokens == slot.prompt.n_tokens() - n_tokens_cur && last.pos_min == pos_min && last.pos_max == pos_max) { + last.id_task = id_task; + return; + } + } + // evict checkpoints within min-step of a previous checkpoint, unless they were // created by the current task int64_t last = -1; @@ -3460,12 +3470,15 @@ private: /* is_prompt = */ true); slot.prompt.tokens.push_back(cur_tok); - // break at the last user message, or at user messages at least min step past the last checkpoint + // break at the last user message, or at user messages at least min step past the last checkpoint. + // a last user message only a few tokens past the last checkpoint is cheaper to re-process than to + // split the batch for (a short follow-up turn would take three small batches instead of two) if (do_checkpoint && spans.is_user_start(slot.prompt.n_tokens())) { const auto pos = slot.prompt.n_tokens(); const auto & checkpoints = slot.prompt.checkpoints; - if (pos == last_user_pos || checkpoints.empty() || pos > checkpoints.back().n_tokens + params_base.checkpoint_min_step) { + const bool last_user_far = pos == last_user_pos && (checkpoints.empty() || pos > checkpoints.back().n_tokens + 256); + if (last_user_far || checkpoints.empty() || pos > checkpoints.back().n_tokens + params_base.checkpoint_min_step) { break; } }