From 0af5bfcd386e2efef1864f730b2c685a7dcdef85 Mon Sep 17 00:00:00 2001 From: tuanzirwar <1281227988@qq.com> Date: Thu, 1 Oct 2026 16:42:05 +0800 Subject: [PATCH] fix: isolate tokenizer truncation for concurrent encode calls --- model2vec/model.py | 53 +++++++++++------- tests/test_concurrent_encoding.py | 91 +++++++++++++++++++++++++++++++ 2 files changed, 124 insertions(+), 20 deletions(-) create mode 100644 tests/test_concurrent_encoding.py diff --git a/model2vec/model.py b/model2vec/model.py index 4e56589..c0b349c 100644 --- a/model2vec/model.py +++ b/model2vec/model.py @@ -173,7 +173,11 @@ def tokenize(self, sentences: Sequence[str]) -> list[list[int]]: :param sentences: The sentences to tokenize. :return: A list of list of tokens. """ - encodings: list[Encoding] = self.tokenizer.encode_batch_fast(sentences, add_special_tokens=False) + return self._tokenize(sentences, self.tokenizer) + + def _tokenize(self, sentences: Sequence[str], tokenizer: Tokenizer) -> list[list[int]]: + """使用本次编码的 tokenizer,避免请求间修改共享截断配置.""" + encodings: list[Encoding] = tokenizer.encode_batch_fast(sentences, add_special_tokens=False) encodings_ids = [encoding.ids for encoding in encodings] @@ -275,22 +279,27 @@ def _encode_dispatch( sentence_batches = list(self._batch(sentences, batch_size)) total_batches = math.ceil(len(sentences) / batch_size) - self._set_max_length_in_tokenizer(max_length) - try: - if use_multiprocessing and len(sentences) > multiprocessing_threshold: - # Disable parallelism for tokenizers - os.environ["TOKENIZERS_PARALLELISM"] = "false" - - results = ProgressParallel( - n_jobs=-1, backend="threading", use_tqdm=show_progress_bar, total=total_batches - )(delayed(batch_fn)(batch, *batch_args) for batch in sentence_batches) + tokenizer = self.tokenizer + if max_length != self.max_length: + # 默认长度复用只读 tokenizer,仅显式覆盖时复制一次供所有批次使用。 + tokenizer = copy.deepcopy(tokenizer) + if max_length is None: + tokenizer.no_truncation() else: - results = [ - batch_fn(batch, *batch_args) - for batch in tqdm(sentence_batches, total=total_batches, disable=not show_progress_bar) - ] - finally: - self._set_max_length_in_tokenizer(self.max_length) + tokenizer.enable_truncation(max_length) + + if use_multiprocessing and len(sentences) > multiprocessing_threshold: + # 禁用 tokenizer 内部并行,沿用外部线程池的批处理方式。 + os.environ["TOKENIZERS_PARALLELISM"] = "false" + + results = ProgressParallel(n_jobs=-1, backend="threading", use_tqdm=show_progress_bar, total=total_batches)( + delayed(batch_fn)(batch, *batch_args, tokenizer=tokenizer) for batch in sentence_batches + ) + else: + results = [ + batch_fn(batch, *batch_args, tokenizer=tokenizer) + for batch in tqdm(sentence_batches, total=total_batches, disable=not show_progress_bar) + ] return results, was_single @@ -368,9 +377,11 @@ def encode_as_sequence( return out_array[0] return out_array - def _encode_batch_as_sequence(self, sentences: Sequence[str]) -> list[np.ndarray]: + def _encode_batch_as_sequence( + self, sentences: Sequence[str], tokenizer: Tokenizer | None = None + ) -> list[np.ndarray]: """Encode a batch of sentences as a sequence.""" - ids = self.tokenize(sentences=sentences) + ids = self._tokenize(sentences, self.tokenizer if tokenizer is None else tokenizer) out: list[np.ndarray] = [] for id_list in ids: if id_list: @@ -454,9 +465,11 @@ def _encode_helper(self, id_list: list[int]) -> np.ndarray: return emb - def _encode_batch(self, sentences: Sequence[str], normalize: bool) -> np.ndarray: + def _encode_batch( + self, sentences: Sequence[str], normalize: bool, tokenizer: Tokenizer | None = None + ) -> np.ndarray: """Encode a batch of sentences.""" - ids = self.tokenize(sentences=sentences) + ids = self._tokenize(sentences, self.tokenizer if tokenizer is None else tokenizer) dtype = self.embedding.dtype if dtype == np.int8: dtype = np.float32 diff --git a/tests/test_concurrent_encoding.py b/tests/test_concurrent_encoding.py new file mode 100644 index 0000000..0e5c700 --- /dev/null +++ b/tests/test_concurrent_encoding.py @@ -0,0 +1,91 @@ +from concurrent.futures import ThreadPoolExecutor +from threading import Event +from typing import Any + +import numpy as np +import pytest +from tokenizers import Tokenizer, models, pre_tokenizers + +from model2vec import StaticModel + + +@pytest.fixture +def concurrent_model() -> StaticModel: + """建立无需下载的真实分词器和静态向量.""" + words = ["[UNK]", "a", "b", "c", "longwordnumberone", "longwordnumbertwo", "longwordnumberthree"] + tokenizer = Tokenizer(models.WordLevel({word: i for i, word in enumerate(words)}, unk_token="[UNK]")) + tokenizer.pre_tokenizer = pre_tokenizers.Whitespace() + vectors = np.arange(len(words) * 2, dtype=np.float32).reshape(len(words), 2) + return StaticModel(vectors, tokenizer, max_length=2) + + +@pytest.mark.parametrize("short_method", ["encode", "encode_as_sequence"]) +@pytest.mark.parametrize("long_method", ["encode", "encode_as_sequence"]) +@pytest.mark.parametrize("long_limit", [3, None]) +@pytest.mark.parametrize("parallel_batches", [False, True]) +def test_concurrent_length_overrides( + concurrent_model: StaticModel, + monkeypatch: pytest.MonkeyPatch, + short_method: str, + long_method: str, + long_limit: int | None, + parallel_batches: bool, +) -> None: + """不同长度及两种编码路径的重叠请求必须与顺序结果一致.""" + short_encode = getattr(concurrent_model, short_method) + long_encode = getattr(concurrent_model, long_method) + batch_options = {"batch_size": 1, "use_multiprocessing": parallel_batches, "multiprocessing_threshold": 0} + expected_short = short_encode(["a b c"] * 2, max_length=1, **batch_options) + expected_long = long_encode(["c b a"] * 2, max_length=long_limit, **batch_options) + entered_short = Event() + finished_long = Event() + original_tokenize = concurrent_model._tokenize + + def scheduled_tokenize(sentences: list[str], tokenizer: Tokenizer) -> list[list[int]]: + if sentences == ["a b c"]: + entered_short.set() + assert finished_long.wait(10), "长请求未完成" + return original_tokenize(sentences, tokenizer) + + monkeypatch.setattr(concurrent_model, "_tokenize", scheduled_tokenize) + with ThreadPoolExecutor(max_workers=1) as executor: + pending = executor.submit(short_encode, ["a b c"] * 2, max_length=1, **batch_options) + try: + assert entered_short.wait(10), "短请求未开始" + actual_long = long_encode(["c b a"] * 2, max_length=long_limit, **batch_options) + finally: + finished_long.set() + actual_short = pending.result(timeout=10) + for expected, actual in zip(expected_short, actual_short): + np.testing.assert_array_equal(actual, expected) + for expected, actual in zip(expected_long, actual_long): + np.testing.assert_array_equal(actual, expected) + assert concurrent_model.tokenizer.truncation["max_length"] == 2 + + +def test_default_encoding_reuses_tokenizer(concurrent_model: StaticModel, monkeypatch: pytest.MonkeyPatch) -> None: + """默认路径不复制 tokenizer,避免给常规推理增加复制开销.""" + observed = [] + original_tokenize = concurrent_model._tokenize + + def observe(sentences: list[str], tokenizer: Tokenizer) -> list[list[int]]: + observed.append(tokenizer) + return original_tokenize(sentences, tokenizer) + + monkeypatch.setattr(concurrent_model, "_tokenize", observe) + concurrent_model.encode(["a b c"], use_multiprocessing=False) + assert observed == [concurrent_model.tokenizer] + + +def test_encoding_error_keeps_default_truncation( + concurrent_model: StaticModel, monkeypatch: pytest.MonkeyPatch +) -> None: + """覆盖长度的编码异常也不能改变模型默认配置.""" + + def fail(*args: Any, **kwargs: Any) -> list[list[int]]: + raise ValueError("分词失败") + + monkeypatch.setattr(concurrent_model, "_tokenize", fail) + with pytest.raises(ValueError, match="分词失败"): + concurrent_model.encode(["a b c"], max_length=None, use_multiprocessing=False) + assert concurrent_model.tokenizer.truncation["max_length"] == 2