[Spec][Ngram] Support multiple SAMs with dynamic HTTP API (#22203)

This commit is contained in:
Liangsheng Yin
2026-04-06 18:49:22 -07:00
committed by GitHub
parent 49cb7d546e
commit e4b1366a46
14 changed files with 685 additions and 121 deletions
@@ -63,27 +63,57 @@ void Ngram::asyncInsert(std::vector<std::vector<int32_t>>&& tokens) {
}
}
// NOTE: staging operations (start/append/finish) are called from a background
// thread during async corpus loading. They do NOT hold mutex_ because
// staging_sam_ is disjoint from sams_ / trie_. Only finishExternalCorpusLoad
// briefly acquires mutex_ when moving the completed SAM into sams_.
void Ngram::startExternalCorpusLoad() {
std::unique_lock<std::mutex> lock(mutex_);
sam_ = std::make_unique<SuffixAutomaton>();
if (staging_sam_) {
throw std::runtime_error("startExternalCorpusLoad called while another load is in progress");
}
staging_sam_ = std::make_unique<SuffixAutomaton>();
}
void Ngram::appendExternalCorpusTokens(const std::vector<int32_t>& tokens) {
std::unique_lock<std::mutex> lock(mutex_);
sam_->appendTokens(tokens);
if (!staging_sam_) {
throw std::runtime_error("appendExternalCorpusTokens called without startExternalCorpusLoad");
}
staging_sam_->appendTokens(tokens);
}
void Ngram::finishExternalCorpusLoad() {
std::unique_lock<std::mutex> lock(mutex_);
sam_->finalize();
if (sam_->empty()) {
sam_.reset();
void Ngram::finishExternalCorpusLoad(const std::string& corpus_id) {
if (!staging_sam_) {
throw std::runtime_error("finishExternalCorpusLoad called without startExternalCorpusLoad");
}
staging_sam_->finalize();
if (staging_sam_->empty()) {
staging_sam_.reset();
throw std::runtime_error("External corpus is empty — no tokens were loaded.");
}
// Only lock briefly to install the completed SAM.
std::unique_lock<std::mutex> lock(mutex_);
sams_[corpus_id] = std::move(staging_sam_);
}
void Ngram::removeExternalCorpus(const std::string& corpus_id) {
std::unique_lock<std::mutex> lock(mutex_);
sams_.erase(corpus_id);
}
void Ngram::clearExternalCorpus() {
std::unique_lock<std::mutex> lock(mutex_);
sam_.reset();
sams_.clear();
staging_sam_.reset();
}
std::vector<std::string> 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);
}
return ids;
}
void Ngram::insertWorker() {
@@ -154,6 +184,14 @@ Result Ngram::batchMatch(
throw std::runtime_error("Unknown match_type: '" + param_.match_type + "'. Must be 'BFS' or 'PROB'.");
}
// All budget values are loop-invariant (mutex_ held, sams_ won't change).
const size_t num_sams = sams_.size();
const auto total_draft_token_num = param_.get_draft_token_num(tokens.size());
const size_t total_sam_budget =
num_sams > 0 ? std::min(param_.external_sam_budget, total_draft_token_num) : size_t{0};
const size_t per_sam_budget = num_sams > 0 ? total_sam_budget / num_sams : size_t{0};
const size_t trie_budget = total_draft_token_num - (per_sam_budget * num_sams);
Result merged;
for (size_t i = 0; i < state_ids.size(); ++i) {
const auto& suffix = tokens[i];
@@ -162,12 +200,8 @@ Result Ngram::batchMatch(
}
auto& state = match_state_[state_ids[i]];
const auto total_draft_token_num = param_.get_draft_token_num(tokens.size());
const auto sam_budget =
sam_ && !sam_->empty() ? std::min(param_.external_sam_budget, total_draft_token_num) : size_t{0};
const auto trie_budget = total_draft_token_num - sam_budget;
if (sam_budget == 0) {
if (total_sam_budget == 0 || per_sam_budget == 0) {
auto res = (trie_.get()->*trie_result_build_fn)(
suffix.data(), suffix.size(), suffix.back(), total_draft_token_num, param_, state, total_lens[i]);
merged.token.insert(merged.token.end(), res.token.begin(), res.token.end());
@@ -175,12 +209,17 @@ Result Ngram::batchMatch(
continue;
}
auto trie_res = (trie_.get()->*trie_result_build_fn)(
auto combined = (trie_.get()->*trie_result_build_fn)(
suffix.data(), suffix.size(), suffix.back(), trie_budget, param_, state, total_lens[i]);
auto sam_res = (sam_.get()->*sam_result_build_fn)(suffix.data(), suffix.size(), suffix.back(), sam_budget, param_);
auto res = combineRootResults_(suffix.back(), static_cast<int>(total_draft_token_num + 1), trie_res, sam_res);
merged.token.insert(merged.token.end(), res.token.begin(), res.token.end());
merged.mask.insert(merged.mask.end(), res.mask.begin(), res.mask.end());
for (const auto& [_, sam] : sams_) {
auto sam_res =
(sam.get()->*sam_result_build_fn)(suffix.data(), suffix.size(), suffix.back(), per_sam_budget, param_);
combined = combineRootResults_(suffix.back(), static_cast<int>(total_draft_token_num + 1), combined, sam_res);
}
merged.token.insert(merged.token.end(), combined.token.begin(), combined.token.end());
merged.mask.insert(merged.mask.end(), combined.mask.begin(), combined.mask.end());
}
return merged;
}
@@ -19,12 +19,16 @@ namespace ngram {
class Ngram {
std::unique_ptr<Trie> trie_;
std::unique_ptr<SuffixAutomaton> sam_;
std::unordered_map<std::string, std::unique_ptr<SuffixAutomaton>> sams_;
// FIXME: single staging slot — only one corpus can be loaded at a time.
// To support concurrent loads, move staging into a per-load local variable.
std::unique_ptr<SuffixAutomaton> staging_sam_;
Param param_;
// NOTE: protects trie_ and pending_count_. Ensures batchMatch never reads
// trie_ while insertWorker is writing. After synchronize(), no pending
// inserts remain so mutex_ contention is effectively zero.
// NOTE: protects trie_, sams_, and pending_count_. staging_sam_ is NOT
// protected by mutex_ — it is only accessed from the corpus loading thread.
// finishExternalCorpusLoad briefly acquires mutex_ to move the completed
// SAM into sams_.
mutable std::mutex mutex_;
mutable std::condition_variable sync_cv_;
// NOTE: tracks inserts from enqueue through trie_->insert() completion,
@@ -46,10 +50,14 @@ class Ngram {
void appendExternalCorpusTokens(const std::vector<int32_t>& tokens);
void finishExternalCorpusLoad();
void finishExternalCorpusLoad(const std::string& corpus_id);
void removeExternalCorpus(const std::string& corpus_id);
void clearExternalCorpus();
std::vector<std::string> listExternalCorpora() const;
Result batchMatch(const std::vector<std::vector<int32_t>>& tokens);
Result batchMatch(
@@ -112,14 +112,28 @@ struct NgramCorpusObj : public tvm::ffi::Object {
ngram_->appendExternalCorpusTokens(tokens);
}
void finish_external_corpus_load() {
ngram_->finishExternalCorpusLoad();
void finish_external_corpus_load(const std::string& corpus_id) {
ngram_->finishExternalCorpusLoad(corpus_id);
}
void remove_external_corpus(const std::string& corpus_id) {
ngram_->removeExternalCorpus(corpus_id);
}
void clear_external_corpus() {
ngram_->clearExternalCorpus();
}
std::string list_external_corpora() {
auto ids = ngram_->listExternalCorpora();
std::string result;
for (size_t i = 0; i < ids.size(); ++i) {
if (i > 0) result += "\n";
result += ids[i];
}
return result;
}
void synchronize() {
ngram_->synchronize();
}
@@ -161,7 +175,9 @@ void register_ngram_corpus() {
.def("start_external_corpus_load", &NgramCorpusObj::start_external_corpus_load)
.def("append_external_corpus_tokens", &NgramCorpusObj::append_external_corpus_tokens)
.def("finish_external_corpus_load", &NgramCorpusObj::finish_external_corpus_load)
.def("remove_external_corpus", &NgramCorpusObj::remove_external_corpus)
.def("clear_external_corpus", &NgramCorpusObj::clear_external_corpus)
.def("list_external_corpora", &NgramCorpusObj::list_external_corpora)
.def("synchronize", &NgramCorpusObj::synchronize)
.def("reset", &NgramCorpusObj::reset);
}