-
Notifications
You must be signed in to change notification settings - Fork 128
fix: isolate tokenizer truncation for concurrent encode calls #385
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Closed
tuanzirwar
wants to merge
1
commit into
MinishLab:main
from
tuanzirwar:codex/isolate-concurrent-truncation
+124
−20
Closed
Changes from all commits
Commits
File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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) | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. |
||
| 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 | ||
|
|
||
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -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 |
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
encode_as_sequence()defaults tomax_length=None, which differs from a model's usual non-Nonedefault. As a result, every such call deep-copies the tokenizer, even for a small request, adding repeated copying cost to sequence encoding.Note: If this suggestion doesn't match your team's coding style, reply to this and let me know. I'll remember it for next time!