diff --git a/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.cpp b/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.cpp index e6f19d346..5644b13e6 100644 --- a/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.cpp +++ b/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.cpp @@ -115,14 +115,14 @@ void Ngram::clearExternalCorpus() { staging_sam_.reset(); } -std::vector Ngram::listExternalCorpora() const { +std::vector> Ngram::listExternalCorpora() const { std::unique_lock lock(mutex_); - std::vector ids; - ids.reserve(sams_.size()); - for (const auto& [id, _] : sams_) { - ids.push_back(id); + std::vector> entries; + entries.reserve(sams_.size()); + for (const auto& [id, sam] : sams_) { + entries.emplace_back(id, sam->tokenCount()); } - return ids; + return entries; } void Ngram::insertWorker() { diff --git a/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.h b/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.h index fffa88ef5..43c0b5915 100644 --- a/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.h +++ b/python/sglang/jit_kernel/csrc/ngram_corpus/ngram.h @@ -59,7 +59,7 @@ class Ngram { void clearExternalCorpus(); - std::vector listExternalCorpora() const; + std::vector> listExternalCorpora() const; Result batchMatch(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 a4e31301b..be059e419 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 @@ -129,11 +129,11 @@ struct NgramCorpusObj : public tvm::ffi::Object { } std::string list_external_corpora() { - auto ids = ngram_->listExternalCorpora(); + auto entries = ngram_->listExternalCorpora(); std::string result; - for (size_t i = 0; i < ids.size(); ++i) { + for (size_t i = 0; i < entries.size(); ++i) { if (i > 0) result += "\n"; - result += ids[i]; + result += entries[i].first + "\t" + std::to_string(entries[i].second); } return result; } diff --git a/python/sglang/jit_kernel/csrc/ngram_corpus/suffix_automaton.h b/python/sglang/jit_kernel/csrc/ngram_corpus/suffix_automaton.h index ebd8f5471..6cecfe5d5 100644 --- a/python/sglang/jit_kernel/csrc/ngram_corpus/suffix_automaton.h +++ b/python/sglang/jit_kernel/csrc/ngram_corpus/suffix_automaton.h @@ -39,6 +39,10 @@ class SuffixAutomaton { return !loaded_; } + int64_t tokenCount() const { + return pos_; + } + Result buildRecency( const int32_t* context, size_t len, int32_t last_token, size_t draft_token_num, const Param& param) const; diff --git a/python/sglang/jit_kernel/ngram_corpus.py b/python/sglang/jit_kernel/ngram_corpus.py index 52e72b17c..da0d56a89 100644 --- a/python/sglang/jit_kernel/ngram_corpus.py +++ b/python/sglang/jit_kernel/ngram_corpus.py @@ -1,7 +1,7 @@ from __future__ import annotations from collections.abc import Iterable, Sequence -from typing import List, Tuple +from typing import Dict, List, Tuple import numpy as np import torch @@ -144,10 +144,14 @@ def get_ngram_corpus_cls(): def remove_corpus(self, corpus_id: str) -> None: self.remove_external_corpus(corpus_id) # type: ignore - def list_corpora(self) -> List[str]: + def list_corpora(self) -> Dict[str, int]: result = self.list_external_corpora() # type: ignore if not result: - return [] - return result.split("\n") + return {} + out: Dict[str, int] = {} + for line in result.split("\n"): + corpus_id, token_count = line.split("\t", 1) + out[corpus_id] = int(token_count) + return out return NgramCorpusFFI diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index c460141c3..3d23a2c41 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -814,7 +814,7 @@ async def list_external_corpora(): return ORJSONResponse( { "success": result.success, - "corpus_ids": result.corpus_ids, + "corpus_token_counts": result.corpus_token_counts, "message": result.message, }, status_code=200 if result.success else HTTPStatus.BAD_REQUEST, diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index b0657777f..f06b40c4e 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -1232,7 +1232,7 @@ class ListExternalCorporaReqInput(BaseReq): @dataclass class ListExternalCorporaReqOutput(BaseReq): success: bool - corpus_ids: List[str] = field(default_factory=list) + corpus_token_counts: Dict[str, int] = field(default_factory=dict) message: str = "" diff --git a/python/sglang/srt/managers/tokenizer_communicator_mixin.py b/python/sglang/srt/managers/tokenizer_communicator_mixin.py index 5b413dcc4..b035cac37 100644 --- a/python/sglang/srt/managers/tokenizer_communicator_mixin.py +++ b/python/sglang/srt/managers/tokenizer_communicator_mixin.py @@ -473,10 +473,12 @@ class TokenizerCommunicatorMixin: ListExternalCorporaReqInput() ) all_success, all_message = _Communicator.merge_results(results) - # Merge corpus IDs from all DP ranks (each rank loads the same set). - corpus_ids = results[0].corpus_ids if all_success else [] + # Merge corpus token counts from all DP ranks (each rank loads the same set). + corpus_token_counts = results[0].corpus_token_counts if all_success else {} return ListExternalCorporaReqOutput( - success=all_success, corpus_ids=corpus_ids, message=all_message + success=all_success, + corpus_token_counts=corpus_token_counts, + message=all_message, ) async def flush_cache( diff --git a/python/sglang/srt/speculative/cpp_ngram/ngram_corpus.py b/python/sglang/srt/speculative/cpp_ngram/ngram_corpus.py index 6cc15a115..4ffc299fd 100644 --- a/python/sglang/srt/speculative/cpp_ngram/ngram_corpus.py +++ b/python/sglang/srt/speculative/cpp_ngram/ngram_corpus.py @@ -90,7 +90,7 @@ class NgramCorpus: old_count = self._corpus_token_counts.pop(corpus_id, 0) self._total_loaded_tokens -= old_count - def list_external_corpora(self) -> List[str]: + def list_external_corpora(self) -> Dict[str, int]: return self._obj.list_corpora() def reset(self): diff --git a/python/sglang/srt/speculative/external_corpus_manager.py b/python/sglang/srt/speculative/external_corpus_manager.py index dd58a0eed..b268af5e1 100644 --- a/python/sglang/srt/speculative/external_corpus_manager.py +++ b/python/sglang/srt/speculative/external_corpus_manager.py @@ -101,7 +101,10 @@ class ExternalCorpusManager: self, recv_req: ListExternalCorporaReqInput ) -> ListExternalCorporaReqOutput: try: - ids = self._worker.list_external_corpora() - return ListExternalCorporaReqOutput(success=True, corpus_ids=ids) + token_counts = self._worker.list_external_corpora() + return ListExternalCorporaReqOutput( + success=True, + corpus_token_counts=token_counts, + ) except Exception as e: return ListExternalCorporaReqOutput(success=False, message=str(e)) diff --git a/python/sglang/srt/speculative/ngram_worker.py b/python/sglang/srt/speculative/ngram_worker.py index 4c0e79503..d219de746 100644 --- a/python/sglang/srt/speculative/ngram_worker.py +++ b/python/sglang/srt/speculative/ngram_worker.py @@ -90,7 +90,7 @@ class NGRAMWorker: def remove_external_corpus(self, corpus_id: str) -> None: self.ngram_corpus.remove_external_corpus(corpus_id) - def list_external_corpora(self) -> list[str]: + def list_external_corpora(self) -> dict[str, int]: return self.ngram_corpus.list_external_corpora() def _efficient_concat_last_n(self, seq1: List[int], seq2: List[int], n: int): diff --git a/test/registered/unit/spec/test_ngram_corpus.py b/test/registered/unit/spec/test_ngram_corpus.py index 1ecc30b4a..97b0906e1 100644 --- a/test/registered/unit/spec/test_ngram_corpus.py +++ b/test/registered/unit/spec/test_ngram_corpus.py @@ -925,8 +925,10 @@ class TestNgramCorpusMultiSam(CustomTestCase): "b", [[10, 20, 30, 40, 50]] ) corpus.commit_external_corpus_load("b", loaded_token_count) - ids = corpus.list_external_corpora() - self.assertEqual(sorted(ids), ["a", "b"]) + token_counts = corpus.list_external_corpora() + self.assertEqual(sorted(token_counts.keys()), ["a", "b"]) + self.assertEqual(token_counts["a"], 5) + self.assertEqual(token_counts["b"], 5) def test_remove(self): corpus = _make_corpus("BFS", draft_token_num=4, external_sam_budget=3) @@ -937,12 +939,12 @@ class TestNgramCorpusMultiSam(CustomTestCase): ) corpus.commit_external_corpus_load("b", loaded_token_count) corpus.remove_external_corpus("a") - self.assertEqual(corpus.list_external_corpora(), ["b"]) + self.assertEqual(list(corpus.list_external_corpora().keys()), ["b"]) def test_remove_nonexistent_is_noop(self): corpus = _make_corpus("BFS", draft_token_num=4, external_sam_budget=3) corpus.remove_external_corpus("nonexistent") - self.assertEqual(corpus.list_external_corpora(), []) + self.assertEqual(corpus.list_external_corpora(), {}) def test_multi_sam_candidates(self): corpus = _make_corpus("BFS", draft_token_num=6, external_sam_budget=4) @@ -983,8 +985,8 @@ class TestNgramCorpusMultiSam(CustomTestCase): external_sam_budget=3, external_corpus_documents=[[1, 2, 3, 4, 5]], ) - ids = corpus.list_external_corpora() - self.assertIn("test_corpus", ids) + token_counts = corpus.list_external_corpora() + self.assertIn("test_corpus", token_counts) def test_remove_frees_token_budget(self): """Removing a corpus should free its tokens from the total budget.""" @@ -1008,7 +1010,7 @@ class TestNgramCorpusMultiSam(CustomTestCase): # Now there's room for a new corpus. loaded_token_count = corpus.load_external_corpus_named("c", [[100, 200, 300]]) corpus.commit_external_corpus_load("c", loaded_token_count) - self.assertEqual(sorted(corpus.list_external_corpora()), ["b", "c"]) + self.assertEqual(sorted(corpus.list_external_corpora().keys()), ["b", "c"]) def test_duplicate_corpus_id_is_rejected(self): """Adding a duplicate corpus_id should fail without replacing the original corpus.""" @@ -1024,7 +1026,7 @@ class TestNgramCorpusMultiSam(CustomTestCase): corpus.load_external_corpus_named("a", [[10, 20, 30]]) self.assertEqual(corpus.remaining_token_budget, 5) - self.assertEqual(corpus.list_external_corpora(), ["a"]) + self.assertEqual(list(corpus.list_external_corpora().keys()), ["a"]) # The original corpus must still be usable for matching. ids, masks = _batch_get(corpus, [[1, 2, 3]]) @@ -1051,7 +1053,7 @@ class TestNgramCorpusMultiSam(CustomTestCase): with self.assertRaises(ValueError): corpus.load_external_corpus_named("b", [[10, 20, 30, 40, 50, 60]]) - self.assertEqual(corpus.list_external_corpora(), ["a"]) + self.assertEqual(list(corpus.list_external_corpora().keys()), ["a"]) self.assertEqual(corpus.remaining_token_budget, 5) # "a" must still be usable for matching. @@ -1106,7 +1108,7 @@ class TestMultiSamHttpMock(CustomTestCase): ) tm.list_external_corpora = AsyncMock( return_value=ListExternalCorporaReqOutput( - success=True, corpus_ids=["a", "b"] + success=True, corpus_token_counts={"a": 100, "b": 200} ) ) set_global_state(mock_state) @@ -1152,7 +1154,7 @@ class TestMultiSamHttpMock(CustomTestCase): self.assertEqual(resp.status_code, 200) data = resp.json() self.assertTrue(data["success"]) - self.assertEqual(sorted(data["corpus_ids"]), ["a", "b"]) + self.assertEqual(data["corpus_token_counts"], {"a": 100, "b": 200}) if __name__ == "__main__":