diff --git a/src/llama-kv-cache.cpp b/src/llama-kv-cache.cpp index 1f773a5a41..7302714c49 100644 --- a/src/llama-kv-cache.cpp +++ b/src/llama-kv-cache.cpp @@ -1130,7 +1130,7 @@ void llama_kv_cache::apply_ubatch(const slot_info & sinfo, const llama_ubatch & cells.pos_set(idx, ubatch.pos[i]); - if (ubatch.is_pos_2d() || ubatch.token) { + if (ubatch.is_pos_2d() || ubatch.token || hparams.ple_n_heads > 0) { llama_kv_cell_ext ext; if (ubatch.is_pos_2d()) { @@ -1140,6 +1140,12 @@ void llama_kv_cache::apply_ubatch(const slot_info & sinfo, const llama_ubatch & if (ubatch.token) { ext.tok = ubatch.token[i]; + } else if (hparams.ple_n_heads > 0) { + // an embd batch has already consumed the image placeholder; store what the + // reference hashes in input_ids at these positions (EOS for old GGUF files) + ext.tok = hparams.ple_image_token_id != 0 + ? (llama_token) hparams.ple_image_token_id + : (llama_token) hparams.ple_eos_token_id; } cells.ext_set(idx, ext); @@ -1815,7 +1821,8 @@ void llama_kv_cache::set_input_v_rot(ggml_tensor * dst) const { } bool llama_kv_cache::has_cell_ext() const { - return hparams.n_pos_per_embd() > 1; + // M-RoPE needs the 2D position, the PLE n-gram hash needs the token id + return hparams.n_pos_per_embd() > 1 || hparams.ple_n_heads > 0; } void llama_kv_cache::get_prev_tokens(const llama_ubatch & ubatch, uint32_t n, std::vector & res) const { diff --git a/src/llama-memory-hybrid-idx.cpp b/src/llama-memory-hybrid-idx.cpp index e7f6e565dc..eb0dd6a51f 100644 --- a/src/llama-memory-hybrid-idx.cpp +++ b/src/llama-memory-hybrid-idx.cpp @@ -139,9 +139,6 @@ void llama_memory_hybrid_idx::clear(bool data) { if (mem_idx) { mem_idx->clear(data); } - - // [TAG_PLE_HISTORY] every sequence is gone, so no window is trusted any more - ple_hist.clear(); } bool llama_memory_hybrid_idx::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) { @@ -155,8 +152,6 @@ bool llama_memory_hybrid_idx::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_po mem_idx->seq_rm(seq_id, p0, p1); } - ple_hist_rm(seq_id, p0, p1); - return get_mem_attn()->seq_rm(seq_id, p0, p1); } @@ -166,8 +161,6 @@ void llama_memory_hybrid_idx::seq_cp(llama_seq_id seq_id_src, llama_seq_id seq_i if (mem_idx) { mem_idx->seq_cp(seq_id_src, seq_id_dst, p0, p1); } - - ple_hist_cp(seq_id_src, seq_id_dst, p0, p1); } void llama_memory_hybrid_idx::seq_keep(llama_seq_id seq_id) { @@ -176,8 +169,6 @@ void llama_memory_hybrid_idx::seq_keep(llama_seq_id seq_id) { if (mem_idx) { mem_idx->seq_keep(seq_id); } - - ple_hist_keep(seq_id); } void llama_memory_hybrid_idx::seq_add(llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift) { @@ -186,8 +177,6 @@ void llama_memory_hybrid_idx::seq_add(llama_seq_id seq_id, llama_pos p0, llama_p if (mem_idx) { mem_idx->seq_add(seq_id, p0, p1, shift); } - - ple_hist_add(seq_id, p0, p1, shift); } void llama_memory_hybrid_idx::seq_div(llama_seq_id seq_id, llama_pos p0, llama_pos p1, int d) { @@ -196,8 +185,6 @@ void llama_memory_hybrid_idx::seq_div(llama_seq_id seq_id, llama_pos p0, llama_p if (mem_idx) { mem_idx->seq_div(seq_id, p0, p1, d); } - - ple_hist_div(seq_id, p0, p1); } std::map llama_memory_hybrid_idx::memory_breakdown() const { @@ -212,251 +199,6 @@ std::map llama_memory_hybrid_idx::memory_bre return mb; } -// -// [TAG_PLE_HISTORY] per-sequence PLE n-gram history -// -// The window is only usable while it is contiguous with the position the sequence decodes next, -// so each operation below either rewrites it exactly or invalidates it with next_pos = -1. -// An invalid window makes set_input pad with EOS, which is what a fresh sequence also gets. -// - -llama_memory_hybrid_idx::ple_history & llama_memory_hybrid_idx::ple_hist_get(llama_seq_id seq_id) const { - return ple_hist[seq_id]; -} - -// first position still remembered by h -static llama_pos ple_hist_beg(const llama_memory_hybrid_idx::ple_history & h) { - return h.next_pos - (llama_pos) h.toks.size(); -} - -static void ple_hist_invalidate(llama_memory_hybrid_idx::ple_history & h) { - h.next_pos = -1; - h.toks.clear(); -} - -// drop everything at position >= p, so the sequence now ends just before p -static void ple_hist_truncate(llama_memory_hybrid_idx::ple_history & h, llama_pos p) { - const llama_pos beg = ple_hist_beg(h); - - if (p <= beg) { - h.toks.clear(); - } else if (p < h.next_pos) { - h.toks.resize((size_t) (p - beg)); - } - - h.next_pos = p; -} - -void llama_memory_hybrid_idx::ple_hist_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) { - if (seq_id < 0) { - // the recursive call erases seq_id's entry when the whole sequence is removed, so - // advance past it first: erase invalidates only the iterator to the erased element - for (auto it = ple_hist.begin(); it != ple_hist.end(); ) { - const llama_seq_id id = it->first; - ++it; - ple_hist_rm(id, p0, p1); - } - return; - } - - auto it = ple_hist.find(seq_id); - if (it == ple_hist.end() || it->second.next_pos < 0) { - return; - } - auto & h = it->second; - - if (p0 <= 0 && p1 < 0) { - // the whole sequence is gone - ple_hist.erase(it); - return; - } - - if (p1 < 0) { - // a rewind: the sequence ends at p0 and the remaining prefix is still contiguous - if (p0 < h.next_pos) { - ple_hist_truncate(h, p0); - } - return; - } - - // a hole in the middle: seq_rm does not renumber what follows, so an overlapping - // window is no longer a run of consecutive positions - if (p1 > ple_hist_beg(h) && p0 < h.next_pos) { - ple_hist_invalidate(h); - } -} - -void llama_memory_hybrid_idx::ple_hist_cp(llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) { - if (seq_id_src == seq_id_dst) { - return; - } - - // whatever the destination had is replaced by the copied range, exactly as its cells are - ple_hist.erase(seq_id_dst); - - auto it = ple_hist.find(seq_id_src); - if (it == ple_hist.end() || it->second.next_pos < 0) { - return; - } - - ple_history h = it->second; - - if (p1 >= 0 && p1 < h.next_pos) { - ple_hist_truncate(h, p1); - } - - // positions below p0 were not copied, so for the destination they are before the - // sequence start, which the hash already reads as EOS - const llama_pos lo = p0 < 0 ? 0 : p0; - if (lo > ple_hist_beg(h)) { - const llama_pos drop = std::min(lo - ple_hist_beg(h), (llama_pos) h.toks.size()); - h.toks.erase(h.toks.begin(), h.toks.begin() + (size_t) drop); - } - - ple_hist[seq_id_dst] = std::move(h); -} - -void llama_memory_hybrid_idx::ple_hist_keep(llama_seq_id seq_id) { - for (auto it = ple_hist.begin(); it != ple_hist.end(); ) { - it = it->first == seq_id ? std::next(it) : ple_hist.erase(it); - } -} - -void llama_memory_hybrid_idx::ple_hist_add(llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift) { - if (seq_id < 0) { - for (auto & it : ple_hist) { - ple_hist_add(it.first, p0, p1, shift); - } - return; - } - - auto it = ple_hist.find(seq_id); - if (it == ple_hist.end() || it->second.next_pos < 0) { - return; - } - auto & h = it->second; - - const llama_pos beg = ple_hist_beg(h); - const llama_pos lo = p0 < 0 ? 0 : p0; - - if (p1 >= 0 && p1 <= beg) { - // entirely below the window: the tokens we remember keep their positions - return; - } - if (lo >= h.next_pos) { - // entirely above the window: nothing we remember moves - return; - } - if (lo <= beg && (p1 < 0 || p1 >= h.next_pos)) { - // the context-shift case: the whole window moves as one and stays consecutive - if (beg + shift < 0) { - ple_hist_invalidate(h); - } else { - h.next_pos += shift; - } - return; - } - - // the shift cuts through the window and breaks its contiguity - ple_hist_invalidate(h); -} - -void llama_memory_hybrid_idx::ple_hist_div(llama_seq_id seq_id, llama_pos p0, llama_pos p1) { - if (seq_id < 0) { - for (auto & it : ple_hist) { - ple_hist_div(it.first, p0, p1); - } - return; - } - - auto it = ple_hist.find(seq_id); - if (it == ple_hist.end() || it->second.next_pos < 0) { - return; - } - auto & h = it->second; - - const llama_pos lo = p0 < 0 ? 0 : p0; - - // division makes the positions non-consecutive, so any overlap ends the window - if ((p1 < 0 || p1 > ple_hist_beg(h)) && lo < h.next_pos) { - ple_hist_invalidate(h); - } -} - -// Serialised as a self-delimiting list so that seq_id == -1 (whole context) and a single -// sequence share one format: -// u32 n_entries -// n_entries * { i32 seq_id, i32 next_pos, u32 n_toks, i32 toks[n_toks] } -// n_toks is at most ple_ngram_size - 1, so an entry is a handful of bytes. -void llama_memory_hybrid_idx::ple_hist_state_write(llama_io_write_i & io, llama_seq_id seq_id) const { - uint32_t n_entries = 0; - for (const auto & it : ple_hist) { - if ((seq_id < 0 || it.first == seq_id) && it.second.next_pos >= 0) { - ++n_entries; - } - } - - io.write(&n_entries, sizeof(n_entries)); - - for (const auto & it : ple_hist) { - if ((seq_id >= 0 && it.first != seq_id) || it.second.next_pos < 0) { - continue; - } - - const int32_t id = it.first; - const int32_t next_pos = it.second.next_pos; - const uint32_t n_toks = (uint32_t) it.second.toks.size(); - - io.write(&id, sizeof(id)); - io.write(&next_pos, sizeof(next_pos)); - io.write(&n_toks, sizeof(n_toks)); - if (n_toks > 0) { - io.write(it.second.toks.data(), n_toks*sizeof(llama_token)); - } - } -} - -void llama_memory_hybrid_idx::ple_hist_state_read(llama_io_read_i & io, llama_seq_id seq_id) { - uint32_t n_entries = 0; - io.read(&n_entries, sizeof(n_entries)); - - // a single-sequence restore replaces one window, a whole-context one replaces them all, - // as the caches around it do - if (seq_id >= 0) { - ple_hist.erase(seq_id); - } else { - ple_hist.clear(); - } - - for (uint32_t i = 0; i < n_entries; ++i) { - int32_t id = 0; - int32_t next_pos = 0; - uint32_t n_toks = 0; - - io.read(&id, sizeof(id)); - io.read(&next_pos, sizeof(next_pos)); - io.read(&n_toks, sizeof(n_toks)); - - // the window is never longer than ple_ngram_size - 1; anything else is a corrupt or - // mismatched blob, and reading it would size an allocation from the file - if (n_toks > LLAMA_MAX_PLE_NGRAM - 1) { - throw std::runtime_error("qwen4exp PLE history: implausible token count in state blob"); - } - - std::vector toks(n_toks); - if (n_toks > 0) { - io.read(toks.data(), n_toks*sizeof(llama_token)); - } - - // a single-sequence restore can target a different seq_id, so the destination wins - const llama_seq_id dst = seq_id >= 0 ? seq_id : (llama_seq_id) id; - - auto & h = ple_hist[dst]; - h.next_pos = next_pos; - h.toks = std::move(toks); - } -} - void llama_memory_hybrid_idx::state_write(llama_io_write_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) const { llama_memory_hybrid::state_write(io, seq_id, flags); @@ -469,10 +211,6 @@ void llama_memory_hybrid_idx::state_write(llama_io_write_i & io, llama_seq_id se } } - // [TAG_PLE_HISTORY] last again, so this section is also a pure suffix. - // It is not under the PARTIAL_ONLY gate: the window is recurrent state, the input the PLE - // conv state comes from, and the recurrent cache is written for partial checkpoints too. - ple_hist_state_write(io, seq_id); } void llama_memory_hybrid_idx::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) { @@ -487,8 +225,6 @@ void llama_memory_hybrid_idx::state_read(llama_io_read_i & io, llama_seq_id seq_ } } - // [TAG_PLE_HISTORY] must mirror the write order above - ple_hist_state_read(io, seq_id); } llama_kv_cache * llama_memory_hybrid_idx::get_mem_idx() const { @@ -567,12 +303,6 @@ uint32_t llama_memory_hybrid_idx_context::get_n_stream() const { return ns_ubatch[i_cur]; } -llama_memory_hybrid_idx::ple_history & llama_memory_hybrid_idx_context::get_ple_hist(llama_seq_id seq_id) const { - GGML_ASSERT(mem != nullptr); - - return mem->ple_hist_get(seq_id); -} - void llama_memory_hybrid_idx_context::set_input_qsa( ggml_tensor * cell_blk, ggml_tensor * blk_cells, diff --git a/src/llama-memory-hybrid-idx.h b/src/llama-memory-hybrid-idx.h index a17e184209..c2567f79c6 100644 --- a/src/llama-memory-hybrid-idx.h +++ b/src/llama-memory-hybrid-idx.h @@ -3,7 +3,6 @@ #include "llama-memory-hybrid.h" #include -#include #include // @@ -78,42 +77,12 @@ public: llama_kv_cache * get_mem_idx() const; // nullptr when the model carries no indexer - // [TAG_PLE_HISTORY] - // The qwen4exp PLE hash of a token mixes in the ple_ngram_size - 1 tokens before it, which - // a decode ubatch does not carry. It lives here because it is per-context per-sequence - // state: it must follow the seq_* operations and the state blob, like the caches next to it. - struct ple_history { - // position the next token of this sequence must have; -1 means the window is not trusted - llama_pos next_pos = -1; - - // the tokens at [next_pos - toks.size(), next_pos), oldest first, at most ple_ngram_size - 1 - // it can be shorter near a sequence start or after a rewind; the caller pads the front with EOS - std::vector toks; - }; - - // history for seq_id, default-constructed (and so untrusted) on first use - // const because set_input updates it through a const memory context - ple_history & ple_hist_get(llama_seq_id seq_id) const; - private: // the indexer cache holds one key head per layer, so it needs its own hparams: // llama_kv_cache keeps a reference to what it is given llama_hparams hparams_idx; const std::unique_ptr mem_idx; - - // [TAG_PLE_HISTORY] empty for every architecture but qwen4exp, the only one that asks for a history - mutable std::unordered_map ple_hist; - - // the seq_* halves of the history bookkeeping, one per llama_memory_i operation - void ple_hist_rm (llama_seq_id seq_id, llama_pos p0, llama_pos p1); - void ple_hist_cp (llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1); - void ple_hist_keep(llama_seq_id seq_id); - void ple_hist_add (llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift); - void ple_hist_div (llama_seq_id seq_id, llama_pos p0, llama_pos p1); - - void ple_hist_state_write(llama_io_write_i & io, llama_seq_id seq_id) const; - void ple_hist_state_read (llama_io_read_i & io, llama_seq_id seq_id); }; class llama_memory_hybrid_idx_context : public llama_memory_hybrid_context { @@ -158,9 +127,6 @@ public: // streams in the current slot info, the `ns` of get_k/get_v; 1 if unified uint32_t get_n_stream() const; - // [TAG_PLE_HISTORY] the per-sequence n-gram history of the owning memory, for set_input - llama_memory_hybrid_idx::ple_history & get_ple_hist(llama_seq_id seq_id) const; - // block-compressed sparse attention (qwen4exp QSA) over the cells of the indexer cache. // Blocks cut the position line, not the cell array, so no caller assumes a contiguous layout: // cell_blk I32 [n_kv, ns] block each cell belongs to diff --git a/src/models/qwen4exp.cpp b/src/models/qwen4exp.cpp index c438710995..e8d390dbd8 100644 --- a/src/models/qwen4exp.cpp +++ b/src/models/qwen4exp.cpp @@ -866,7 +866,7 @@ ggml_tensor * llama_model_qwen4exp::graph::build_layer_ffn(ggml_tensor * cur, co class llm_graph_input_ple : public llm_graph_input_i { public: llm_graph_input_ple(const llama_model_qwen4exp & pmodel, - const llama_memory_hybrid_idx_context * mctx) : pmodel(pmodel), mctx(mctx) {} + const llama_kv_cache_context * mctx) : pmodel(pmodel), mctx(mctx) {} virtual ~llm_graph_input_ple() = default; void set_input(const llama_ubatch * ubatch) override; @@ -875,8 +875,11 @@ public: const llama_model_qwen4exp & pmodel; - // the token history lives on the memory, so it is per context and part of the state blob - const llama_memory_hybrid_idx_context * mctx; + // the predecessor tokens live in the attention KV cells (ext.tok) + const llama_kv_cache_context * mctx; + + // scratch, reused across set_input() calls + std::vector prev; }; void llm_graph_input_ple::set_input(const llama_ubatch * ubatch) { @@ -899,66 +902,33 @@ void llm_graph_input_ple::set_input(const llama_ubatch * ubatch) { const int64_t n_heads = hp.ple_n_heads; const int64_t per_gram = hp.ple_heads_per_ngram; const int64_t eos = hp.ple_eos_token_id; + const int64_t n_prev = n_gram - 1; std::vector idx(n_heads * n_tokens); - // missing predecessors come from the per-sequence history, but only when it is - // contiguous with the incoming position; otherwise the window is EOS-padded GGML_ASSERT(mctx != nullptr); - // snapshot the history first, so a token cannot read an earlier token of this same ubatch - // the snapshot is always n_gram - 1 long and EOS-padded at the front: prev() puts the most recent token last - std::unordered_map> snap; for (int64_t i = 0; i < n_tokens; ++i) { - const llama_seq_id seq = ubatch->seq_id[i][0]; - if (snap.count(seq)) { - continue; - } - auto & h = mctx->get_ple_hist(seq); - if (h.next_pos != ubatch->pos[i]) { - h.next_pos = ubatch->pos[i]; - h.toks.clear(); - } - if ((int64_t) h.toks.size() > n_gram - 1) { - h.toks.erase(h.toks.begin(), h.toks.end() - (n_gram - 1)); - } - - std::vector padded(n_gram - 1, (llama_token) eos); - std::copy(h.toks.begin(), h.toks.end(), padded.end() - (int64_t) h.toks.size()); - snap[seq] = std::move(padded); + // the preceding tokens would be ambiguous, see get_prev_tokens() + GGML_ASSERT(ubatch->n_seq_id[i] == 1 && "PLE n-gram embeddings do not support tokens shared by multiple sequences"); } + // predecessors come from the KV cells (ext.tok); apply_ubatch() has already stored the + // current ubatch, so predecessors within this very ubatch are covered as well + mctx->get_prev_tokens(*ubatch, n_prev, prev); + for (int64_t i = 0; i < n_tokens; ++i) { - const llama_seq_id seq = ubatch->seq_id[i][0]; - const llama_pos pos = ubatch->pos[i]; - - const auto & hist = snap[seq]; - - // predecessor s (1-based) of this token, EOS past a segment boundary - auto prev = [&](int64_t s) -> int64_t { - const int64_t j = i - s; - if (j >= 0 && ubatch->seq_id[j][0] == seq && ubatch->pos[j] == pos - s) { - return tok_of(j); - } - // s - i positions before this ubatch started, most recent last - const int64_t back = s - i; - const int64_t k = (int64_t) hist.size() - back; - if (back > 0 && k >= 0 && k < (int64_t) hist.size() && pos - s >= 0) { - return hist[k]; - } - return eos; - }; - - // an EOS in the window resets everything at or before it + // an EOS in the window resets everything at or before it, and a missing predecessor + // (before the sequence start, or no cached cell) reads as EOS // the EOS of the token itself does not cut its own context, as in the reference std::vector ctx(n_gram); ctx[0] = tok_of(i); bool cut = false; for (int64_t s = 1; s < n_gram; ++s) { - ctx[s] = cut ? eos : prev(s); - if (ctx[s] == eos) { - cut = true; - } + // predecessor s positions back; prev[] is oldest-first, missing entries are LLAMA_TOKEN_NULL + const llama_token t = cut ? LLAMA_TOKEN_NULL : prev[i*n_prev + (n_prev - s)]; + cut = cut || t < 0 || t == eos; + ctx[s] = cut ? eos : t; } for (int64_t n = 2; n <= n_gram; ++n) { @@ -973,13 +943,6 @@ void llm_graph_input_ple::set_input(const llama_ubatch * ubatch) { (int32_t) (mixed % hp.ple_head_vocab_sizes[h_i] + hp.ple_head_offsets[h_i]); } } - - auto & h = mctx->get_ple_hist(seq); - h.toks.push_back(tok_of(i)); - if ((int64_t) h.toks.size() > n_gram - 1) { - h.toks.erase(h.toks.begin(), h.toks.end() - (n_gram - 1)); - } - h.next_pos = pos + 1; } ggml_backend_tensor_set(rows, idx.data(), 0, idx.size()*ggml_element_size(rows)); @@ -1048,8 +1011,9 @@ ggml_tensor * llama_model_qwen4exp::graph::build_ple( const int64_t hc_dim = hc * n_embd; const int64_t n_heads = hparams.ple_n_heads; + // the attention cells see every ubatch regardless of the layer types auto ple_inp = std::make_unique( - static_cast(model), mctx_hyb); + static_cast(model), mctx_hyb->get_attn()); ple_inp->rows = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_heads * n_tokens); ggml_set_input(ple_inp->rows);