From b2008bf9e0821cc5d466bbce1b4ab87417944823 Mon Sep 17 00:00:00 2001 From: Khoa Pham Date: Sun, 5 Apr 2026 22:21:05 -0700 Subject: [PATCH] [Spec][Ngram] 5/N: Store and advance anchor match state across decode steps (#21243) --- .../jit_kernel/csrc/ngram_corpus/ngram.cpp | 58 ++++++- .../jit_kernel/csrc/ngram_corpus/ngram.h | 13 +- .../csrc/ngram_corpus/ngram_corpus_ffi.cpp | 56 +++++-- .../jit_kernel/csrc/ngram_corpus/trie.cpp | 122 ++++++++++++-- .../jit_kernel/csrc/ngram_corpus/trie.h | 77 ++++++++- python/sglang/jit_kernel/ngram_corpus.py | 27 ++++ .../srt/speculative/cpp_ngram/ngram_corpus.py | 40 ++++- python/sglang/srt/speculative/ngram_worker.py | 14 +- .../spec/utils/test_ngram_corpus.py | 152 +++++++++++++++--- 9 files changed, 499 insertions(+), 60 deletions(-) diff --git a/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.cpp b/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.cpp index 0b90fa812..c79d7cd7d 100644 --- a/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.cpp +++ b/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.cpp @@ -77,10 +77,10 @@ void Ngram::insertWorker() { } } -Result Ngram::batchMatch(const std::vector>& tokens) const { +Result Ngram::batchMatch(const std::vector>& tokens) { std::unique_lock lock(mutex_); - using BuildFn = Result (Trie::*)(const int32_t*, size_t, int32_t, size_t, const Param&) const; + using BuildFn = Result (Trie::*)(const int32_t*, size_t, int32_t, size_t, const Param&, MatchState&, size_t) const; BuildFn build_fn; if (param_.match_type == "BFS") { build_fn = &Trie::buildRecency; @@ -91,13 +91,63 @@ Result Ngram::batchMatch(const std::vector>& tokens) const } Result merged; - for (const auto& suffix : tokens) { + for (size_t i = 0; i < tokens.size(); ++i) { + const auto& suffix = tokens[i]; + if (suffix.empty()) { + throw std::runtime_error("batchMatch received an empty token tail"); + } + MatchState temp_state; auto draft_token_num = param_.get_draft_token_num(tokens.size()); - auto res = (trie_.get()->*build_fn)(suffix.data(), suffix.size(), suffix.back(), draft_token_num, param_); + auto res = (trie_.get()->*build_fn)( + suffix.data(), suffix.size(), suffix.back(), draft_token_num, param_, temp_state, suffix.size()); merged.token.insert(merged.token.end(), res.token.begin(), res.token.end()); merged.mask.insert(merged.mask.end(), res.mask.begin(), res.mask.end()); } return merged; } +Result Ngram::batchMatch( + const std::vector& state_ids, + const std::vector>& tokens, + const std::vector& total_lens) { + if (state_ids.size() != tokens.size() || state_ids.size() != total_lens.size()) { + throw std::runtime_error("batchMatch expects state_ids, tokens, and total_lens to match in size"); + } + + std::unique_lock lock(mutex_); + + using BuildFn = Result (Trie::*)(const int32_t*, size_t, int32_t, size_t, const Param&, MatchState&, size_t) const; + BuildFn build_fn; + if (param_.match_type == "BFS") { + build_fn = &Trie::buildRecency; + } else if (param_.match_type == "PROB") { + build_fn = &Trie::buildFrequency; + } else { + throw std::runtime_error("Unknown match_type: '" + param_.match_type + "'. Must be 'BFS' or 'PROB'."); + } + + Result merged; + for (size_t i = 0; i < state_ids.size(); ++i) { + const auto& suffix = tokens[i]; + if (suffix.empty()) { + throw std::runtime_error("batchMatch received an empty token tail"); + } + + auto& state = match_state_[state_ids[i]]; + auto draft_token_num = param_.get_draft_token_num(tokens.size()); + auto res = (trie_.get()->*build_fn)( + suffix.data(), suffix.size(), suffix.back(), draft_token_num, param_, state, total_lens[i]); + merged.token.insert(merged.token.end(), res.token.begin(), res.token.end()); + merged.mask.insert(merged.mask.end(), res.mask.begin(), res.mask.end()); + } + return merged; +} + +void Ngram::eraseMatchState(const std::vector& state_ids) { + std::unique_lock lock(mutex_); + for (const auto& sid : state_ids) { + match_state_.erase(sid); + } +} + } // namespace ngram diff --git a/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.h b/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.h index fb1461d9a..d1b404974 100644 --- a/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.h +++ b/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.h @@ -9,7 +9,9 @@ #include #include #include +#include #include +#include #include namespace ngram { @@ -28,6 +30,7 @@ class Ngram { size_t pending_count_ = 0; utils::Queue> insert_queue_; std::thread insert_worker_; + std::unordered_map match_state_; public: Ngram(size_t capacity, const Param& param); @@ -37,13 +40,21 @@ class Ngram { void asyncInsert(std::vector>&& tokens); - Result batchMatch(const std::vector>& tokens) const; + Result batchMatch(const std::vector>& tokens); + + Result batchMatch( + const std::vector& state_ids, + const std::vector>& tokens, + const std::vector& total_lens); + + void eraseMatchState(const std::vector& state_ids); void reset() { std::unique_lock lock(mutex_); if (trie_) { trie_->reset(); } + match_state_.clear(); } const Param& param() const { diff --git a/python/sglang/jit_kernel/csrc/ngram_corpus/ngram_corpus_ffi.cpp b/python/sglang/jit_kernel/csrc/ngram_corpus/ngram_corpus_ffi.cpp index e1797e1fc..bbd815c51 100644 --- a/python/sglang/jit_kernel/csrc/ngram_corpus/ngram_corpus_ffi.cpp +++ b/python/sglang/jit_kernel/csrc/ngram_corpus/ngram_corpus_ffi.cpp @@ -62,7 +62,52 @@ struct NgramCorpusObj : public tvm::ffi::Object { } auto result = ngram_->batchMatch(tokens); + write_result_(result, out_tokens, out_mask); + } + void batch_match_stateful( + const tvm::ffi::TensorView state_ids_tv, + const tvm::ffi::TensorView tokens_flat, + const tvm::ffi::TensorView offsets, + const tvm::ffi::TensorView total_lens_tv, + const tvm::ffi::TensorView out_tokens, + const tvm::ffi::TensorView out_mask) { + auto* sid = static_cast(state_ids_tv.data_ptr()); + auto* data = static_cast(tokens_flat.data_ptr()); + auto* offs = static_cast(offsets.data_ptr()); + auto* tlens = static_cast(total_lens_tv.data_ptr()); + int64_t batch_size = offsets.size(0) - 1; + + std::vector state_ids(sid, sid + batch_size); + std::vector> tokens(batch_size); + std::vector total_lens(batch_size); + for (int64_t i = 0; i < batch_size; ++i) { + tokens[i].assign(data + offs[i], data + offs[i + 1]); + total_lens[i] = static_cast(tlens[i]); + } + + auto result = ngram_->batchMatch(state_ids, tokens, total_lens); + write_result_(result, out_tokens, out_mask); + } + + void erase_match_state(const tvm::ffi::TensorView state_ids_tv) { + auto* sid = static_cast(state_ids_tv.data_ptr()); + int64_t n = state_ids_tv.size(0); + std::vector state_ids(sid, sid + n); + ngram_->eraseMatchState(state_ids); + } + + void synchronize() { + ngram_->synchronize(); + } + + void reset() { + ngram_->reset(); + } + + private: + void write_result_( + const ngram::Result& result, const tvm::ffi::TensorView& out_tokens, const tvm::ffi::TensorView& out_mask) { auto* out_tok = static_cast(out_tokens.data_ptr()); auto* out_msk = static_cast(out_mask.data_ptr()); if (result.token.size() > static_cast(out_tokens.size(0))) { @@ -79,15 +124,6 @@ struct NgramCorpusObj : public tvm::ffi::Object { std::memcpy(out_msk, result.mask.data(), result.mask.size() * sizeof(uint8_t)); } - void synchronize() { - ngram_->synchronize(); - } - - void reset() { - ngram_->reset(); - } - - private: std::unique_ptr ngram_; }; @@ -97,6 +133,8 @@ void register_ngram_corpus() { .def(refl::init(), "__init__") .def("async_insert", &NgramCorpusObj::async_insert) .def("batch_match", &NgramCorpusObj::batch_match) + .def("batch_match_stateful", &NgramCorpusObj::batch_match_stateful) + .def("erase_match_state", &NgramCorpusObj::erase_match_state) .def("synchronize", &NgramCorpusObj::synchronize) .def("reset", &NgramCorpusObj::reset); } diff --git a/python/sglang/jit_kernel/csrc/ngram_corpus/trie.cpp b/python/sglang/jit_kernel/csrc/ngram_corpus/trie.cpp index 67058eccb..a429a5576 100644 --- a/python/sglang/jit_kernel/csrc/ngram_corpus/trie.cpp +++ b/python/sglang/jit_kernel/csrc/ngram_corpus/trie.cpp @@ -84,12 +84,14 @@ void Trie::squeeze(size_t count) { last->parent->lru.erase(last->parent_lru_pos); last->parent->sorted_children.erase(last); last->parent->child.erase(last->token); + retireNode(last); node_pool_[free_node_count_++] = last; } } void Trie::reset() { + ++trie_epoch_; global_lru_.clear(); path_.clear(); node_pool_.clear(); @@ -100,11 +102,31 @@ void Trie::reset() { root_ = getNode(); } -std::vector> Trie::match(const int32_t* context, size_t len) const { - std::vector> result; +const TrieNode* Trie::resolve(const MatchState& state, const NodeRef& ref) const { + if (ref.ptr == nullptr || state.trie_epoch != trie_epoch_ || ref.ptr->version != ref.version) { + return nullptr; + } + return ref.ptr; +} + +bool Trie::validateMatchState_(const MatchState& state) const { + if (state.trie_epoch != trie_epoch_) { + return false; + } + for (const auto& ref : state.anchors) { + if (ref.ptr && !resolve(state, ref)) { + return false; + } + } + return true; +} + +void Trie::rebuildMatchState_(const int32_t* context, size_t len, MatchState& state, size_t total_len) const { const auto max_match_depth = std::min(len, param_.max_trie_depth); - result.reserve(max_match_depth); - for (size_t match_depth = max_match_depth; match_depth > 0; --match_depth) { + state.trie_epoch = trie_epoch_; + state.processed_total_len = total_len; + state.anchors.assign(max_match_depth, {}); + for (size_t match_depth = 1; match_depth <= max_match_depth; ++match_depth) { auto start = context + len - match_depth; auto end = start + match_depth; auto cursor = root_; @@ -117,17 +139,88 @@ std::vector> Trie::match(const int32_t* context, s ++start; cursor = iter->second; } - if (cursor != nullptr && !cursor->child.empty()) { - result.emplace_back(cursor, static_cast(match_depth)); + if (cursor != nullptr) { + state.anchors[match_depth - 1] = capture(cursor); + } + } +} + +bool Trie::advanceMatchState_(MatchState& state, const int32_t* tokens, size_t len, size_t total_len) const { + if (!validateMatchState_(state)) { + return false; + } + + for (size_t i = 0; i < len; ++i) { + const auto next_depth = std::min(state.anchors.size() + 1, param_.max_trie_depth); + std::vector next(next_depth); + + const auto root_ref = rootRef(); + const auto root = resolve(state, root_ref); + if (root == nullptr) { + return false; + } + if (auto iter = root->child.find(tokens[i]); iter != root->child.end()) { + next[0] = capture(iter->second); + } + + for (size_t depth = 1; depth < next_depth; ++depth) { + const auto& prev_ref = state.anchors[depth - 1]; + if (prev_ref.ptr == nullptr) { + continue; + } + const auto prev_node = resolve(state, prev_ref); + if (prev_node == nullptr) { + return false; + } + if (auto iter = prev_node->child.find(tokens[i]); iter != prev_node->child.end()) { + next[depth] = capture(iter->second); + } + } + + state.anchors.swap(next); + } + + state.processed_total_len = total_len; + return true; +} + +std::vector> Trie::getExpandableAnchors_(const MatchState& state) const { + std::vector> result; + result.reserve(state.anchors.size()); + for (size_t depth = state.anchors.size(); depth > 0; --depth) { + const auto node = resolve(state, state.anchors[depth - 1]); + if (node != nullptr && !node->child.empty()) { + result.emplace_back(node, static_cast(depth)); } } return result; } -Result Trie::buildRecency( - const int32_t* context, size_t len, int32_t last_token, size_t draft_token_num, const Param& param) const { - auto anchors = match(context, len); +std::vector> +Trie::match(const int32_t* context, size_t len, MatchState& state, size_t total_len) const { + const bool has_forward_progress = total_len >= state.processed_total_len; + const auto appended_len = has_forward_progress ? total_len - state.processed_total_len : 0; + const auto expected_prev_depth = std::min(state.processed_total_len, param_.max_trie_depth); + const bool can_advance = state.trie_epoch == trie_epoch_ && has_forward_progress && appended_len <= len && + state.anchors.size() == expected_prev_depth; + if (can_advance && advanceMatchState_(state, context + len - appended_len, appended_len, total_len)) { + return getExpandableAnchors_(state); + } + + rebuildMatchState_(context, len, state, total_len); + return getExpandableAnchors_(state); +} + +Result Trie::buildRecency( + const int32_t* context, + size_t len, + int32_t last_token, + size_t draft_token_num, + const Param& param, + MatchState& state, + size_t total_len) const { + auto anchors = match(context, len, state, total_len); const auto max_match_depth = std::max(1, static_cast(param.max_trie_depth - 1)); double bfs_breadth_scale = double(param.max_bfs_breadth - param.min_bfs_breadth) / max_match_depth; @@ -166,9 +259,14 @@ Result Trie::buildRecency( } Result Trie::buildFrequency( - const int32_t* context, size_t len, int32_t last_token, size_t draft_token_num, const Param& param) const { - auto anchors = match(context, len); - + const int32_t* context, + size_t len, + int32_t last_token, + size_t draft_token_num, + const Param& param, + MatchState& state, + size_t total_len) const { + auto anchors = match(context, len, state, total_len); struct CompareByLastDouble { bool operator()( const std::tuple& a, diff --git a/python/sglang/jit_kernel/csrc/ngram_corpus/trie.h b/python/sglang/jit_kernel/csrc/ngram_corpus/trie.h index bd555597d..909ffae8f 100644 --- a/python/sglang/jit_kernel/csrc/ngram_corpus/trie.h +++ b/python/sglang/jit_kernel/csrc/ngram_corpus/trie.h @@ -22,6 +22,9 @@ struct TrieNode { TrieNode* parent; std::list lru; int32_t freq = 0; + // Logical generation of this TrieNode. retireNode() bumps it before the node + // goes back to the pool so stale NodeRefs fail validation after reuse. + uint64_t version = 1; struct CompareByFreq { bool operator()(TrieNode* a, TrieNode* b) const { @@ -31,6 +34,23 @@ struct TrieNode { std::multiset sorted_children; }; +// By-value handle to a logical trie location, cached in MatchState. +// We cannot cache TrieNode* alone across decode steps: squeeze() may evict a +// node, and getNode() may later recycle the same address for a different node. +struct NodeRef { + TrieNode* ptr = nullptr; + uint64_t version = 0; +}; + +// Per-request cached anchors. anchors[d - 1] caches the trie match for the +// length-d suffix ending at the current last token; processed_total_len records +// the full request length covered by those cached anchors. +struct MatchState { + uint64_t trie_epoch = 0; + size_t processed_total_len = 0; + std::vector anchors; +}; + class Trie { public: Trie(size_t capacity, const Param& param); @@ -38,22 +58,72 @@ class Trie { void insert(const int32_t* tokens, size_t len); Result buildRecency( - const int32_t* context, size_t len, int32_t last_token, size_t draft_token_num, const Param& param) const; + const int32_t* context, + size_t len, + int32_t last_token, + size_t draft_token_num, + const Param& param, + MatchState& state, + size_t total_len) const; Result buildFrequency( - const int32_t* context, size_t len, int32_t last_token, size_t draft_token_num, const Param& param) const; + const int32_t* context, + size_t len, + int32_t last_token, + size_t draft_token_num, + const Param& param, + MatchState& state, + size_t total_len) const; void squeeze(size_t count); void reset(); private: - std::vector> match(const int32_t* context, size_t len) const; + // Stateful suffix matcher. If `state` still represents the previous step for + // this request, infer the newly appended suffix from (`context`, `total_len`) + // and advance anchors incrementally; otherwise rebuild the cached anchors from + // `context`. Returns only the suffix matches that are currently expandable. + std::vector> + match(const int32_t* context, size_t len, MatchState& state, size_t total_len) const; + // Recompute all cached anchors from the current tail. After this, for every + // d in [1, min(len, max_trie_depth)], anchors[d - 1] represents the suffix of + // length d ending at context[len - 1]. + void rebuildMatchState_(const int32_t* context, size_t len, MatchState& state, size_t total_len) const; + // Advance the cached anchors by consuming the newly appended suffix one + // token at a time, without re-walking all suffixes from root. + bool advanceMatchState_(MatchState& state, const int32_t* tokens, size_t len, size_t total_len) const; + // Check that every non-empty cached NodeRef in MatchState still resolves to + // the same logical trie node under the current trie_epoch_. + bool validateMatchState_(const MatchState& state) const; + // MatchState keeps all live suffix matches, including leaves. This helper + // filters the cached anchors down to the suffixes that currently have children and + // therefore can seed BFS / PROB draft construction. + std::vector> getExpandableAnchors_(const MatchState& state) const; + // Resolve a cached NodeRef back to a live trie node. nullptr means the + // cached location went stale and the caller should rebuild from context. + const TrieNode* resolve(const MatchState& state, const NodeRef& ref) const; + NodeRef rootRef() const { + return NodeRef{root_, root_->version}; + } + NodeRef capture(TrieNode* node) const { + if (node == nullptr) { + return {}; + } + return NodeRef{node, node->version}; + } + void retireNode(TrieNode* node) { + if (node != nullptr) { + ++node->version; + } + } TrieNode* getNode() { auto node = node_pool_[--free_node_count_]; + auto version = node->version; node->~TrieNode(); new (node) TrieNode(); + node->version = version; return node; } @@ -64,6 +134,7 @@ class Trie { TrieNode* root_; std::vector path_; Param param_; + uint64_t trie_epoch_ = 1; }; } // namespace ngram diff --git a/python/sglang/jit_kernel/ngram_corpus.py b/python/sglang/jit_kernel/ngram_corpus.py index 42b6babab..2182c5c34 100644 --- a/python/sglang/jit_kernel/ngram_corpus.py +++ b/python/sglang/jit_kernel/ngram_corpus.py @@ -85,4 +85,31 @@ def get_ngram_corpus_cls(): np.int64 ) + def match_stateful( + self, + state_ids: List[int], + batch_tokens: List[List[int]], + total_lens: List[int], + ) -> Tuple[np.ndarray, np.ndarray]: + tokens_flat, offsets = _to_csr(batch_tokens) + batch_size = len(batch_tokens) + d = self._draft_token_num + + state_ids_t = torch.tensor(state_ids, dtype=torch.int64) + total_lens_t = torch.tensor(total_lens, dtype=torch.int64) + out_tokens = torch.zeros(batch_size * d, dtype=torch.int32) + out_mask = torch.zeros(batch_size * d * d, dtype=torch.uint8) + + self.batch_match_stateful( # type: ignore + state_ids_t, tokens_flat, offsets, total_lens_t, out_tokens, out_mask + ) + + return out_tokens.numpy().astype(np.int64), out_mask.numpy().astype( + np.int64 + ) + + def erase_states(self, state_ids: List[int]) -> None: + state_ids_t = torch.tensor(state_ids, dtype=torch.int64) + self.erase_match_state(state_ids_t) # type: ignore + return NgramCorpusFFI diff --git a/python/sglang/srt/speculative/cpp_ngram/ngram_corpus.py b/python/sglang/srt/speculative/cpp_ngram/ngram_corpus.py index 0eb6bd71c..0795e3bf4 100644 --- a/python/sglang/srt/speculative/cpp_ngram/ngram_corpus.py +++ b/python/sglang/srt/speculative/cpp_ngram/ngram_corpus.py @@ -1,7 +1,7 @@ # -*- coding: utf-8 -*- import logging -from typing import List, Tuple +from typing import Dict, List, Tuple import numpy as np @@ -31,6 +31,16 @@ class NgramCorpus: ) self.default_mask = np.ones((1, 1), dtype=np.int64) self.draft_token_num = draft_token_num + self._req_id_to_state_id: Dict[str, int] = {} + self._next_state_id: int = 0 + + def _get_state_id(self, req_id: str) -> int: + sid = self._req_id_to_state_id.get(req_id) + if sid is None: + sid = self._next_state_id + self._next_state_id += 1 + self._req_id_to_state_id[req_id] = sid + return sid def batch_put(self, batch_tokens: List[List[int]]): self._obj.insert(batch_tokens) @@ -40,9 +50,26 @@ class NgramCorpus: def reset(self): self._obj.reset() # type: ignore + self._req_id_to_state_id.clear() + self._next_state_id = 0 - def batch_get(self, batch_tokens: List[List[int]]) -> Tuple[np.ndarray, np.ndarray]: - return self._obj.match(batch_tokens) + def batch_get( + self, + req_ids: List[str], + batch_tokens: List[List[int]], + total_lens: List[int], + ) -> Tuple[np.ndarray, np.ndarray]: + state_ids = [self._get_state_id(rid) for rid in req_ids] + return self._obj.match_stateful(state_ids, batch_tokens, total_lens) + + def erase_match_state(self, req_ids: List[str]): + state_ids = [] + for rid in req_ids: + sid = self._req_id_to_state_id.pop(rid, None) + if sid is not None: + state_ids.append(sid) + if state_ids: + self._obj.erase_states(state_ids) def leaf_paths_from_mask( self, tokens: List[int], tree_mask: List[List[int]] @@ -119,6 +146,11 @@ if __name__ == "__main__": corpus.batch_put(token_ids) corpus.synchronize() - decoding_ids, decoding_masks = corpus.batch_get([[1, 2, 3], [3, 44], [3, 6, 999]]) + queries = [[1, 2, 3], [3, 44], [3, 6, 999]] + decoding_ids, decoding_masks = corpus.batch_get( + req_ids=[f"query-{i}" for i in range(len(queries))], + batch_tokens=queries, + total_lens=[len(q) for q in queries], + ) corpus.debug_result(decoding_ids, decoding_masks) diff --git a/python/sglang/srt/speculative/ngram_worker.py b/python/sglang/srt/speculative/ngram_worker.py index 098338c5e..bf332ebeb 100644 --- a/python/sglang/srt/speculative/ngram_worker.py +++ b/python/sglang/srt/speculative/ngram_worker.py @@ -123,13 +123,19 @@ class NGRAMWorker: bs = batch.batch_size() self.ngram_corpus.synchronize() + req_ids = [] batch_tokens = [] + total_lens = [] for req in batch.reqs: check_token = self._efficient_concat_last_n( req.origin_input_ids, req.output_ids, self.max_trie_depth ) + req_ids.append(req.rid) batch_tokens.append(check_token) - req_drafts, mask = self.ngram_corpus.batch_get(batch_tokens) + total_lens.append(len(req.origin_input_ids) + len(req.output_ids)) + req_drafts, mask = self.ngram_corpus.batch_get( + req_ids, batch_tokens, total_lens + ) total_draft_token_num = len(req_drafts) # Check if speculative decoding is needed; here we always enforce it @@ -263,6 +269,12 @@ class NGRAMWorker: if batch.return_logprob: add_output_logprobs_for_spec_v1(batch, verify_input, logits_output) self._update_ngram_corpus(batch) + finished_req_ids = [] + for req in batch.reqs: + if req.finished() or req.is_retracted: + finished_req_ids.append(req.rid) + if finished_req_ids: + self.ngram_corpus.erase_match_state(finished_req_ids) batch.forward_mode = ForwardMode.DECODE else: diff --git a/test/registered/spec/utils/test_ngram_corpus.py b/test/registered/spec/utils/test_ngram_corpus.py index e8d9fc026..f0169306d 100644 --- a/test/registered/spec/utils/test_ngram_corpus.py +++ b/test/registered/spec/utils/test_ngram_corpus.py @@ -1,4 +1,5 @@ import unittest +import uuid import numpy as np @@ -22,6 +23,26 @@ def _make_corpus(match_type="BFS", **kwargs): return NgramCorpus(**defaults) +def _batch_get( + corpus: NgramCorpus, + batch_tokens: list[list[int]], +): + return corpus.batch_get( + req_ids=[uuid.uuid4().hex for _ in range(len(batch_tokens))], + batch_tokens=batch_tokens, + total_lens=[len(tokens) for tokens in batch_tokens], + ) + + +def _batch_get_with_state( + corpus: NgramCorpus, + req_id: str, + current_tokens: list[int], + total_len: int, +): + return corpus.batch_get([req_id], [current_tokens], [total_len]) + + SEED_SEQUENCES = [ [1, 2, 3, 4, 5, 6, 7, 8, 9, 10], [1, 2, 3, 44, 55, 66, 77, 88, 99, 100], @@ -116,7 +137,7 @@ class TestNgramCorpusBFS(CustomTestCase): cls.corpus = _make_corpus("BFS") cls.corpus.batch_put(SEED_SEQUENCES) cls.corpus.synchronize() - ids, masks = cls.corpus.batch_get(QUERY_SEQUENCES) + ids, masks = _batch_get(cls.corpus, QUERY_SEQUENCES) draft = 8 cls.ids = ids.reshape(-1, draft) cls.masks = masks.reshape(-1, draft, draft) @@ -142,7 +163,7 @@ class TestNgramCorpusProb(CustomTestCase): cls.corpus = _make_corpus("PROB") cls.corpus.batch_put(SEED_SEQUENCES) cls.corpus.synchronize() - ids, masks = cls.corpus.batch_get(QUERY_SEQUENCES) + ids, masks = _batch_get(cls.corpus, QUERY_SEQUENCES) cls.ids = ids.reshape(-1, 8) cls.masks = masks.reshape(-1, 8, 8) @@ -166,7 +187,7 @@ class TestNgramCorpusReset(CustomTestCase): corpus.batch_put(SEED_SEQUENCES) corpus.synchronize() - ids_before, _ = corpus.batch_get([[1, 2, 3]]) + ids_before, _ = _batch_get(corpus, [[1, 2, 3]]) self.assertTrue( any(t != 0 for t in ids_before.tolist()[1:]), "Expected non-trivial draft tokens before reset", @@ -174,7 +195,7 @@ class TestNgramCorpusReset(CustomTestCase): corpus.reset() - ids_after, _ = corpus.batch_get([[1, 2, 3]]) + ids_after, _ = _batch_get(corpus, [[1, 2, 3]]) self.assertEqual( ids_after.tolist(), [3, 0, 0, 0, 0, 0, 0, 0], @@ -190,7 +211,7 @@ class TestNgramCorpusNoMatch(CustomTestCase): corpus.batch_put([[10, 20, 30, 40, 50]]) corpus.synchronize() - ids, masks = corpus.batch_get([[999, 888, 777]]) + ids, masks = _batch_get(corpus, [[999, 888, 777]]) ids_list = ids.tolist() self.assertEqual(ids_list[0], 777, "First token should be last context token") self.assertTrue( @@ -200,7 +221,7 @@ class TestNgramCorpusNoMatch(CustomTestCase): def test_empty_corpus(self): corpus = _make_corpus("BFS") - ids, masks = corpus.batch_get([[1, 2, 3]]) + ids, masks = _batch_get(corpus, [[1, 2, 3]]) ids_list = ids.tolist() self.assertEqual(ids_list[0], 3) self.assertTrue(all(t == 0 for t in ids_list[1:])) @@ -217,7 +238,7 @@ class TestNgramCorpusMultipleInserts(CustomTestCase): corpus.batch_put([[1, 2, 3, 44, 55]]) corpus.synchronize() - ids, _ = corpus.batch_get([[1, 2, 3]]) + ids, _ = _batch_get(corpus, [[1, 2, 3]]) ids_list = ids.tolist() self.assertIn(4, ids_list, "Token 4 from first insert should still match") @@ -233,7 +254,7 @@ class TestNgramCorpusSqueeze(CustomTestCase): corpus.batch_put([long_seq]) corpus.synchronize() - ids, masks = corpus.batch_get([[50, 51, 52]]) + ids, masks = _batch_get(corpus, [[50, 51, 52]]) self.assertEqual(len(ids), 8, "Should still produce draft_token_num outputs") def test_eviction_preserves_recent(self): @@ -247,7 +268,7 @@ class TestNgramCorpusSqueeze(CustomTestCase): corpus.batch_put([recent_seq]) corpus.synchronize() - ids, _ = corpus.batch_get([[2000, 2001, 2002]]) + ids, _ = _batch_get(corpus, [[2000, 2001, 2002]]) ids_list = ids.tolist() self.assertEqual(ids_list[0], 2002, "Last context token should be first") self.assertIn(2003, ids_list, "Recent sequence should still be matchable") @@ -294,13 +315,13 @@ class TestNgramCorpusBatchConsistency(CustomTestCase): corpus.batch_put(SEED_SEQUENCES) corpus.synchronize() - batch_ids, batch_masks = corpus.batch_get(QUERY_SEQUENCES) + batch_ids, batch_masks = _batch_get(corpus, QUERY_SEQUENCES) draft = 8 batch_ids = batch_ids.reshape(-1, draft) batch_masks = batch_masks.reshape(-1, draft, draft) for i, query in enumerate(QUERY_SEQUENCES): - single_ids, single_masks = corpus.batch_get([query]) + single_ids, single_masks = _batch_get(corpus, [query]) single_ids = single_ids.reshape(-1, draft) single_masks = single_masks.reshape(-1, draft, draft) @@ -329,7 +350,7 @@ class TestMaskValidity(CustomTestCase): corpus = _make_corpus("BFS") corpus.batch_put(SEED_SEQUENCES) corpus.synchronize() - _, masks = corpus.batch_get(QUERY_SEQUENCES) + _, masks = _batch_get(corpus, QUERY_SEQUENCES) masks = masks.reshape(-1, 8, 8) for i in range(masks.shape[0]): self._check_mask(masks[i].tolist()) @@ -338,7 +359,7 @@ class TestMaskValidity(CustomTestCase): corpus = _make_corpus("PROB") corpus.batch_put(SEED_SEQUENCES) corpus.synchronize() - _, masks = corpus.batch_get(QUERY_SEQUENCES) + _, masks = _batch_get(corpus, QUERY_SEQUENCES) masks = masks.reshape(-1, 8, 8) for i in range(masks.shape[0]): self._check_mask(masks[i].tolist()) @@ -362,7 +383,7 @@ class TestFrequencyBoosting(CustomTestCase): corpus.batch_put([[1, 2, 3, 20, 21]]) corpus.synchronize() - ids, _ = corpus.batch_get([[1, 2, 3]]) + ids, _ = _batch_get(corpus, [[1, 2, 3]]) ids_list = ids.tolist() self.assertEqual( @@ -388,7 +409,7 @@ class TestRecencyOrdering(CustomTestCase): corpus.batch_put([[1, 2, 3, 20, 21]]) corpus.synchronize() - ids, _ = corpus.batch_get([[1, 2, 3]]) + ids, _ = _batch_get(corpus, [[1, 2, 3]]) ids_list = ids.tolist() self.assertEqual( ids_list[1], @@ -406,7 +427,7 @@ class TestOverlappingSuffixes(CustomTestCase): corpus.batch_put([[300, 400, 7, 8, 9, 60, 61]]) corpus.synchronize() - ids, _ = corpus.batch_get([[7, 8, 9]]) + ids, _ = _batch_get(corpus, [[7, 8, 9]]) ids_list = ids.tolist() self.assertIn(50, ids_list, "Continuation from first sequence missing") self.assertIn(60, ids_list, "Continuation from second sequence missing") @@ -420,7 +441,7 @@ class TestSingleTokenContext(CustomTestCase): corpus.batch_put([[5, 10, 20, 30]]) corpus.synchronize() - ids, masks = corpus.batch_get([[5]]) + ids, masks = _batch_get(corpus, [[5]]) ids_list = ids.tolist() self.assertEqual(ids_list[0], 5, "First token should be last context token") self.assertIn(10, ids_list, "Should match continuation after single token 5") @@ -436,7 +457,7 @@ class TestLongContext(CustomTestCase): corpus.synchronize() long_query = list(range(1, 16)) - ids, masks = corpus.batch_get([long_query]) + ids, masks = _batch_get(corpus, [long_query]) ids_list = ids.tolist() self.assertEqual(ids_list[0], 15, "First token should be last context token") self.assertIn(16, ids_list, "Should match via suffix despite long context") @@ -447,7 +468,7 @@ class TestLongContext(CustomTestCase): corpus.batch_put([[99, 3, 4, 5, 6, 8]]) corpus.synchronize() - ids, _ = corpus.batch_get([[2, 3, 4, 5, 6]]) + ids, _ = _batch_get(corpus, [[2, 3, 4, 5, 6]]) ids_list = ids.tolist() self.assertIn( 7, ids_list, "Longest stored suffix should contribute a continuation" @@ -468,7 +489,7 @@ class TestDraftBudgetSaturation(CustomTestCase): corpus.batch_put([seq]) corpus.synchronize() - ids, _ = corpus.batch_get([[1, 2, 3]]) + ids, _ = _batch_get(corpus, [[1, 2, 3]]) ids_list = ids.tolist() self.assertEqual(len(ids_list), 8) non_zero = [t for t in ids_list[1:] if t != 0] @@ -487,7 +508,7 @@ class TestTruncate(CustomTestCase): corpus.batch_put(SEED_SEQUENCES) corpus.synchronize() - ids, masks = corpus.batch_get([[1, 2, 3]]) + ids, _ = _batch_get(corpus, [[1, 2, 3]]) ids = ids.reshape(8) self.assertEqual(len(ids), 8) @@ -501,7 +522,7 @@ class TestTruncate(CustomTestCase): corpus.batch_put(SEED_SEQUENCES) corpus.synchronize() - ids, masks = corpus.batch_get([[1, 2, 3]]) + _, masks = _batch_get(corpus, [[1, 2, 3]]) n = 8 full_mask = masks.reshape(n, n) @@ -530,14 +551,14 @@ class TestResetAndReinsert(CustomTestCase): corpus.batch_put([[10, 20, 30, 40, 50]]) corpus.synchronize() - ids_old, _ = corpus.batch_get([[1, 2, 3]]) + ids_old, _ = _batch_get(corpus, [[1, 2, 3]]) ids_old_list = ids_old.tolist() self.assertTrue( all(t == 0 for t in ids_old_list[1:]), f"Old data should not match after reset+reinsert, got {ids_old_list}", ) - ids_new, _ = corpus.batch_get([[10, 20, 30]]) + ids_new, _ = _batch_get(corpus, [[10, 20, 30]]) ids_new_list = ids_new.tolist() self.assertEqual(ids_new_list[0], 30) self.assertIn(40, ids_new_list, "New data should match after reset+reinsert") @@ -553,7 +574,7 @@ class TestSqueezeEvictsOld(CustomTestCase): corpus.batch_put([old_seq]) corpus.synchronize() - ids_before, _ = corpus.batch_get([[5000, 5001, 5002]]) + ids_before, _ = _batch_get(corpus, [[5000, 5001, 5002]]) self.assertIn( 5003, ids_before.tolist(), @@ -565,7 +586,7 @@ class TestSqueezeEvictsOld(CustomTestCase): corpus.batch_put([new_seq]) corpus.synchronize() - ids_after, _ = corpus.batch_get([[5000, 5001, 5002]]) + ids_after, _ = _batch_get(corpus, [[5000, 5001, 5002]]) ids_after_list = ids_after.tolist() self.assertNotIn( 5003, @@ -574,5 +595,84 @@ class TestSqueezeEvictsOld(CustomTestCase): ) +class TestNgramCorpusIncremental(CustomTestCase): + """Verify the incremental matching path matches the stateless path.""" + + def _assert_incremental_matches_stateless(self, match_type: str): + corpus = _make_corpus(match_type, max_trie_depth=4, draft_token_num=4) + corpus.batch_put([[1, 2, 3, 4, 5, 6], [9, 3, 4, 7, 8]]) + corpus.synchronize() + + req_id = f"req-{match_type.lower()}" + + steps = [ + [1, 2, 3], + [1, 2, 3, 4], + [1, 2, 3, 4, 5, 6], + ] + for full_sequence in steps: + current_tail = full_sequence[-4:] + inc_ids, inc_masks = _batch_get_with_state( + corpus, + req_id, + current_tail, + len(full_sequence), + ) + full_ids, full_masks = _batch_get(corpus, [current_tail]) + np.testing.assert_array_equal(inc_ids, full_ids) + np.testing.assert_array_equal(inc_masks, full_masks) + + def test_incremental_matches_stateless_bfs(self): + self._assert_incremental_matches_stateless("BFS") + + def test_incremental_matches_stateless_prob(self): + self._assert_incremental_matches_stateless("PROB") + + def test_leaf_anchor_becomes_expandable(self): + corpus = _make_corpus("BFS", max_trie_depth=4, draft_token_num=4) + corpus.batch_put([[1, 2, 3]]) + corpus.synchronize() + + req_id = "leaf-anchor" + ids_before, _ = _batch_get_with_state(corpus, req_id, [2, 3], 2) + self.assertTrue( + all(t == 0 for t in ids_before.tolist()[1:]), + f"Expected only the last token before extension, got {ids_before.tolist()}", + ) + + corpus.batch_put([[9, 2, 3, 4]]) + corpus.synchronize() + + inc_ids, inc_masks = _batch_get_with_state(corpus, req_id, [2, 3], 2) + full_ids, full_masks = _batch_get(corpus, [[2, 3]]) + np.testing.assert_array_equal(inc_ids, full_ids) + np.testing.assert_array_equal(inc_masks, full_masks) + self.assertIn( + 4, + inc_ids.tolist(), + f"Expected token 4 after extension, got {inc_ids.tolist()}", + ) + + def test_stale_state_rebuilds_after_eviction(self): + corpus = _make_corpus("BFS", capacity=150, max_trie_depth=6, draft_token_num=4) + corpus.batch_put([list(range(5000, 5030))]) + corpus.synchronize() + + req_id = "evicted" + _batch_get_with_state(corpus, req_id, [5000, 5001, 5002], 3) + + for i in range(5): + new_seq = list(range(6000 + i * 30, 6000 + i * 30 + 30)) + corpus.batch_put([new_seq]) + corpus.synchronize() + + inc_ids, inc_masks = _batch_get_with_state( + corpus, req_id, [5000, 5001, 5002], 3 + ) + full_ids, full_masks = _batch_get(corpus, [[5000, 5001, 5002]]) + np.testing.assert_array_equal(inc_ids, full_ids) + np.testing.assert_array_equal(inc_masks, full_masks) + + if __name__ == "__main__": unittest.main(verbosity=3)