This repository has no description
Something went wrong. Try again.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204diff --git a/common/common.cpp b/common/common.cppindex 4045903a..860545f4 100644--- a/common/common.cpp+++ b/common/common.cpp@@ -22,6 +22,7 @@ #include <fstream> #include <iostream> #include <iterator>+#include <mutex> #include <regex> #include <sstream> #include <string>@@ -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<std::vector<uint8_t>> ckpt_pool;++static void ckpt_buffer_take(std::vector<uint8_t> & buf, size_t size) {+ if (buf.capacity() >= size) {+ return;+ }+ {+ std::lock_guard<std::mutex> 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<uint8_t> & buf) {+ if (buf.capacity() < (1u << 20)) {+ return;+ }+ std::lock_guard<std::mutex> 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.hindex 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<uint8_t> 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.cppindex 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.hindex 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<llama_seq_id, std::string> sampler_sigs_reserved;+ ggml_backend_t backend_cpu = nullptr; std::vector<ggml_backend_ptr> backends; diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cppindex 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; } }