Skip to content
Merged
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
10 changes: 7 additions & 3 deletions FlagEmbedding/abc/finetune/reranker/AbsDataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)

Expand Down Expand Up @@ -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',
Expand Down Expand Up @@ -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',
Expand All @@ -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',
Expand Down
19 changes: 13 additions & 6 deletions FlagEmbedding/inference/reranker/decoder_only/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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',
Expand All @@ -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',
Expand Down Expand Up @@ -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',
Expand All @@ -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',
Expand Down Expand Up @@ -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',
Expand All @@ -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',
Expand Down
7 changes: 5 additions & 2 deletions FlagEmbedding/inference/reranker/decoder_only/layerwise.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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',
Expand Down Expand Up @@ -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',
Expand Down
7 changes: 5 additions & 2 deletions FlagEmbedding/inference/reranker/decoder_only/lightweight.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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',
Expand Down Expand Up @@ -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',
Expand Down
4 changes: 3 additions & 1 deletion FlagEmbedding/inference/reranker/encoder_only/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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',
Expand Down
118 changes: 118 additions & 0 deletions FlagEmbedding/utils/tokenizer_compat.py
Original file line number Diff line number Diff line change
@@ -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
14 changes: 10 additions & 4 deletions examples/inference/reranker/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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',
Expand Down Expand Up @@ -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:
Expand All @@ -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',
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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',
Expand Down Expand Up @@ -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},
}
```
```
4 changes: 4 additions & 0 deletions tests/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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/
```
Expand Down
Loading
Loading