# Applies to llama.cpp commit 3737e41 # Includes MTP Compact Rollback commit 06c1414f; excludes cumulative ROCm, pipeline, HIP VEC, and MoE changes # Tested with llama.cpp version 1266 (CPU/default build) # Parser plus Qwen, Nemotron-H, and DeepSeek-V4 recurrent rollback tests passed diff --git a/common/arg.cpp b/common/arg.cpp index 86f8610a..30b98654 100644 --- a/common/arg.cpp +++ b/common/arg.cpp @@ -1291,6 +1291,9 @@ bool common_params_parse(int argc, char ** argv, common_params & params, llama_e ctx_arg.params = params_org; return false; } + if (ctx_arg.params.speculative.draft.n_rs_seq > ctx_arg.params.speculative.draft.n_max) { + throw std::invalid_argument("--spec-mtp-cr-depth must not exceed --spec-draft-n-max"); + } if (ctx_arg.params.usage) { common_params_print_usage(ctx_arg); if (ctx_arg.print_usage) { @@ -4102,6 +4105,16 @@ common_params_context common_params_parser_init(common_params & params, llama_ex params.speculative.draft.n_max = value; } ).set_spec().set_examples({LLAMA_EXAMPLE_SPECULATIVE, LLAMA_EXAMPLE_LOOKUP, LLAMA_EXAMPLE_SERVER, LLAMA_EXAMPLE_CLI}).set_env("LLAMA_ARG_SPEC_DRAFT_N_MAX")); + add_opt(common_arg( + {"--spec-mtp-cr-depth"}, "N", + "MTP Compact Rollback depth; lower values save memory but replay accepted tokens after deep rejection (default: --spec-draft-n-max)", + [](common_params & params, int value) { + if (value < 1) { + throw std::invalid_argument("--spec-mtp-cr-depth must be at least 1"); + } + params.speculative.draft.n_rs_seq = value; + } + ).set_spec().set_examples({LLAMA_EXAMPLE_SPECULATIVE, LLAMA_EXAMPLE_SERVER, LLAMA_EXAMPLE_CLI}).set_env("LLAMA_ARG_SPEC_MTP_CR_DEPTH")); add_opt(common_arg( {"--spec-draft-n-min"}, "N", string_format("minimum number of draft tokens to use for speculative decoding (default: %d)", params.speculative.draft.n_min), diff --git a/common/common.cpp b/common/common.cpp index 3d54bd60..3ab2598e 100644 --- a/common/common.cpp +++ b/common/common.cpp @@ -2268,6 +2268,8 @@ void common_prompt_checkpoint::clear() { data_tgt.clear(); data_dft.clear(); data_spec.clear(); + data_tgt_on_device = false; + data_dft_on_device = false; } void common_prompt_checkpoint::update_pos( @@ -2279,6 +2281,50 @@ void common_prompt_checkpoint::update_pos( this->pos_max = pos_max; } +static void common_prompt_checkpoint_save( + std::vector & data, + bool & on_device, + llama_context * ctx, + llama_seq_id seq_id, + llama_state_seq_flags flags, + const char * label) { + const bool requested_on_device = flags & LLAMA_STATE_SEQ_FLAGS_ON_DEVICE; + on_device = false; + + auto save = [&](llama_state_seq_flags save_flags) { + const size_t ckpt_size = llama_state_seq_get_size_ext(ctx, seq_id, save_flags); + if (ckpt_size == 0) { + return false; + } + + std::vector saved(ckpt_size); + const size_t n = llama_state_seq_get_data_ext(ctx, saved.data(), ckpt_size, seq_id, save_flags); + if (n != ckpt_size) { + return false; + } + + data.swap(saved); + return true; + }; + + bool saved = save(flags); + if (requested_on_device && !saved) { + COM_WRN("%s: ON_DEVICE %s checkpoint save failed; retrying with host storage\n", + __func__, label); + flags &= ~LLAMA_STATE_SEQ_FLAGS_ON_DEVICE; + saved = save(flags); + } + + if (!saved) { + if (flags & LLAMA_STATE_SEQ_FLAGS_ON_DEVICE) { + GGML_ABORT("checkpoint device save failed for %s\n", label); + } + GGML_ABORT("checkpoint size mismatch while saving %s\n", label); + } + + on_device = flags & LLAMA_STATE_SEQ_FLAGS_ON_DEVICE; +} + void common_prompt_checkpoint::update_tgt( llama_context * ctx, llama_seq_id seq_id, @@ -2287,14 +2333,7 @@ void common_prompt_checkpoint::update_tgt( return; } - const size_t ckpt_size = llama_state_seq_get_size_ext(ctx, seq_id, flags); - - data_tgt.resize(ckpt_size); - - const size_t n = llama_state_seq_get_data_ext(ctx, data_tgt.data(), ckpt_size, seq_id, flags); - if (n != ckpt_size) { - GGML_ABORT("checkpoint size mismatch: expected %zu, got %zu\n", ckpt_size, n); - } + common_prompt_checkpoint_save(data_tgt, data_tgt_on_device, ctx, seq_id, flags, "target"); } void common_prompt_checkpoint::update_dft( @@ -2305,14 +2344,7 @@ void common_prompt_checkpoint::update_dft( return; } - const size_t ckpt_size = llama_state_seq_get_size_ext(ctx, seq_id, flags); - - data_dft.resize(ckpt_size); - - const size_t n = llama_state_seq_get_data_ext(ctx, data_dft.data(), ckpt_size, seq_id, flags); - if (n != ckpt_size) { - GGML_ABORT("checkpoint size mismatch: expected %zu, got %zu\n", ckpt_size, n); - } + common_prompt_checkpoint_save(data_dft, data_dft_on_device, ctx, seq_id, flags, "draft"); } void common_prompt_checkpoint::load_tgt( @@ -2327,6 +2359,9 @@ void common_prompt_checkpoint::load_tgt( return; } + flags = (flags & ~LLAMA_STATE_SEQ_FLAGS_ON_DEVICE) | + (data_tgt_on_device ? LLAMA_STATE_SEQ_FLAGS_ON_DEVICE : 0); + const size_t n = llama_state_seq_set_data_ext(ctx, data_tgt.data(), data_tgt.size(), seq_id, flags); if (n != data_tgt.size()) { GGML_ABORT("checkpoint size mismatch: expected %zu, got %zu\n", data_tgt.size(), n); @@ -2345,6 +2380,9 @@ void common_prompt_checkpoint::load_dft( return; } + flags = (flags & ~LLAMA_STATE_SEQ_FLAGS_ON_DEVICE) | + (data_dft_on_device ? LLAMA_STATE_SEQ_FLAGS_ON_DEVICE : 0); + const size_t n = llama_state_seq_set_data_ext(ctx, data_dft.data(), data_dft.size(), seq_id, flags); if (n != data_dft.size()) { GGML_ABORT("checkpoint size mismatch: expected %zu, got %zu\n", data_dft.size(), n); @@ -2353,9 +2391,11 @@ void common_prompt_checkpoint::load_dft( void common_prompt_checkpoint::clear_tgt() { data_tgt.clear(); + data_tgt_on_device = false; } void common_prompt_checkpoint::clear_dft() { data_dft.clear(); data_spec.clear(); + data_dft_on_device = false; } diff --git a/common/common.h b/common/common.h index de49dac9..acb13ba8 100644 --- a/common/common.h +++ b/common/common.h @@ -324,6 +324,7 @@ struct common_params_model { struct common_params_speculative_draft { int32_t n_max = 3; // maximum number of tokens to draft during speculative decoding int32_t n_min = 0; // minimum number of draft tokens to use for speculative decoding + int32_t n_rs_seq = -1; // MTP Compact Rollback depth (-1 = n_max) float p_split = 0.1f; // speculative decoding split probability float p_min = 0.0f; // minimum speculative decoding probability (greedy) @@ -384,11 +385,18 @@ struct common_params_speculative { } uint32_t need_n_rs_seq() const { - bool needs_rs_seq = std::any_of(types.begin(), types.end(), [&](auto t) { - return t == COMMON_SPECULATIVE_TYPE_DRAFT_MTP || t == COMMON_SPECULATIVE_TYPE_DRAFT_EAGLE3 || t == COMMON_SPECULATIVE_TYPE_DRAFT_DFLASH || t == COMMON_SPECULATIVE_TYPE_DRAFT_DSPARK; + const bool needs_mtp = std::find(types.begin(), types.end(), COMMON_SPECULATIVE_TYPE_DRAFT_MTP) != types.end(); + const bool needs_other_rs = std::any_of(types.begin(), types.end(), [&](auto t) { + return t == COMMON_SPECULATIVE_TYPE_DRAFT_EAGLE3 || t == COMMON_SPECULATIVE_TYPE_DRAFT_DFLASH || t == COMMON_SPECULATIVE_TYPE_DRAFT_DSPARK; }); - return needs_rs_seq ? draft.n_max : 0u; + if (needs_other_rs) { + return draft.n_max; + } + if (needs_mtp) { + return draft.n_rs_seq >= 0 ? draft.n_rs_seq : draft.n_max; + } + return 0u; } }; @@ -1146,6 +1154,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/common/fit.cpp b/common/fit.cpp index c601fe40..e85287e0 100644 --- a/common/fit.cpp +++ b/common/fit.cpp @@ -34,7 +34,8 @@ static std::vector common_get_device_memory_data_impl( uint32_t & hp_ngl, uint32_t & hp_n_ctx_train, uint32_t & hp_n_expert, - ggml_log_level log_level) { + ggml_log_level log_level, + std::vector * checkpoint_sizes = nullptr) { struct user_data_t { struct { ggml_log_callback callback; @@ -96,6 +97,28 @@ static std::vector common_get_device_memory_data_impl( } } + if (checkpoint_sizes) { + checkpoint_sizes->assign(nd + 1, 0); + const auto checkpoint_breakdown = llama_get_state_seq_device_buffer_sizes(ctx); + for (const auto & [buft, size] : checkpoint_breakdown) { + if (ggml_backend_buft_is_host(buft)) { + checkpoint_sizes->back() += size; + continue; + } + + ggml_backend_dev_t dev = ggml_backend_buft_get_device(buft); + if (!dev) { + continue; + } + for (size_t i = 0; i < nd; ++i) { + if (dev == llama_model_get_device(model, i)) { + (*checkpoint_sizes)[i] += size; + break; + } + } + } + } + { ggml_backend_dev_t cpu_dev = ggml_backend_dev_by_type(GGML_BACKEND_DEVICE_TYPE_CPU); if (cpu_dev == nullptr) { @@ -161,8 +184,10 @@ common_device_memory_data_vec common_get_device_memory_data( uint32_t & hp_n_ctx_train, uint32_t & hp_n_expert, ggml_log_level log_level) { + std::vector checkpoint_sizes; std::vector impl = common_get_device_memory_data_impl( - path_model, mparams, cparams, devs, hp_ngl, hp_n_ctx_train, hp_n_expert, log_level); + path_model, mparams, cparams, devs, hp_ngl, hp_n_ctx_train, hp_n_expert, log_level, + &checkpoint_sizes); common_device_memory_data_vec ret(impl.size()); for (size_t i = 0; i < impl.size(); i++) { @@ -171,6 +196,7 @@ common_device_memory_data_vec common_get_device_memory_data( ret[i].model = impl[i].mb.model; ret[i].context = impl[i].mb.context; ret[i].compute = impl[i].mb.compute; + ret[i].checkpoint = checkpoint_sizes[i]; } return ret; } diff --git a/common/fit.h b/common/fit.h index 824d386b..7022746a 100644 --- a/common/fit.h +++ b/common/fit.h @@ -51,6 +51,7 @@ struct common_device_memory_data { size_t model; size_t context; size_t compute; + size_t checkpoint; }; using common_device_memory_data_vec = std::vector; diff --git a/include/llama.h b/include/llama.h index a04177f9..03076b26 100644 --- a/include/llama.h +++ b/include/llama.h @@ -920,6 +920,11 @@ extern "C" { llama_seq_id seq_id, llama_state_seq_flags flags); + // Allocate and pin all configured per-sequence device checkpoint buffers. + // Returns false when the context has no device checkpoint layout or an + // allocation fails. + LLAMA_API bool llama_state_seq_reserve_device_buffers(struct llama_context * ctx); + LLAMA_API size_t llama_state_seq_set_data_ext( struct llama_context * ctx, const uint8_t * src, diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 0402044d..cb8f6c2a 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -2715,10 +2715,12 @@ private: class llama_io_write_device : public llama_io_write_i { public: - llama_io_write_device(uint8_t * p, size_t len, llama_memory_buffers & mbufs) : ptr(p), buf_size(len), mbufs(mbufs) { + llama_io_write_device( + uint8_t * p, size_t len, llama_memory_buffers & mbufs, bool copy_tensors = true) : + ptr(p), buf_size(len), mbufs(mbufs), copy_tensors(copy_tensors) { } - ~llama_io_write_device() { + void finish() { llama_memory_buffers mbufs_new; for (const auto & winfo : winfos) { @@ -2780,11 +2782,21 @@ public: if (need_alloc) { if (!mbuf_cur.buf || mbuf_cur.total_size != mbuf.total_size) { - mbuf_cur = std::move(mbuf); + ggml_backend_buffer_ptr buf { + ggml_backend_alloc_ctx_tensors_from_buft(mbuf.ctx.get(), buft) + }; + if (!buf) { + throw std::runtime_error( + std::string("failed to allocate device checkpoint buffer from '") + + ggml_backend_buft_name(buft) + "'"); + } - mbuf_cur.buf.reset(ggml_backend_alloc_ctx_tensors_from_buft(mbuf_cur.ctx.get(), buft)); + const size_t allocated_size = ggml_backend_buffer_get_size(buf.get()); + mbuf.buf = std::move(buf); + mbuf_cur = std::move(mbuf); - LLAMA_LOG_INFO("%s: allocated '%s' buffer %.3f MiB\n", __func__, ggml_backend_buft_name(buft), mbuf.total_size/1024.0/1024.0); + LLAMA_LOG_INFO("%s: allocated '%s' buffer %.3f MiB\n", __func__, + ggml_backend_buft_name(buft), allocated_size/1024.0/1024.0); } else { //LLAMA_LOG_INFO("%s: reallocating tensors in '%s' buffer %.3f MiB\n", __func__, ggml_backend_buft_name(buft), mbuf.total_size/1024.0/1024.0); @@ -2804,8 +2816,10 @@ public: } } - for (size_t i = 0; i < mbuf_cur.org.size(); ++i) { - ggml_backend_tensor_copy(mbuf_cur.org[i], mbuf_cur.cpy[i]); + if (copy_tensors) { + for (size_t i = 0; i < mbuf_cur.org.size(); ++i) { + ggml_backend_tensor_copy(mbuf_cur.org[i], mbuf_cur.cpy[i]); + } } } } @@ -2814,14 +2828,16 @@ public: if (size > buf_size) { throw std::runtime_error("unexpectedly reached end of buffer"); } - memcpy(ptr, src, size); - ptr += size; + if (ptr != nullptr) { + memcpy(ptr, src, size); + ptr += size; + } size_written += size; buf_size -= size; } void write_tensor(ggml_tensor * tensor, size_t offset, size_t size) override { - // save the write for later during destruction + // save the write until finish(), after all state tensors are known winfos.push_back({tensor, ptr, size, offset}); } @@ -2843,6 +2859,7 @@ private: std::vector winfos; llama_memory_buffers & mbufs; + const bool copy_tensors; }; class llama_io_read_device : public llama_io_read_i { @@ -2981,8 +2998,11 @@ size_t llama_context::state_seq_get_size(llama_seq_id seq_id, llama_state_seq_fl size_t llama_context::state_seq_get_data(llama_seq_id seq_id, uint8_t * dst, size_t size, llama_state_seq_flags flags) { std::unique_ptr io; + llama_io_write_device * io_device = nullptr; if (flags & LLAMA_STATE_SEQ_FLAGS_ON_DEVICE) { - io = std::make_unique(dst, size, mem_storage[seq_id]); + auto device = std::make_unique(dst, size, mem_storage[seq_id]); + io_device = device.get(); + io = std::move(device); } else { io = std::make_unique(dst, size); } @@ -2991,7 +3011,11 @@ size_t llama_context::state_seq_get_data(llama_seq_id seq_id, uint8_t * dst, siz io->write(&io_magic, sizeof(io_magic)); io->write(&seq_id, sizeof(seq_id)); - return state_seq_write_data(*io, seq_id, flags); + const size_t n = state_seq_write_data(*io, seq_id, flags); + if (io_device) { + io_device->finish(); + } + return n; } catch (const std::exception & err) { LLAMA_LOG_ERROR("%s: error saving state: %s\n", __func__, err.what()); return 0; @@ -3284,6 +3308,30 @@ llama_memory_breakdown llama_context::memory_breakdown() const { return ret; } +std::map llama_context::state_seq_device_buffer_sizes() const { + return memory ? memory->state_seq_device_buffer_sizes() + : std::map {}; +} + +bool llama_context::state_seq_reserve_device_buffers() { + if (!memory || memory->state_seq_device_buffer_sizes().empty()) { + return false; + } + + try { + for (llama_seq_id seq_id = 0; seq_id < static_cast(cparams.n_seq_max); ++seq_id) { + llama_io_write_device io(nullptr, std::numeric_limits::max(), mem_storage[seq_id], false); + memory->state_seq_write_device_layout(io); + io.finish(); + } + return true; + } catch (const std::exception & err) { + mem_storage.clear(); + LLAMA_LOG_ERROR("%s: failed to reserve device checkpoint buffers: %s\n", __func__, err.what()); + return false; + } +} + // // training // @@ -4210,6 +4258,15 @@ llama_memory_breakdown llama_get_memory_breakdown(const struct llama_context * c return ctx->memory_breakdown(); } +std::map llama_get_state_seq_device_buffer_sizes( + const struct llama_context * ctx) { + return ctx->state_seq_device_buffer_sizes(); +} + +bool llama_state_seq_reserve_device_buffers(struct llama_context * ctx) { + return ctx->state_seq_reserve_device_buffers(); +} + llama_context * llama_get_ctx_other(struct llama_context * ctx) { return ctx->get_cparams().ctx_other; } diff --git a/src/llama-context.h b/src/llama-context.h index bf91daa8..83273234 100644 --- a/src/llama-context.h +++ b/src/llama-context.h @@ -188,6 +188,8 @@ struct llama_context { void perf_reset(); llama_memory_breakdown memory_breakdown() const; + std::map state_seq_device_buffer_sizes() const; + bool state_seq_reserve_device_buffers(); // // training diff --git a/src/llama-ext.h b/src/llama-ext.h index 35d6e58a..1eb87f53 100644 --- a/src/llama-ext.h +++ b/src/llama-ext.h @@ -90,6 +90,11 @@ LLAMA_API ggml_backend_dev_t llama_model_get_device(const struct llama_model * m LLAMA_API llama_memory_breakdown llama_get_memory_breakdown(const struct llama_context * ctx); +// Exact backend-buffer sizes needed for one device-resident sequence-state +// checkpoint per configured sequence. +LLAMA_API std::map llama_get_state_seq_device_buffer_sizes( + const struct llama_context * ctx); + // Set whether the context outputs nextn embeddings or not // If masked == true, output the embeddings only for the tokens with batch.logits != 0 // If masked == false, output the embeddings for all tokens in the batch regardless of batch.logits diff --git a/src/llama-memory-hybrid-iswa.cpp b/src/llama-memory-hybrid-iswa.cpp index 06f7fd54..408982b7 100644 --- a/src/llama-memory-hybrid-iswa.cpp +++ b/src/llama-memory-hybrid-iswa.cpp @@ -192,6 +192,14 @@ std::map llama_memory_hybrid_iswa::memory_br return mb; } +std::map llama_memory_hybrid_iswa::state_seq_device_buffer_sizes() const { + return mem_recr->state_seq_device_buffer_sizes(); +} + +void llama_memory_hybrid_iswa::state_seq_write_device_layout(llama_io_write_i & io) const { + mem_recr->state_seq_write_device_layout(io); +} + void llama_memory_hybrid_iswa::state_write(llama_io_write_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) const { mem_attn->state_write(io, seq_id, flags); mem_recr->state_write(io, seq_id, flags); diff --git a/src/llama-memory-hybrid-iswa.h b/src/llama-memory-hybrid-iswa.h index c9d3f9f5..02e5ce33 100644 --- a/src/llama-memory-hybrid-iswa.h +++ b/src/llama-memory-hybrid-iswa.h @@ -70,6 +70,8 @@ public: llama_pos seq_pos_max(llama_seq_id seq_id) const override; std::map memory_breakdown() const override; + std::map state_seq_device_buffer_sizes() const override; + void state_seq_write_device_layout(llama_io_write_i & io) const override; // state write/load diff --git a/src/llama-memory-hybrid.cpp b/src/llama-memory-hybrid.cpp index 42c7381a..b05cb1a9 100644 --- a/src/llama-memory-hybrid.cpp +++ b/src/llama-memory-hybrid.cpp @@ -187,6 +187,14 @@ std::map llama_memory_hybrid::memory_breakdo return mb; } +std::map llama_memory_hybrid::state_seq_device_buffer_sizes() const { + return mem_recr->state_seq_device_buffer_sizes(); +} + +void llama_memory_hybrid::state_seq_write_device_layout(llama_io_write_i & io) const { + mem_recr->state_seq_write_device_layout(io); +} + void llama_memory_hybrid::state_write(llama_io_write_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) const { if ((flags & LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY) == 0) { mem_attn->state_write(io, seq_id, flags); diff --git a/src/llama-memory-hybrid.h b/src/llama-memory-hybrid.h index 484eafb7..5818fab4 100644 --- a/src/llama-memory-hybrid.h +++ b/src/llama-memory-hybrid.h @@ -70,6 +70,8 @@ public: llama_pos seq_pos_max(llama_seq_id seq_id) const override; std::map memory_breakdown() const override; + std::map state_seq_device_buffer_sizes() const override; + void state_seq_write_device_layout(llama_io_write_i & io) const override; // state write/load diff --git a/src/llama-memory-recurrent.cpp b/src/llama-memory-recurrent.cpp index e2990972..0b87407b 100644 --- a/src/llama-memory-recurrent.cpp +++ b/src/llama-memory-recurrent.cpp @@ -1,5 +1,6 @@ #include "llama-memory-recurrent.h" +#include "ggml-alloc.h" #include "ggml-backend.h" #include "llama-impl.h" #include "llama-io.h" @@ -414,6 +415,64 @@ std::map llama_memory_recurrent::memory_brea return ret; } +std::map llama_memory_recurrent::state_seq_device_buffer_sizes() const { + struct layout { + size_t n_tensors = 0; + std::vector> tensors; + }; + + std::map layouts; + const uint32_t n_layer = hparams.n_layer(); + + auto add_row = [&](ggml_tensor * tensor, uint32_t n_embd) { + if (tensor == nullptr) { + return; + } + const uint64_t row_size = ggml_row_size(tensor->type, n_embd); + auto * buft = ggml_backend_buffer_get_type(tensor->buffer); + auto & cur = layouts[buft]; + cur.n_tensors++; + cur.tensors.emplace_back(tensor->type, row_size / ggml_element_size(tensor)); + }; + + for (uint32_t il = 0; il < n_layer; ++il) { + add_row(r_l[il], hparams.n_embd_r()); + } + for (uint32_t il = 0; il < n_layer; ++il) { + add_row(s_l[il], hparams.n_embd_s()); + } + + std::map ret; + for (const auto & [buft, layout] : layouts) { + ggml_init_params params = { + /*.mem_size =*/ layout.n_tensors * ggml_tensor_overhead(), + /*.mem_buffer =*/ nullptr, + /*.no_alloc =*/ true, + }; + ggml_context_ptr ctx { ggml_init(params) }; + if (!ctx) { + throw std::runtime_error("failed to create recurrent checkpoint sizing context"); + } + for (const auto & [type, n] : layout.tensors) { + ggml_new_tensor_1d(ctx.get(), type, n); + } + + const size_t one_seq = ggml_backend_alloc_ctx_tensors_from_buft_size(ctx.get(), buft); + if (n_seq_max != 0 && one_seq > std::numeric_limits::max() / n_seq_max) { + throw std::runtime_error("recurrent checkpoint size overflow"); + } + ret[buft] = one_seq * n_seq_max; + } + return ret; +} + +void llama_memory_recurrent::state_seq_write_device_layout(llama_io_write_i & io) const { + // One logical recurrent row is the complete per-sequence checkpoint. The + // actual save may select a rollback plane, but its tensor shapes and buffer + // sizes are identical to row zero. + state_write_data(io, {{0, 1}}); +} + llama_memory_context_ptr llama_memory_recurrent::init_batch(llama_batch_allocr & balloc, uint32_t n_ubatch, bool embd_all) { do { balloc.split_reset(); @@ -799,7 +858,7 @@ void llama_memory_recurrent::state_write(llama_io_write_i & io, llama_seq_id seq } if ((flags & LLAMA_STATE_SEQ_FLAGS_ON_DEVICE) && cell_ranges.size() > 1) { - GGML_ABORT("cannot save/load multiple ranges of cells to/from device memory\n"); + throw std::runtime_error("cannot save/load multiple ranges of cells to/from device memory"); } // DEBUG CHECK: Sum of cell counts in ranges should equal the total cell count diff --git a/src/llama-memory-recurrent.h b/src/llama-memory-recurrent.h index b13b7b74..270a4857 100644 --- a/src/llama-memory-recurrent.h +++ b/src/llama-memory-recurrent.h @@ -53,6 +53,8 @@ public: llama_pos seq_pos_max(llama_seq_id seq_id) const override; std::map memory_breakdown() const override; + std::map state_seq_device_buffer_sizes() const override; + void state_seq_write_device_layout(llama_io_write_i & io) const override; bool prepare(const std::vector & ubatches); diff --git a/src/llama-memory.h b/src/llama-memory.h index db825396..374d49d4 100644 --- a/src/llama-memory.h +++ b/src/llama-memory.h @@ -118,6 +118,15 @@ struct llama_memory_i { virtual std::map memory_breakdown() const = 0; + // Exact backend-buffer sizes needed to keep one sequence-state checkpoint + // per configured sequence. Memory types without device checkpoint data + // return an empty map. + virtual std::map state_seq_device_buffer_sizes() const { return {}; } + + // Describe the tensor data for one device-resident sequence checkpoint. + // This is used to allocate the real checkpoint buffers before evaluation. + virtual void state_seq_write_device_layout(llama_io_write_i & io) const { GGML_UNUSED(io); } + // // state write/read // diff --git a/tests/test-arg-parser.cpp b/tests/test-arg-parser.cpp index ba58f852..8f5aa4a9 100644 --- a/tests/test-arg-parser.cpp +++ b/tests/test-arg-parser.cpp @@ -3,6 +3,7 @@ #include "download.h" #include "llama.h" #include "speculative.h" +#include "ggml-backend.h" #include #include @@ -197,6 +198,132 @@ static void test(void) { assert(true == common_params_parse(argv.size(), list_str_to_char(argv).data(), params, LLAMA_EXAMPLE_SPECULATIVE)); assert(params.speculative.draft.n_max == 123); + { + common_params params_mtp; + argv = {"binary_name", "-m", "abc.gguf", "--spec-draft-n-max", "5", "--spec-mtp-cr-depth", "1"}; + assert(true == common_params_parse(argv.size(), list_str_to_char(argv).data(), params_mtp, LLAMA_EXAMPLE_SPECULATIVE)); + assert(params_mtp.speculative.draft.n_max == 5); + assert(params_mtp.speculative.draft.n_rs_seq == 1); + + params_mtp.speculative.types = { COMMON_SPECULATIVE_TYPE_DRAFT_MTP }; + assert(params_mtp.speculative.need_n_rs_seq() == 1); + + params_mtp.speculative.types.push_back(COMMON_SPECULATIVE_TYPE_DRAFT_EAGLE3); + assert(params_mtp.speculative.need_n_rs_seq() == 5); + } + + { + common_params params_mtp; + argv = {"binary_name", "-m", "abc.gguf", "--spec-draft-n-max", "5"}; + assert(true == common_params_parse(argv.size(), list_str_to_char(argv).data(), params_mtp, LLAMA_EXAMPLE_SPECULATIVE)); + assert(params_mtp.speculative.draft.n_rs_seq == -1); + + params_mtp.speculative.types = { COMMON_SPECULATIVE_TYPE_DRAFT_MTP }; + assert(params_mtp.speculative.need_n_rs_seq() == 5); + } + + { + common_params params_mtp; + argv = {"binary_name", "-m", "abc.gguf", "--spec-mtp-cr-depth", "1", "--spec-draft-n-max", "5"}; + assert(true == common_params_parse(argv.size(), list_str_to_char(argv).data(), params_mtp, LLAMA_EXAMPLE_SPECULATIVE)); + assert(params_mtp.speculative.draft.n_rs_seq == 1); + } + + { + common_params params_mtp; + argv = {"binary_name", "-m", "abc.gguf", "--spec-draft-n-max", "5", "--spec-mtp-cr-depth", "0"}; + assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), params_mtp, LLAMA_EXAMPLE_SPECULATIVE)); + } + + { + common_params params_mtp; + argv = {"binary_name", "-m", "abc.gguf", "--spec-draft-n-max", "5", "--spec-mtp-cr-depth=-1"}; + assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), params_mtp, LLAMA_EXAMPLE_SPECULATIVE)); + } + + { + common_params params_mtp; + argv = {"binary_name", "-m", "abc.gguf", "--spec-mtp-cr-depth", "6", "--spec-draft-n-max", "5"}; + assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), params_mtp, LLAMA_EXAMPLE_SPECULATIVE)); + } + + { + printf("test-arg-parser: test recurrent checkpoint reservation inputs\n\n"); + + common_params fit_params; + assert(fit_params.fit_params); + assert(fit_params.fit_params_target.size() == llama_max_devices()); + for (const size_t target : fit_params.fit_params_target) { + assert(target == 1024 * 1024 * 1024ULL); + } + + argv = {"binary_name", "-m", "abc.gguf", "--fit", "off", "--split-mode", "none", "--main-gpu", "1", + "--fit-target", "256,64", "--tensor-split", "3,1"}; + assert(true == common_params_parse(argv.size(), list_str_to_char(argv).data(), fit_params, LLAMA_EXAMPLE_COMMON)); + assert(!fit_params.fit_params); + assert(fit_params.split_mode == LLAMA_SPLIT_MODE_NONE); + assert(fit_params.main_gpu == 1); + assert(fit_params.fit_params_target[0] == 256 * 1024 * 1024ULL); + assert(fit_params.fit_params_target[1] == 64 * 1024 * 1024ULL); + assert(fit_params.tensor_split[0] == 3.0f); + assert(fit_params.tensor_split[1] == 1.0f); + + auto checkpoint_reservation_needed = [](const common_params & p) { + return p.speculative.need_n_rs_seq() < static_cast(p.speculative.draft.n_max); + }; + + common_params mtp_full; + argv = {"binary_name", "-m", "abc.gguf", "--spec-draft-n-max", "5"}; + assert(true == common_params_parse(argv.size(), list_str_to_char(argv).data(), mtp_full, LLAMA_EXAMPLE_SPECULATIVE)); + mtp_full.speculative.types = { COMMON_SPECULATIVE_TYPE_DRAFT_MTP }; + assert(mtp_full.speculative.need_n_rs_seq() == 5); + assert(!checkpoint_reservation_needed(mtp_full)); + + common_params mtp_reduced; + argv = {"binary_name", "-m", "abc.gguf", "--spec-draft-n-max", "5", "--spec-mtp-cr-depth", "2"}; + assert(true == common_params_parse(argv.size(), list_str_to_char(argv).data(), mtp_reduced, LLAMA_EXAMPLE_SPECULATIVE)); + mtp_reduced.speculative.types = { COMMON_SPECULATIVE_TYPE_DRAFT_MTP }; + assert(mtp_reduced.speculative.need_n_rs_seq() == 2); + assert(checkpoint_reservation_needed(mtp_reduced)); + + common_params mtp_with_other_rs; + argv = {"binary_name", "-m", "abc.gguf", "--spec-draft-n-max", "5", "--spec-mtp-cr-depth", "2"}; + assert(true == common_params_parse(argv.size(), list_str_to_char(argv).data(), mtp_with_other_rs, LLAMA_EXAMPLE_SPECULATIVE)); + mtp_with_other_rs.speculative.types = { + COMMON_SPECULATIVE_TYPE_DRAFT_MTP, + COMMON_SPECULATIVE_TYPE_DRAFT_EAGLE3, + }; + assert(mtp_with_other_rs.speculative.need_n_rs_seq() == 5); + assert(!checkpoint_reservation_needed(mtp_with_other_rs)); + + ggml_backend_load_all(); + std::vector gpu_devices; + for (size_t i = 0; i < ggml_backend_dev_count(); ++i) { + ggml_backend_dev_t dev = ggml_backend_dev_get(i); + if (ggml_backend_dev_type(dev) != GGML_BACKEND_DEVICE_TYPE_CPU) { + gpu_devices.push_back(dev); + } + } + if (gpu_devices.size() >= 2) { + const std::string device_arg = string_format("%s,%s", + ggml_backend_dev_name(gpu_devices[0]), ggml_backend_dev_name(gpu_devices[1])); + common_params multi_gpu; + argv = {"binary_name", "-m", "abc.gguf", "--device", device_arg, "--split-mode", "layer", + "--fit-target", "128,256", "--tensor-split", "3,1"}; + assert(true == common_params_parse(argv.size(), list_str_to_char(argv).data(), multi_gpu, LLAMA_EXAMPLE_COMMON)); + assert(multi_gpu.devices.size() == 3); + assert(multi_gpu.devices[0] == gpu_devices[0]); + assert(multi_gpu.devices[1] == gpu_devices[1]); + assert(multi_gpu.devices[2] == nullptr); + assert(multi_gpu.fit_params_target[0] == 128 * 1024 * 1024ULL); + assert(multi_gpu.fit_params_target[1] == 256 * 1024 * 1024ULL); + assert(multi_gpu.tensor_split[0] == 3.0f); + assert(multi_gpu.tensor_split[1] == 1.0f); + } else { + printf("test-arg-parser: skip asymmetric multi-GPU input test (fewer than two non-CPU devices)\n\n"); + } + } + argv = {"binary_name", "-lm", "none"}; assert(true == common_params_parse(argv.size(), list_str_to_char(argv).data(), params, LLAMA_EXAMPLE_COMMON)); assert(params.load_mode == LLAMA_LOAD_MODE_NONE); diff --git a/tests/test-recurrent-state-rollback.cpp b/tests/test-recurrent-state-rollback.cpp index c6f599e5..ae972958 100644 --- a/tests/test-recurrent-state-rollback.cpp +++ b/tests/test-recurrent-state-rollback.cpp @@ -8,15 +8,112 @@ #include #include -static llama_context * make_ctx(const common_params & params, llama_model * model) { +static llama_context * make_ctx( + const common_params & params, + llama_model * model, + uint32_t n_rs_seq = 8, + uint32_t n_seq_max = 1) { auto cparams = common_context_params_to_llama(params); - cparams.n_seq_max = 1; - cparams.n_rs_seq = 8; + cparams.n_seq_max = n_seq_max; + cparams.n_rs_seq = n_rs_seq; cparams.n_batch = std::max(cparams.n_batch, (uint32_t) (cparams.n_rs_seq + 1)); cparams.n_ubatch = std::max(cparams.n_ubatch, (uint32_t) (cparams.n_rs_seq + 1)); return llama_init_from_model(model, cparams); } +static bool test_fragmented_on_device_fallback( + const common_params & params, + llama_model * model, + const std::vector & tokens) { + llama_context * ctx = make_ctx(params, model, 1, 3); + if (ctx == nullptr) { + fprintf(stderr, "%s : failed to initialize fragmented context\n", __func__); + return false; + } + + llama_batch batch = llama_batch_init(3, 0, 3); + for (llama_seq_id seq_id = 0; seq_id < 3; ++seq_id) { + common_batch_add(batch, tokens[seq_id % tokens.size()], 0, { seq_id }, seq_id == 2); + } + bool ok = llama_decode(ctx, batch) == 0; + llama_batch_free(batch); + + if (ok) { + ok = llama_memory_seq_rm(llama_get_memory(ctx), 1, -1, -1); + } + + common_prompt_checkpoint ckpt; + constexpr llama_state_seq_flags flags = + LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY | LLAMA_STATE_SEQ_FLAGS_ON_DEVICE; + if (ok) { + ckpt.update_tgt(ctx, -1, flags); + if (ckpt.data_tgt_on_device) { + // Some recurrent layouts compact the surviving sequences into one + // device range. That is valid and does not require host fallback. + fprintf(stderr, "%s : checkpoint remained device-contiguous\n", __func__); + } else { + fprintf(stderr, "%s : fragmented checkpoint fell back to host storage\n", __func__); + } + } + if (ok) { + // load_tgt() must use the actual storage mode recorded when the + // checkpoint was saved, whether the layout stayed contiguous or fell + // back to host storage. + ckpt.load_tgt(ctx, -1, flags); + } + + llama_free(ctx); + return ok; +} + +static bool test_preallocated_device_checkpoints( + const common_params & params, + llama_model * model, + const std::vector & tokens) { + constexpr int32_t n_seq = 2; + llama_context * ctx = make_ctx(params, model, 1, n_seq); + if (ctx == nullptr) { + fprintf(stderr, "%s : failed to initialize multi-sequence context\n", __func__); + return false; + } + + bool ok = llama_state_seq_reserve_device_buffers(ctx); + if (!ok) { + // Architectures without recurrent tensor rows have no device layout to + // reserve. The server handles this by using host checkpoints. + fprintf(stderr, "%s : no device checkpoint layout; using host checkpoints\n", __func__); + llama_free(ctx); + return true; + } + + llama_batch batch = llama_batch_init(n_seq, 0, n_seq); + for (llama_seq_id seq_id = 0; seq_id < n_seq; ++seq_id) { + common_batch_add(batch, tokens[seq_id % tokens.size()], 0, { seq_id }, true); + } + if (ok) { + ok = llama_decode(ctx, batch) == 0; + } + llama_batch_free(batch); + + constexpr llama_state_seq_flags flags = + LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY | LLAMA_STATE_SEQ_FLAGS_ON_DEVICE; + std::vector checkpoints(n_seq); + for (llama_seq_id seq_id = 0; ok && seq_id < n_seq; ++seq_id) { + checkpoints[seq_id].update_tgt(ctx, seq_id, flags); + if (!checkpoints[seq_id].data_tgt_on_device) { + fprintf(stderr, "%s : sequence %d did not reuse its preallocated device checkpoint\n", + __func__, seq_id); + ok = false; + } + } + for (llama_seq_id seq_id = 0; ok && seq_id < n_seq; ++seq_id) { + checkpoints[seq_id].load_tgt(ctx, seq_id, flags); + } + + llama_free(ctx); + return ok; +} + static bool decode_tokens(llama_context * ctx, const std::vector & tokens, uint32_t count) { llama_batch batch = llama_batch_init(count, 0, 1); for (uint32_t pos = 0; pos < count; ++pos) { @@ -207,6 +304,156 @@ static bool test_multi_seq_split_replay(const common_params & params, llama_mode return true; } +static bool decode_range( + llama_context * ctx, + const std::vector & tokens, + uint32_t first, + uint32_t count) { + if (count == 0) { + return true; + } + + llama_batch batch = llama_batch_init(count, 0, 1); + for (uint32_t i = 0; i < count; ++i) { + common_batch_add(batch, tokens[first + i], first + i, { 0 }, i + 1 == count); + } + const bool ok = llama_decode(ctx, batch) == 0; + llama_batch_free(batch); + return ok; +} + +static bool compare_logits( + llama_context * ctx_full, + llama_context * ctx_depth, + int n_vocab, + uint32_t n_accepted, + const char * stage, + float eps) { + const float * logits_full = llama_get_logits(ctx_full); + const float * logits_depth = llama_get_logits(ctx_depth); + if (logits_full == nullptr || logits_depth == nullptr) { + fprintf(stderr, "%s : missing %s logits after accepting %u draft tokens\n", __func__, stage, n_accepted); + return false; + } + + int argmax_full = 0; + int argmax_depth = 0; + float max_diff = 0.0f; + for (int token = 0; token < n_vocab; ++token) { + if (logits_full[token] > logits_full[argmax_full]) { + argmax_full = token; + } + if (logits_depth[token] > logits_depth[argmax_depth]) { + argmax_depth = token; + } + + const float diff = std::fabs(logits_full[token] - logits_depth[token]); + max_diff = std::max(max_diff, diff); + if (eps >= 0.0f && diff > eps) { + fprintf(stderr, "%s : %s logits mismatch after accepting %u draft tokens, token %d (%g != %g)\n", + __func__, stage, n_accepted, token, (double) logits_full[token], (double) logits_depth[token]); + return false; + } + } + if (eps < 0.0f && argmax_full != argmax_depth) { + fprintf(stderr, "%s : %s greedy token mismatch after accepting %u draft tokens (%d != %d, max logit diff %g)\n", + __func__, stage, n_accepted, argmax_full, argmax_depth, (double) max_diff); + return false; + } + if (eps < 0.0f) { + fprintf(stderr, "%s : %s greedy token %d preserved after accepting %u draft tokens (max logit diff %g)\n", + __func__, stage, argmax_full, n_accepted, (double) max_diff); + } + return true; +} + +static bool test_recompute_fallback( + const common_params & params, + llama_model * model, + const std::vector & input_tokens, + int n_vocab) { + constexpr uint32_t n_prefix = 4; + constexpr uint32_t n_draft = 5; + constexpr uint32_t n_verify = n_draft + 1; + + std::vector tokens = input_tokens; + tokens.resize(n_prefix + n_verify + 2, tokens.back()); + + for (uint32_t n_accepted = 0; n_accepted < n_draft; ++n_accepted) { + llama_context * ctx_full = make_ctx(params, model, n_draft); + llama_context * ctx_depth = make_ctx(params, model, 1); + if (ctx_full == nullptr || ctx_depth == nullptr) { + fprintf(stderr, "%s : failed to initialize fallback contexts\n", __func__); + llama_free(ctx_full); + llama_free(ctx_depth); + return false; + } + + bool ok = decode_range(ctx_full, tokens, 0, n_prefix) && + decode_range(ctx_depth, tokens, 0, n_prefix); + if (ok) { + ok = compare_logits(ctx_full, ctx_depth, n_vocab, n_accepted, "prefix", 1e-5f); + } + + constexpr llama_state_seq_flags partial_flags = + LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY | LLAMA_STATE_SEQ_FLAGS_ON_DEVICE; + common_prompt_checkpoint ckpt; + if (ok) { + ckpt.update_tgt(ctx_depth, 0, partial_flags); + if (!ckpt.data_tgt_on_device) { + fprintf(stderr, "%s : contiguous recurrent checkpoint unexpectedly fell back to host storage\n", __func__); + ok = false; + } + } + if (ok) { + ok = decode_range(ctx_full, tokens, n_prefix, n_verify) && + decode_range(ctx_depth, tokens, n_prefix, n_verify); + } + + const llama_pos rollback_pos = n_prefix + 1 + n_accepted; + const uint32_t n_rollback = n_draft - n_accepted; + if (ok) { + ok = llama_memory_seq_rm(llama_get_memory(ctx_full), 0, rollback_pos, -1); + } + if (ok && n_rollback <= 1) { + ok = llama_memory_seq_rm(llama_get_memory(ctx_depth), 0, rollback_pos, -1); + } else if (ok) { + ckpt.load_tgt(ctx_depth, 0, partial_flags); + ok = llama_memory_seq_rm(llama_get_memory(ctx_depth), 0, n_prefix, -1); + } + + const llama_pos replacement_pos = rollback_pos; + if (ok && n_rollback <= 1) { + ok = decode_one(ctx_full, tokens[n_prefix + n_verify], replacement_pos) && + decode_one(ctx_depth, tokens[n_prefix + n_verify], replacement_pos); + } else if (ok) { + std::vector replay_tokens = tokens; + replay_tokens[replacement_pos] = tokens[n_prefix + n_verify]; + ok = decode_one(ctx_full, tokens[n_prefix + n_verify], replacement_pos) && + decode_range(ctx_depth, replay_tokens, n_prefix, 2 + n_accepted); + } + if (ok) { + ok = compare_logits(ctx_full, ctx_depth, n_vocab, n_accepted, "replacement", -1.0f); + } + if (ok) { + ok = decode_one(ctx_full, tokens[n_prefix + n_verify + 1], replacement_pos + 1) && + decode_one(ctx_depth, tokens[n_prefix + n_verify + 1], replacement_pos + 1) && + compare_logits(ctx_full, ctx_depth, n_vocab, n_accepted, "continuation", -1.0f); + } + + llama_free(ctx_full); + llama_free(ctx_depth); + + if (!ok) { + fprintf(stderr, "%s : fallback validation failed after accepting %u draft tokens\n", __func__, n_accepted); + return false; + } + } + + fprintf(stderr, "%s : depth-1 recomputation preserves depth-5 greedy decisions at every rejection position\n", __func__); + return true; +} + int main(int argc, char ** argv) { std::setlocale(LC_NUMERIC, "C"); @@ -237,6 +484,26 @@ int main(int argc, char ** argv) { const llama_vocab * vocab = llama_model_get_vocab(model); const int n_vocab = llama_vocab_n_tokens(vocab); + std::vector tokens; + if (llama_vocab_type(vocab) == LLAMA_VOCAB_TYPE_NONE) { + tokens = { 1, 2, 3, 4, 5, 6, 7, 8, 9 }; + } else { + tokens = common_tokenize(vocab, "The quick brown fox jumps over the lazy dog", true); + } + if (tokens.empty()) { + fprintf(stderr, "%s : not enough prompt tokens\n", __func__); + return 1; + } + if (!test_fragmented_on_device_fallback(params, model, tokens)) { + return 1; + } + if (!test_preallocated_device_checkpoints(params, model, tokens)) { + return 1; + } + if (!test_recompute_fallback(params, model, tokens, n_vocab)) { + return 1; + } + llama_context * ctx_src = make_ctx(params, model); llama_context * ctx_dst = make_ctx(params, model); if (ctx_src == nullptr || ctx_dst == nullptr) { @@ -251,12 +518,6 @@ int main(int argc, char ** argv) { return 0; } - std::vector tokens; - if (llama_vocab_type(vocab) == LLAMA_VOCAB_TYPE_NONE) { - tokens = { 1, 2, 3, 4, 5, 6, 7, 8, 9 }; - } else { - tokens = common_tokenize(ctx_src, "The quick brown fox jumps over the lazy dog", true); - } const uint32_t n_rs_seq = llama_n_rs_seq(ctx_src); constexpr uint32_t n_rollback = 3; if (n_rs_seq < n_rollback) { @@ -265,10 +526,6 @@ int main(int argc, char ** argv) { llama_free(ctx_dst); return 0; } - if (tokens.empty()) { - fprintf(stderr, "%s : not enough prompt tokens\n", __func__); - return 1; - } tokens.resize(n_rs_seq + 1, tokens.back()); const uint32_t n_tokens = tokens.size(); @@ -333,6 +590,10 @@ int main(int argc, char ** argv) { constexpr llama_state_seq_flags partial_flags = LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY; common_prompt_checkpoint ckpt_partial; ckpt_partial.update_tgt(ctx_src, 0, partial_flags); + if (ckpt_partial.data_tgt_on_device) { + fprintf(stderr, "%s : host recurrent checkpoint recorded device storage\n", __func__); + return 1; + } ckpt_partial.load_tgt(ctx_dst, 0, partial_flags); if (!replay_and_compare("partial")) { diff --git a/tools/server/README.md b/tools/server/README.md index 93736c3e..b6727e34 100644 --- a/tools/server/README.md +++ b/tools/server/README.md @@ -255,6 +255,7 @@ For the full list of features, please refer to [server's changelog](https://gith | `--spec-draft-cpu-moe, -cmoed, --cpu-moe-draft` | keep all Mixture of Experts (MoE) weights in the CPU for the draft model
(env: LLAMA_ARG_SPEC_DRAFT_CPU_MOE) | | `--spec-draft-n-cpu-moe, --spec-draft-ncmoe, -ncmoed, --n-cpu-moe-draft N` | keep the Mixture of Experts (MoE) weights of the first N layers in the CPU for the draft model
(env: LLAMA_ARG_SPEC_DRAFT_N_CPU_MOE) | | `--spec-draft-n-max N` | number of tokens to draft for speculative decoding (default: 3)
(env: LLAMA_ARG_SPEC_DRAFT_N_MAX) | +| `--spec-mtp-cr-depth N` | MTP Compact Rollback depth; lower values save memory but replay accepted tokens after deep rejection (default: `--spec-draft-n-max`)
(env: LLAMA_ARG_SPEC_MTP_CR_DEPTH) | | `--spec-draft-n-min N` | minimum number of draft tokens to use for speculative decoding (default: 0)
(env: LLAMA_ARG_SPEC_DRAFT_N_MIN) | | `--spec-draft-p-split, --draft-p-split P` | speculative decoding split probability (default: 0.10)
(env: LLAMA_ARG_SPEC_DRAFT_P_SPLIT) | | `--spec-draft-p-min, --draft-p-min P` | minimum speculative decoding probability (greedy) (default: 0.00)
(env: LLAMA_ARG_SPEC_DRAFT_P_MIN) | diff --git a/tools/server/server-common.cpp b/tools/server/server-common.cpp index 4f5b8202..7d2593b7 100644 --- a/tools/server/server-common.cpp +++ b/tools/server/server-common.cpp @@ -83,6 +83,10 @@ json server_slot_stats::to_json() const { base["draft_n"] = n_draft_tokens; base["draft_n_accepted"] = n_draft_accepted; } + if (n_draft_replay_count > 0) { + base["draft_replay_count"] = n_draft_replay_count; + base["draft_replay_n"] = n_draft_replay_tokens; + } return base; } diff --git a/tools/server/server-common.h b/tools/server/server-common.h index f8ea82ef..1d5f28ec 100644 --- a/tools/server/server-common.h +++ b/tools/server/server-common.h @@ -352,6 +352,8 @@ struct server_slot_stats { uint64_t n_draft_tokens = 0; uint64_t n_draft_accepted = 0; uint64_t n_draft_verif_steps = 0; + uint64_t n_draft_replay_count = 0; + uint64_t n_draft_replay_tokens = 0; // these are absolute timestamps (in us) // note: must be signed - they are subtracted before the later ones are set diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index a9edbd7b..e6947c09 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -634,6 +634,12 @@ struct server_slot { " acc per pos = (%s)\n", acceptance_rates_per_pos.c_str()); } + if (stats.n_draft_replay_count > 0) { + SLT_INF(*this, + " MTP replays = %10" PRIu64 " events / %5" PRIu64 " tokens\n", + stats.n_draft_replay_count, stats.n_draft_replay_tokens); + } + common_speculative_print_stats(spec); } @@ -844,6 +850,7 @@ private: common_context_seq_rm_type ctx_tgt_seq_rm_type = COMMON_CONTEXT_SEQ_RM_TYPE_NO; common_context_seq_rm_type ctx_dft_seq_rm_type = COMMON_CONTEXT_SEQ_RM_TYPE_NO; + bool spec_mtp_device_checkpoint = false; common_speculative_ptr spec; @@ -1040,6 +1047,90 @@ private: } } + // The on-device recurrent speculative checkpoint is allocated lazily, so reserve + // the additional recurrent-state group before the joint target / MTP fit. + if (params_base.fit_params && spec_mtp) { + const uint32_t n_rs_seq = params_base.speculative.need_n_rs_seq(); + const int32_t draft_n_max = params_base.speculative.draft.n_max; + if (draft_n_max < 0) { + SRV_ERR("%s", "[spec] invalid negative MTP draft.n_max while fitting\n"); + return false; + } + if (n_rs_seq < static_cast(draft_n_max)) { + try { + common_params params_rs = params_base; + auto mparams_rs = common_model_params_to_llama(params_rs); + auto cparams_rs = common_context_params_to_llama(params_rs); + + std::vector devs_rs; + uint32_t hp_ngl_rs = 0; + uint32_t hp_nct_rs = 0; + uint32_t hp_nex_rs = 0; + cparams_rs.n_rs_seq = n_rs_seq; + const auto dmd_rs = common_get_device_memory_data( + params_base.model.path.c_str(), &mparams_rs, &cparams_rs, + devs_rs, hp_ngl_rs, hp_nct_rs, hp_nex_rs, GGML_LOG_LEVEL_ERROR); + + std::vector devs_rs_next; + uint32_t hp_ngl_rs_next = 0; + uint32_t hp_nct_rs_next = 0; + uint32_t hp_nex_rs_next = 0; + cparams_rs.n_rs_seq = n_rs_seq + 1; + const auto dmd_rs_next = common_get_device_memory_data( + params_base.model.path.c_str(), &mparams_rs, &cparams_rs, + devs_rs_next, hp_ngl_rs_next, hp_nct_rs_next, hp_nex_rs_next, GGML_LOG_LEVEL_ERROR); + + if (devs_rs.size() != devs_rs_next.size() || dmd_rs.size() < devs_rs.size() || + dmd_rs_next.size() < devs_rs_next.size()) { + throw std::runtime_error("target context device count changed during recurrent-state measurement"); + } + + // common_fit_params() consumes fit_params_target in the measured model-device + // order. The configured device list can contain a nullptr sentinel, CPU/ACCEL + // devices, or be reduced by split-mode none, so it is not a safe index map. + if (params_base.fit_params_target.size() < devs_rs.size()) { + throw std::runtime_error("fit_params_target has no entry for every target device"); + } + + size_t total = 0; + for (size_t j = 0; j < devs_rs.size(); ++j) { + const auto next_dev = std::find(devs_rs_next.begin(), devs_rs_next.end(), devs_rs[j]); + if (next_dev == devs_rs_next.end()) { + throw std::runtime_error("target context device mapping changed during recurrent-state measurement"); + } + const size_t next_index = next_dev - devs_rs_next.begin(); + if (next_index != j) { + throw std::runtime_error("target context device order changed during recurrent-state measurement"); + } + if (dmd_rs_next[j].context < dmd_rs[j].context) { + throw std::runtime_error("recurrent-state context measurement decreased"); + } + const size_t delta = dmd_rs_next[j].context - dmd_rs[j].context; + const size_t checkpoint = dmd_rs[j].checkpoint; + if (checkpoint == 0 && delta != 0) { + throw std::runtime_error("recurrent checkpoint sizing produced no device allocation"); + } + params_base.fit_params_target[j] += checkpoint; + total += checkpoint; + SRV_INF("[spec] recurrent checkpoint fit reservation: device %s, %.2f MiB " + "(exact backend allocation; recurrent-plane delta %.2f MiB, context %.2f -> %.2f MiB)\n", + ggml_backend_dev_name(devs_rs[j]), checkpoint / (1024.0 * 1024.0), + delta / (1024.0 * 1024.0), + dmd_rs[j].context / (1024.0 * 1024.0), + dmd_rs_next[j].context / (1024.0 * 1024.0)); + } + if (total == 0) { + throw std::runtime_error("recurrent-state measurement produced no device reservation"); + } + SRV_INF("[spec] recurrent checkpoint fit reservation: %.2f MiB total\n", + total / (1024.0 * 1024.0)); + } catch (const std::exception & e) { + SRV_ERR("[spec] failed to reserve recurrent-state memory before fitting: %s\n", e.what()); + return false; + } + } + } + // note: the draft / MTP context is fitted together with the target model, see common_fit_extra_model // attach a progress callback @@ -1202,6 +1293,16 @@ private: model_dft = nullptr; } + spec_mtp_device_checkpoint = false; + if (spec && ctx_tgt_seq_rm_type == COMMON_CONTEXT_SEQ_RM_TYPE_RS) { + spec_mtp_device_checkpoint = llama_state_seq_reserve_device_buffers(ctx_tgt); + if (spec_mtp_device_checkpoint) { + SRV_INF("%s", "[spec] reserved device recurrent checkpoints before evaluation\n"); + } else { + SRV_WRN("%s", "[spec] device recurrent checkpoint reservation failed; using host checkpoints\n"); + } + } + for (int i = 0; i < params_base.n_parallel; i++) { server_slot & slot = slots[i]; @@ -2960,9 +3061,23 @@ private: (ctx_dft_seq_rm_type == COMMON_CONTEXT_SEQ_RM_TYPE_RS && draft.size() > llama_n_rs_seq(ctx_dft)); if (use_ckpt_tgt) { + llama_state_seq_flags ckpt_flags = LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY; + if (ctx_tgt_seq_rm_type == COMMON_CONTEXT_SEQ_RM_TYPE_RS && spec_mtp_device_checkpoint) { + // Phase 1 proof of concept: avoid copying the recurrent fallback checkpoint + // through host memory on every speculative round. The device buffer is + // allocated during load_model(), before evaluation can consume the + // fitted headroom reserved for the recurrent-state checkpoint. + ckpt_flags |= LLAMA_STATE_SEQ_FLAGS_ON_DEVICE; + } + //const int64_t t_start = ggml_time_us(); - ckpt.update_tgt(ctx_tgt, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); + ckpt.update_tgt(ctx_tgt, slot.id, ckpt_flags); + + if ((ckpt_flags & LLAMA_STATE_SEQ_FLAGS_ON_DEVICE) && !ckpt.data_tgt_on_device) { + spec_mtp_device_checkpoint = false; + SRV_WRN("%s", "[spec] device recurrent checkpoint save failed; using host checkpoints for the rest of this server run\n"); + } //const int64_t t_total = ggml_time_us() - t_start; //printf("checkpoint total: %f ms\n", t_total / 1000.0); @@ -3806,6 +3921,10 @@ private: ctx_tgt_seq_rm_type == COMMON_CONTEXT_SEQ_RM_TYPE_FULL || (ctx_tgt_seq_rm_type == COMMON_CONTEXT_SEQ_RM_TYPE_RS && n_rollback > llama_n_rs_seq(ctx_tgt)); + const bool use_rs_replay = + ctx_tgt_seq_rm_type == COMMON_CONTEXT_SEQ_RM_TYPE_RS && + n_rollback > llama_n_rs_seq(ctx_tgt); + // check for partial draft acceptance if (n_rollback > 0) { if (use_ckpt_tgt) { @@ -3817,11 +3936,20 @@ private: slot.spec_is_replay = true; slot.spec_draft = std::move(accepted); + if (use_rs_replay) { + slot.stats.n_draft_replay_count += 1; + slot.stats.n_draft_replay_tokens += slot.spec_draft.size() + 1; + } + const auto & ckpt = slot.spec_ckpt; SLT_DBG(slot, "restoring speculative checkpoint (pos_min = %d, pos_max = %d, size = %zu)\n", ckpt.pos_min, ckpt.pos_max, ckpt.size()); - ckpt.load_tgt(slot.ctx_tgt, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); + llama_state_seq_flags ckpt_flags = LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY; + if (use_rs_replay) { + ckpt_flags |= LLAMA_STATE_SEQ_FLAGS_ON_DEVICE; + } + ckpt.load_tgt(slot.ctx_tgt, slot.id, ckpt_flags); if (slot.ctx_dft) { ckpt.load_dft(slot.ctx_dft, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); diff --git a/tools/server/tests/utils.py b/tools/server/tests/utils.py index a0d2dfa3..b3cac346 100644 --- a/tools/server/tests/utils.py +++ b/tools/server/tests/utils.py @@ -99,6 +99,7 @@ class ServerProcess: spec_type: str | None = None spec_draft_n_min: int | None = None spec_draft_n_max: int | None = None + spec_mtp_cr_depth: int | None = None no_ui: bool | None = None jinja: bool | None = None reasoning_format: Literal['deepseek', 'none', 'nothink'] | None = None @@ -245,6 +246,8 @@ class ServerProcess: server_args.extend(["--spec-draft-n-max", self.spec_draft_n_max]) if self.spec_draft_n_min: server_args.extend(["--spec-draft-n-min", self.spec_draft_n_min]) + if self.spec_mtp_cr_depth: + server_args.extend(["--spec-mtp-cr-depth", self.spec_mtp_cr_depth]) if self.no_ui: server_args.append("--no-ui") if self.no_models_autoload: