diff --git a/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.cpp b/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.cpp index 5644b13e6..f4c234268 100644 --- a/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.cpp +++ b/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.cpp @@ -139,35 +139,6 @@ void Ngram::insertWorker() { } } -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&, 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 < 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_, 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, diff --git a/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.h b/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.h index 43c0b5915..72d972306 100644 --- a/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.h +++ b/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.h @@ -61,8 +61,6 @@ class Ngram { std::vector> listExternalCorpora() const; - Result batchMatch(const std::vector>& tokens); - Result batchMatch( const std::vector& state_ids, const std::vector>& tokens, 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 be059e419..7e6163874 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 @@ -51,24 +51,6 @@ struct NgramCorpusObj : public tvm::ffi::Object { ngram_->asyncInsert(std::move(tokens)); } - void batch_match( - const tvm::ffi::TensorView tokens_flat, - const tvm::ffi::TensorView offsets, - const tvm::ffi::TensorView out_tokens, - const tvm::ffi::TensorView out_mask) { - auto* data = static_cast(tokens_flat.data_ptr()); - auto* offs = static_cast(offsets.data_ptr()); - int64_t batch_size = offsets.size(0) - 1; - - std::vector> tokens(batch_size); - for (int64_t i = 0; i < batch_size; ++i) { - tokens[i].assign(data + offs[i], data + offs[i + 1]); - } - - 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, @@ -173,7 +155,6 @@ void register_ngram_corpus() { refl::ObjectDef() .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("start_external_corpus_load", &NgramCorpusObj::start_external_corpus_load) diff --git a/python/sglang/jit_kernel/ngram_corpus.py b/python/sglang/jit_kernel/ngram_corpus.py index da0d56a89..d2121417c 100644 --- a/python/sglang/jit_kernel/ngram_corpus.py +++ b/python/sglang/jit_kernel/ngram_corpus.py @@ -74,23 +74,6 @@ def get_ngram_corpus_cls(): tokens_flat, offsets = _to_csr(batch_tokens) self.async_insert(tokens_flat, offsets) # type: ignore - def match( - self, - batch_tokens: List[List[int]], - ) -> Tuple[np.ndarray, np.ndarray]: - tokens_flat, offsets = _to_csr(batch_tokens) - batch_size = len(batch_tokens) - d = self._draft_token_num - - out_tokens = torch.zeros(batch_size * d, dtype=torch.int32) - out_mask = torch.zeros(batch_size * d * d, dtype=torch.uint8) - - self.batch_match(tokens_flat, offsets, out_tokens, out_mask) # type: ignore - - return out_tokens.numpy().astype(np.int64), out_mask.numpy().astype( - np.int64 - ) - def match_stateful( self, state_ids: List[int],