[Spec][Ngram] Clean up unused stateless batchMatch (#22487)
This commit is contained in:
@@ -139,35 +139,6 @@ void Ngram::insertWorker() {
|
||||
}
|
||||
}
|
||||
|
||||
Result Ngram::batchMatch(const std::vector<std::vector<int32_t>>& tokens) {
|
||||
std::unique_lock<std::mutex> 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<int64_t>& state_ids,
|
||||
const std::vector<std::vector<int32_t>>& tokens,
|
||||
|
||||
@@ -61,8 +61,6 @@ class Ngram {
|
||||
|
||||
std::vector<std::pair<std::string, int64_t>> listExternalCorpora() const;
|
||||
|
||||
Result batchMatch(const std::vector<std::vector<int32_t>>& tokens);
|
||||
|
||||
Result batchMatch(
|
||||
const std::vector<int64_t>& state_ids,
|
||||
const std::vector<std::vector<int32_t>>& tokens,
|
||||
|
||||
@@ -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<const int32_t*>(tokens_flat.data_ptr());
|
||||
auto* offs = static_cast<const int64_t*>(offsets.data_ptr());
|
||||
int64_t batch_size = offsets.size(0) - 1;
|
||||
|
||||
std::vector<std::vector<int32_t>> 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<NgramCorpusObj>()
|
||||
.def(refl::init<int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t, int64_t>(), "__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)
|
||||
|
||||
@@ -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],
|
||||
|
||||
Reference in New Issue
Block a user