From 0e924f52059e27df066b29433e88629146c24453 Mon Sep 17 00:00:00 2001 From: hanhainebula <2512674094@qq.com> Date: Sun, 23 Aug 2026 20:17:50 +0800 Subject: [PATCH] fix(reranker): support transformers v5 tokenizer preparation --- .../abc/finetune/reranker/AbsDataset.py | 10 +- .../inference/reranker/decoder_only/base.py | 19 ++- .../reranker/decoder_only/layerwise.py | 7 +- .../reranker/decoder_only/lightweight.py | 7 +- .../inference/reranker/encoder_only/base.py | 4 +- FlagEmbedding/utils/tokenizer_compat.py | 118 +++++++++++++++ examples/inference/reranker/README.md | 14 +- tests/README.md | 4 + tests/test_reranker_tokenizer_compat.py | 136 ++++++++++++++++++ 9 files changed, 301 insertions(+), 18 deletions(-) create mode 100644 FlagEmbedding/utils/tokenizer_compat.py create mode 100644 tests/test_reranker_tokenizer_compat.py diff --git a/FlagEmbedding/abc/finetune/reranker/AbsDataset.py b/FlagEmbedding/abc/finetune/reranker/AbsDataset.py index 73830bbb..0150c7e9 100644 --- a/FlagEmbedding/abc/finetune/reranker/AbsDataset.py +++ b/FlagEmbedding/abc/finetune/reranker/AbsDataset.py @@ -16,6 +16,7 @@ from typing import List from .AbsArguments import AbsRerankerDataArguments +from FlagEmbedding.utils.tokenizer_compat import prepare_for_model_compat logger = logging.getLogger(__name__) @@ -115,7 +116,8 @@ def create_one_example(self, qry_encoding: str, doc_encoding: str): """ qry_inputs = self.tokenizer.encode(qry_encoding, truncation=True, max_length=self.args.query_max_len + self.args.passage_max_len // 4, add_special_tokens=False) doc_inputs = self.tokenizer.encode(doc_encoding, truncation=True, max_length=self.args.passage_max_len + self.args.query_max_len // 2, add_special_tokens=False) - item = self.tokenizer.prepare_for_model( + item = prepare_for_model_compat( + self.tokenizer, qry_inputs, doc_inputs, truncation='only_second', @@ -302,7 +304,8 @@ def __getitem__(self, item) -> List[BatchEncoding]: add_special_tokens=False ) if self.tokenizer.bos_token_id is not None and self.tokenizer.bos_token_id != self.tokenizer.pad_token_id: - item = self.tokenizer.prepare_for_model( + item = prepare_for_model_compat( + self.tokenizer, [self.tokenizer.bos_token_id] + query_inputs['input_ids'], self.sep_inputs + passage_inputs['input_ids'], truncation='only_second', @@ -313,7 +316,8 @@ def __getitem__(self, item) -> List[BatchEncoding]: add_special_tokens=False ) else: - item = self.tokenizer.prepare_for_model( + item = prepare_for_model_compat( + self.tokenizer, query_inputs['input_ids'], self.sep_inputs + passage_inputs['input_ids'], truncation='only_second', diff --git a/FlagEmbedding/inference/reranker/decoder_only/base.py b/FlagEmbedding/inference/reranker/decoder_only/base.py index 4d5b26ec..b4aa8f50 100644 --- a/FlagEmbedding/inference/reranker/decoder_only/base.py +++ b/FlagEmbedding/inference/reranker/decoder_only/base.py @@ -10,6 +10,7 @@ from FlagEmbedding.abc.inference import AbsReranker from FlagEmbedding.inference.reranker.encoder_only.base import sigmoid +from FlagEmbedding.utils.tokenizer_compat import prepare_for_model_compat def last_logit_pool(logits: Tensor, @@ -89,7 +90,8 @@ def __getitem__(self, item): query_inputs = self.all_queries_inputs[item] passage_inputs = self.all_passages_inputs[item] if self.tokenizer.bos_token_id is not None and self.tokenizer.bos_token_id != self.tokenizer.pad_token_id: - item = self.tokenizer.prepare_for_model( + item = prepare_for_model_compat( + self.tokenizer, [self.tokenizer.bos_token_id] + query_inputs['input_ids'], self.sep_inputs + passage_inputs['input_ids'], truncation='only_second', @@ -100,7 +102,8 @@ def __getitem__(self, item): add_special_tokens=False ) else: - item = self.tokenizer.prepare_for_model( + item = prepare_for_model_compat( + self.tokenizer, query_inputs['input_ids'], self.sep_inputs + passage_inputs['input_ids'], truncation='only_second', @@ -371,7 +374,8 @@ def compute_score_single_gpu( all_passages_inputs_sorted[:min(len(all_passages_inputs_sorted), batch_size)] ): if self.tokenizer.bos_token_id is not None and self.tokenizer.bos_token_id != self.tokenizer.pad_token_id: - item = self.tokenizer.prepare_for_model( + item = prepare_for_model_compat( + self.tokenizer, [self.tokenizer.bos_token_id] + query_inputs['input_ids'], sep_inputs + passage_inputs['input_ids'], truncation='only_second', @@ -382,7 +386,8 @@ def compute_score_single_gpu( add_special_tokens=False ) else: - item = self.tokenizer.prepare_for_model( + item = prepare_for_model_compat( + self.tokenizer, query_inputs['input_ids'], sep_inputs + passage_inputs['input_ids'], truncation='only_second', @@ -452,7 +457,8 @@ def compute_score_single_gpu( batch_inputs = [] for query_inputs, passage_inputs in zip(queries_inputs, passages_inputs): if self.tokenizer.bos_token_id is not None and self.tokenizer.bos_token_id != self.tokenizer.pad_token_id: - item = self.tokenizer.prepare_for_model( + item = prepare_for_model_compat( + self.tokenizer, [self.tokenizer.bos_token_id] + query_inputs['input_ids'], sep_inputs + passage_inputs['input_ids'], truncation='only_second', @@ -463,7 +469,8 @@ def compute_score_single_gpu( add_special_tokens=False ) else: - item = self.tokenizer.prepare_for_model( + item = prepare_for_model_compat( + self.tokenizer, query_inputs['input_ids'], sep_inputs + passage_inputs['input_ids'], truncation='only_second', diff --git a/FlagEmbedding/inference/reranker/decoder_only/layerwise.py b/FlagEmbedding/inference/reranker/decoder_only/layerwise.py index 4b75da36..9f7fb98d 100644 --- a/FlagEmbedding/inference/reranker/decoder_only/layerwise.py +++ b/FlagEmbedding/inference/reranker/decoder_only/layerwise.py @@ -10,6 +10,7 @@ from FlagEmbedding.abc.inference import AbsReranker from FlagEmbedding.inference.reranker.encoder_only.base import sigmoid +from FlagEmbedding.utils.tokenizer_compat import prepare_for_model_compat from FlagEmbedding.inference.reranker.decoder_only.base import DatasetForReranker, Collater from .models.modeling_minicpm_reranker import LayerWiseMiniCPMForCausalLM @@ -252,7 +253,8 @@ def compute_score_single_gpu( all_queries_inputs_sorted[:min(len(all_queries_inputs_sorted), batch_size)], all_passages_inputs_sorted[:min(len(all_passages_inputs_sorted), batch_size)] ): - item = self.tokenizer.prepare_for_model( + item = prepare_for_model_compat( + self.tokenizer, [self.tokenizer.bos_token_id] + query_inputs['input_ids'], sep_inputs + passage_inputs['input_ids'], truncation='only_second', @@ -329,7 +331,8 @@ def compute_score_single_gpu( batch_inputs = [] for query_inputs, passage_inputs in zip(queries_inputs, passages_inputs): - item = self.tokenizer.prepare_for_model( + item = prepare_for_model_compat( + self.tokenizer, [self.tokenizer.bos_token_id] + query_inputs['input_ids'], sep_inputs + passage_inputs['input_ids'], truncation='only_second', diff --git a/FlagEmbedding/inference/reranker/decoder_only/lightweight.py b/FlagEmbedding/inference/reranker/decoder_only/lightweight.py index 000478af..44422c52 100644 --- a/FlagEmbedding/inference/reranker/decoder_only/lightweight.py +++ b/FlagEmbedding/inference/reranker/decoder_only/lightweight.py @@ -10,6 +10,7 @@ from FlagEmbedding.abc.inference import AbsReranker from FlagEmbedding.inference.reranker.encoder_only.base import sigmoid +from FlagEmbedding.utils.tokenizer_compat import prepare_for_model_compat def last_logit_pool_lightweight(logits: Tensor, @@ -333,7 +334,8 @@ def compute_score_single_gpu( all_queries_inputs_sorted[:min(len(all_queries_inputs_sorted), batch_size)], all_passages_inputs_sorted[:min(len(all_passages_inputs_sorted), batch_size)] ): - item = self.tokenizer.prepare_for_model( + item = prepare_for_model_compat( + self.tokenizer, [self.tokenizer.bos_token_id] + query_inputs['input_ids'], sep_inputs + passage_inputs['input_ids'], truncation='only_second', @@ -388,7 +390,8 @@ def compute_score_single_gpu( query_lengths = [] prompt_lengths = [] for query_inputs, passage_inputs in zip(queries_inputs, passages_inputs): - item = self.tokenizer.prepare_for_model( + item = prepare_for_model_compat( + self.tokenizer, [self.tokenizer.bos_token_id] + query_inputs['input_ids'], sep_inputs + passage_inputs['input_ids'], truncation='only_second', diff --git a/FlagEmbedding/inference/reranker/encoder_only/base.py b/FlagEmbedding/inference/reranker/encoder_only/base.py index 1a4d8b6a..15686dcc 100644 --- a/FlagEmbedding/inference/reranker/encoder_only/base.py +++ b/FlagEmbedding/inference/reranker/encoder_only/base.py @@ -5,6 +5,7 @@ from transformers import AutoModelForSequenceClassification, AutoTokenizer from FlagEmbedding.abc.inference import AbsReranker +from FlagEmbedding.utils.tokenizer_compat import prepare_for_model_compat def sigmoid(x): @@ -144,7 +145,8 @@ def compute_score_single_gpu( **kwargs )['input_ids'] for q_inp, d_inp in zip(queries_inputs_batch, passages_inputs_batch): - item = self.tokenizer.prepare_for_model( + item = prepare_for_model_compat( + self.tokenizer, q_inp, d_inp, truncation='only_second', diff --git a/FlagEmbedding/utils/tokenizer_compat.py b/FlagEmbedding/utils/tokenizer_compat.py new file mode 100644 index 00000000..955c4308 --- /dev/null +++ b/FlagEmbedding/utils/tokenizer_compat.py @@ -0,0 +1,118 @@ +"""Tokenizer compatibility helpers for supported Transformers versions.""" + +from typing import Any, List, Optional, Sequence + + +def _decode_token_ids(tokenizer: Any, token_ids: Sequence[int]) -> str: + """Decode token ids without dropping unknown or other special tokens.""" + return tokenizer.decode( + list(token_ids), + skip_special_tokens=False, + clean_up_tokenization_spaces=False, + ) + + +def _truncate_second_sequence( + tokenizer: Any, + first_ids: Sequence[int], + second_ids: Sequence[int], + max_length: Optional[int], + truncation: Any, +) -> List[int]: + """Apply the old ``only_second`` truncation behavior to token ids.""" + second_ids = list(second_ids) + truncation = getattr(truncation, "value", truncation) + if max_length is None or truncation not in (True, "only_second"): + return second_ids + + tokens_to_remove = len(first_ids) + len(second_ids) - max_length + if tokens_to_remove <= 0 or len(second_ids) <= tokens_to_remove: + # This matches the legacy tokenizer behavior: only the second sequence + # may be truncated, so an oversized first sequence is left untouched. + return second_ids + + if getattr(tokenizer, "truncation_side", "right") == "left": + return second_ids[tokens_to_remove:] + return second_ids[:-tokens_to_remove] + + +def prepare_for_model_compat( + tokenizer: Any, + first_ids: Sequence[int], + second_ids: Optional[Sequence[int]] = None, + *, + truncation: Any = None, + max_length: Optional[int] = None, + padding: Any = False, + return_attention_mask: Optional[bool] = None, + return_token_type_ids: Optional[bool] = None, + add_special_tokens: bool = True, + **kwargs: Any, +) -> dict: + """Prepare a pair of tokenized sequences across Transformers v4 and v5. + + Transformers v5 removed the id-level ``prepare_for_model`` API from + tokenizers. For encoder-only rerankers, the supported replacement is the + tokenizer call with a text pair. The ids are decoded with special tokens + preserved so ``unk_token_id`` and model-specific special tokens are not + silently lost before the pair is tokenized again. + + Decoder-only rerankers use this helper with ``add_special_tokens=False``. + In that mode no text round-trip is needed; the second sequence is truncated + directly while respecting ``tokenizer.truncation_side``. + """ + has_pair = second_ids is not None + if second_ids is None: + second_ids = [] + + legacy_prepare_for_model = getattr(tokenizer, "prepare_for_model", None) + if callable(legacy_prepare_for_model): + return legacy_prepare_for_model( + list(first_ids), + list(second_ids) if has_pair else None, + truncation=truncation, + max_length=max_length, + padding=padding, + return_attention_mask=return_attention_mask, + return_token_type_ids=return_token_type_ids, + add_special_tokens=add_special_tokens, + **kwargs, + ) + + if add_special_tokens: + tokenizer_kwargs = dict( + truncation=truncation, + max_length=max_length, + padding=padding, + return_attention_mask=return_attention_mask, + return_token_type_ids=return_token_type_ids, + add_special_tokens=True, + ) + tokenizer_kwargs.update(kwargs) + return tokenizer( + _decode_token_ids(tokenizer, first_ids), + _decode_token_ids(tokenizer, second_ids) if has_pair else None, + **tokenizer_kwargs, + ) + + first_ids = list(first_ids) + second_ids = _truncate_second_sequence( + tokenizer, + first_ids, + second_ids, + max_length, + truncation, + ) + input_ids = first_ids + second_ids + result = {"input_ids": input_ids} + + if return_attention_mask is None: + return_attention_mask = "attention_mask" in getattr(tokenizer, "model_input_names", []) + if return_token_type_ids is None: + return_token_type_ids = "token_type_ids" in getattr(tokenizer, "model_input_names", []) + if return_attention_mask: + result["attention_mask"] = [1] * len(input_ids) + if return_token_type_ids: + result["token_type_ids"] = [0] * len(first_ids) + [0] * len(second_ids) + + return result diff --git a/examples/inference/reranker/README.md b/examples/inference/reranker/README.md index 19fb63a1..9dc836f3 100644 --- a/examples/inference/reranker/README.md +++ b/examples/inference/reranker/README.md @@ -199,6 +199,7 @@ It supports `BAAI/bge-reranker-v2-gemma`: ```python import torch from transformers import AutoModelForCausalLM, AutoTokenizer +from FlagEmbedding.utils.tokenizer_compat import prepare_for_model_compat def get_inputs(pairs, tokenizer, prompt=None, max_length=1024): if prompt is None: @@ -222,7 +223,8 @@ def get_inputs(pairs, tokenizer, prompt=None, max_length=1024): add_special_tokens=False, max_length=max_length, truncation=True) - item = tokenizer.prepare_for_model( + item = prepare_for_model_compat( + tokenizer, [tokenizer.bos_token_id] + query_inputs['input_ids'], sep_inputs + passage_inputs['input_ids'], truncation='only_second', @@ -262,6 +264,7 @@ It supports `BAAI/bge-reranker-v2-minicpm-layerwise`: ```python import torch from transformers import AutoModelForCausalLM, AutoTokenizer +from FlagEmbedding.utils.tokenizer_compat import prepare_for_model_compat def get_inputs(pairs, tokenizer, prompt=None, max_length=1024): if prompt is None: @@ -285,7 +288,8 @@ def get_inputs(pairs, tokenizer, prompt=None, max_length=1024): add_special_tokens=False, max_length=max_length, truncation=True) - item = tokenizer.prepare_for_model( + item = prepare_for_model_compat( + tokenizer, [tokenizer.bos_token_id] + query_inputs['input_ids'], sep_inputs + passage_inputs['input_ids'], truncation='only_second', @@ -326,6 +330,7 @@ It supports `BAAI/bge-reranker-v2.5-gemma2-lightweight`: ```python import torch from transformers import AutoModelForCausalLM, AutoTokenizer +from FlagEmbedding.utils.tokenizer_compat import prepare_for_model_compat def last_logit_pool(logits: torch.Tensor, attention_mask: torch.Tensor) -> torch.Tensor: @@ -361,7 +366,8 @@ def get_inputs(pairs, tokenizer, prompt=None, max_length=1024): add_special_tokens=False, max_length=max_length, truncation=True) - item = tokenizer.prepare_for_model( + item = prepare_for_model_compat( + tokenizer, [tokenizer.bos_token_id] + query_inputs['input_ids'], sep_inputs + passage_inputs['input_ids'], truncation='only_second', @@ -469,4 +475,4 @@ If you find this repository useful, please consider giving a star :star: and cit primaryClass={cs.IR}, url={https://arxiv.org/abs/2409.15700}, } -``` \ No newline at end of file +``` diff --git a/tests/README.md b/tests/README.md index 79939ae9..fcf55e37 100644 --- a/tests/README.md +++ b/tests/README.md @@ -6,6 +6,7 @@ This directory contains tests for the FlagEmbedding library, including compatibi - `test_imports_v5.py`: Tests that imports work with Transformers v5, particularly the compatibility layer for `is_torch_fx_available`. - `test_finetune_trainer_compat.py`: Tests the Transformers Trainer API migration, including `processing_class`/legacy `tokenizer` construction, reranker runner arguments, and processor checkpoint saving. +- `test_reranker_tokenizer_compat.py`: Tests reranker pair preparation across Transformers v4/v5, including the v5 removal of `prepare_for_model`. - `test_infer_embedder_basic.py`: Tests basic functionality of BGE embedder models with a small public checkpoint. - `test_infer_reranker_basic.py`: Tests basic functionality of reranker models. @@ -28,6 +29,9 @@ pytest tests/test_imports_v5.py # Run the fine-tuning Trainer compatibility tests pytest tests/test_finetune_trainer_compat.py +# Run the reranker tokenizer compatibility tests +pytest tests/test_reranker_tokenizer_compat.py + # Run with verbose output pytest -v tests/ ``` diff --git a/tests/test_reranker_tokenizer_compat.py b/tests/test_reranker_tokenizer_compat.py new file mode 100644 index 00000000..41f96819 --- /dev/null +++ b/tests/test_reranker_tokenizer_compat.py @@ -0,0 +1,136 @@ +import os +from types import SimpleNamespace + +import pytest +from transformers import AutoTokenizer + +from FlagEmbedding.abc.finetune.reranker.AbsDataset import AbsRerankerTrainDataset +from FlagEmbedding.utils.tokenizer_compat import prepare_for_model_compat + + +RERANKER_MODEL_PATH = os.environ.get( + "FLAGEMBEDDING_RERANKER_MODEL", + "/share/project/Search-SWE/model/bge-reranker-base", +) + + +class LegacyTokenizer: + model_input_names = ["input_ids", "token_type_ids", "attention_mask"] + truncation_side = "right" + + def __init__(self): + self.calls = [] + + def prepare_for_model(self, first_ids, second_ids, **kwargs): + self.calls.append((first_ids, second_ids, kwargs)) + return {"input_ids": list(first_ids) + list(second_ids)} + + +@pytest.mark.skipif( + not os.path.isdir(RERANKER_MODEL_PATH), + reason="local reranker model is not available", +) +def test_encoder_pair_preserves_v5_tokenizer_input_ids(): + tokenizer = AutoTokenizer.from_pretrained( + RERANKER_MODEL_PATH, + local_files_only=True, + ) + + query = "๐Ÿ˜€๐Ÿงช" + passage = "unknown token" + query_ids = tokenizer(query, add_special_tokens=False)["input_ids"] + passage_ids = tokenizer(passage, add_special_tokens=False)["input_ids"] + + actual = prepare_for_model_compat( + tokenizer, + query_ids, + passage_ids, + truncation="only_second", + max_length=64, + padding=False, + ) + expected = tokenizer( + query, + passage, + truncation="only_second", + max_length=64, + padding=False, + ) + + assert not hasattr(tokenizer, "prepare_for_model") + assert actual["input_ids"] == expected["input_ids"] + assert actual["attention_mask"] == expected["attention_mask"] + + +@pytest.mark.skipif( + not os.path.isdir(RERANKER_MODEL_PATH), + reason="local reranker model is not available", +) +def test_encoder_training_example_works_without_prepare_for_model(): + tokenizer = AutoTokenizer.from_pretrained( + RERANKER_MODEL_PATH, + local_files_only=True, + ) + dataset = AbsRerankerTrainDataset.__new__(AbsRerankerTrainDataset) + dataset.tokenizer = tokenizer + dataset.args = SimpleNamespace(query_max_len=32, passage_max_len=128) + + actual = dataset.create_one_example("What is AI?", "AI is artificial intelligence.") + expected = tokenizer( + "What is AI?", + "AI is artificial intelligence.", + truncation="only_second", + max_length=160, + padding=False, + ) + + assert actual["input_ids"] == expected["input_ids"] + assert actual["attention_mask"] == expected["attention_mask"] + + +def test_legacy_tokenizer_path_is_preserved(): + tokenizer = LegacyTokenizer() + actual = prepare_for_model_compat( + tokenizer, + [1, 2], + [3, 4], + truncation="only_second", + max_length=8, + padding=False, + ) + + assert actual["input_ids"] == [1, 2, 3, 4] + assert len(tokenizer.calls) == 1 + assert tokenizer.calls[0][2]["add_special_tokens"] is True + + +def test_decoder_only_truncation_respects_tokenizer_side(): + tokenizer = LegacyTokenizer() + tokenizer.prepare_for_model = None + + right = prepare_for_model_compat( + tokenizer, + [1, 2], + [3, 4, 5], + truncation="only_second", + max_length=4, + padding=False, + return_attention_mask=False, + return_token_type_ids=False, + add_special_tokens=False, + ) + assert right["input_ids"] == [1, 2, 3, 4] + + tokenizer.truncation_side = "left" + left = prepare_for_model_compat( + tokenizer, + [1, 2], + [3, 4, 5], + truncation="only_second", + max_length=4, + padding=False, + return_attention_mask=False, + return_token_type_ids=False, + add_special_tokens=False, + ) + assert left["input_ids"] == [1, 2, 4, 5]