Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 33 additions & 20 deletions model2vec/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]

Expand Down Expand Up @@ -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)
Comment on lines +283 to +285

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Default sequence calls copy tokenizer encode_as_sequence() defaults to max_length=None, which differs from a model's usual non-None default. 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!

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

Expand Down Expand Up @@ -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)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Tokenize overrides are bypassed If a StaticModel subclass overrides the public tokenize() method to customize token IDs, encode_as_sequence() now calls _tokenize() instead. encode() makes the same change, so both methods silently ignore the override and can return different embeddings than before.

out: list[np.ndarray] = []
for id_list in ids:
if id_list:
Expand Down Expand Up @@ -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
Expand Down
91 changes: 91 additions & 0 deletions tests/test_concurrent_encoding.py
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