[Spec][Ngram] Return token counts in list_external_corpora API (#22471)

Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Khoa Pham
2026-04-10 21:50:02 -07:00
committed by GitHub
co-authored by Claude Opus 4.6
parent 3c46ff2ac5
commit 04bd8e1218
12 changed files with 49 additions and 34 deletions
@@ -115,14 +115,14 @@ void Ngram::clearExternalCorpus() {
staging_sam_.reset();
}
std::vector<std::string> Ngram::listExternalCorpora() const {
std::vector<std::pair<std::string, int64_t>> Ngram::listExternalCorpora() const {
std::unique_lock<std::mutex> lock(mutex_);
std::vector<std::string> ids;
ids.reserve(sams_.size());
for (const auto& [id, _] : sams_) {
ids.push_back(id);
std::vector<std::pair<std::string, int64_t>> entries;
entries.reserve(sams_.size());
for (const auto& [id, sam] : sams_) {
entries.emplace_back(id, sam->tokenCount());
}
return ids;
return entries;
}
void Ngram::insertWorker() {
@@ -59,7 +59,7 @@ class Ngram {
void clearExternalCorpus();
std::vector<std::string> listExternalCorpora() const;
std::vector<std::pair<std::string, int64_t>> listExternalCorpora() const;
Result batchMatch(const std::vector<std::vector<int32_t>>& tokens);
@@ -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;
}
@@ -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;
+8 -4
View File
@@ -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
+1 -1
View File
@@ -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,
+1 -1
View File
@@ -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 = ""
@@ -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(
@@ -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):
@@ -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))
@@ -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):