[Spec][Ngram] Support multiple SAMs with dynamic HTTP API (#22203)
This commit is contained in:
@@ -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);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user