Skip to content

Commit 465419f

Browse files
Make cross-encoder abstract (#76)
* fix: enhance documentation for SetEncoder config, model and tokenizer * fix: enhance documentation for T5 cross-encoder configuration, model and tokenizer * fix: update documentation in SetEncoderConfig for clarity on document embeddings * feat: implement MonoModel for mono cross-encoder and update T5CrossEncoderModel to inherit from MonoModel * fix pretty printing results for run files * flake8 + black + add pylate dependency * fix flake8 * iosrt * fix dependencies in test ci * add option to pass model kwargs to modules * fix pretty printing for skipped inference datasets * update config with model kwargs when necessary * remove quotes when writing run * register monoelectra rank-distillm data * remove flash * make model registration public * update seismic test index * remove test indexes * Add MonoConfig and MonoModel classes; update model registration and tests * finish making cross-encoder abstract --------- Co-authored-by: Rayk Kretzschmar <112646288+RaykKretzschmar@users.noreply.github.com>
1 parent 35ad129 commit 465419f

11 files changed

Lines changed: 319 additions & 189 deletions

lightning_ir/cross_encoder/cross_encoder_model.py

Lines changed: 4 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
This module defines the model class used to implement cross-encoder models.
55
"""
66

7+
from abc import ABC, abstractmethod
78
from dataclasses import dataclass
89
from typing import Type
910

@@ -23,7 +24,7 @@ class CrossEncoderOutput(LightningIROutput):
2324
"""Joint query-document embeddings"""
2425

2526

26-
class CrossEncoderModel(LightningIRModel):
27+
class CrossEncoderModel(LightningIRModel, ABC):
2728
config_class: Type[CrossEncoderConfig] = CrossEncoderConfig
2829
"""Configuration class for cross-encoder models."""
2930

@@ -36,9 +37,9 @@ def __init__(self, config: CrossEncoderConfig, *args, **kwargs):
3637
"""
3738
super().__init__(config, *args, **kwargs)
3839
self.config: CrossEncoderConfig
39-
self.linear = torch.nn.Linear(config.hidden_size, 1, bias=config.linear_bias)
4040

4141
@batch_encoding_wrapper
42+
@abstractmethod
4243
def forward(self, encoding: BatchEncoding) -> CrossEncoderOutput:
4344
"""Computes contextualized embeddings for the joint query-document input sequence and computes a relevance
4445
score.
@@ -48,9 +49,4 @@ def forward(self, encoding: BatchEncoding) -> CrossEncoderOutput:
4849
:return: Output of the model
4950
:rtype: CrossEncoderOutput
5051
"""
51-
embeddings = self._backbone_forward(**encoding).last_hidden_state
52-
embeddings = self.pooling(
53-
embeddings, encoding.get("attention_mask", None), pooling_strategy=self.config.pooling_strategy
54-
)
55-
scores = self.linear(embeddings).view(-1)
56-
return CrossEncoderOutput(scores=scores, embeddings=embeddings)
52+
pass

lightning_ir/cross_encoder/cross_encoder_tokenizer.py

Lines changed: 12 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,9 @@ class CrossEncoderTokenizer(LightningIRTokenizer):
1717
config_class: Type[CrossEncoderConfig] = CrossEncoderConfig
1818
"""Configuration class for the tokenizer."""
1919

20-
def __init__(self, *args, query_length: int = 32, doc_length: int = 512, **kwargs):
20+
def __init__(
21+
self, *args, query_length: int = 32, doc_length: int = 512, tokenizer_pattern: str | None = None, **kwargs
22+
):
2123
""":class:`.LightningIRTokenizer` for cross-encoder models. Encodes queries and documents jointly and ensures
2224
that the input sequences are of the correct length.
2325
@@ -27,7 +29,10 @@ def __init__(self, *args, query_length: int = 32, doc_length: int = 512, **kwarg
2729
:type doc_length: int, optional
2830
:type doc_length: int, optional
2931
"""
30-
super().__init__(*args, query_length=query_length, doc_length=doc_length, **kwargs)
32+
super().__init__(
33+
*args, query_length=query_length, doc_length=doc_length, tokenizer_pattern=tokenizer_pattern, **kwargs
34+
)
35+
self.tokenizer_pattern = tokenizer_pattern
3136

3237
def _truncate(self, text: Sequence[str], max_length: int) -> List[str]:
3338
"""Encodes a list of texts, truncates them to a maximum number of tokens and decodes them to strings."""
@@ -98,22 +103,17 @@ def tokenize(
98103
raise ValueError("Both queries and docs must be provided.")
99104
if isinstance(docs, str) and not isinstance(queries, str):
100105
raise ValueError("Queries and docs must be both lists or both strings.")
101-
is_string_queries = False
102-
is_string_docs = False
103106
if isinstance(queries, str):
104107
queries = [queries]
105-
is_string_queries = True
106108
if isinstance(docs, str):
107109
docs = [docs]
108-
is_string_docs = True
109-
is_string_both = is_string_queries and is_string_docs
110110
num_docs = self._process_num_docs(queries, docs, num_docs)
111111
queries, docs = self._preprocess(queries, docs, num_docs)
112-
return_tensors = kwargs.get("return_tensors", None)
113-
if return_tensors is not None:
114-
kwargs["pad_to_multiple_of"] = 8
115-
if is_string_both:
116-
encoding = self(queries[0], docs[0], **kwargs)
112+
113+
if self.tokenizer_pattern is not None:
114+
input_texts = [self.tokenizer_pattern.format(query=query, doc=doc) for query, doc in zip(queries, docs)]
115+
encoding = self(input_texts, **kwargs)
117116
else:
118117
encoding = self(queries, docs, **kwargs)
118+
119119
return {"encoding": encoding}

lightning_ir/models/__init__.py

Lines changed: 3 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,21 +1,20 @@
11
from .col import ColConfig, ColModel, ColTokenizer
22
from .dpr import DprConfig, DprModel
3+
from .mono import MonoConfig, MonoModel
34
from .set_encoder import SetEncoderConfig, SetEncoderModel, SetEncoderTokenizer
45
from .splade import SpladeConfig, SpladeModel
5-
from .t5_cross_encoder import T5CrossEncoderConfig, T5CrossEncoderModel, T5CrossEncoderTokenizer
66

77
__all__ = [
88
"ColConfig",
99
"ColModel",
1010
"ColTokenizer",
1111
"DprConfig",
1212
"DprModel",
13+
"MonoConfig",
14+
"MonoModel",
1315
"SetEncoderConfig",
1416
"SetEncoderModel",
1517
"SetEncoderTokenizer",
1618
"SpladeConfig",
1719
"SpladeModel",
18-
"T5CrossEncoderConfig",
19-
"T5CrossEncoderModel",
20-
"T5CrossEncoderTokenizer",
2120
]

lightning_ir/models/mono.py

Lines changed: 117 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,117 @@
1+
"""
2+
Model implementation for mono cross-encoder models. Originally introduced in
3+
`Passage Re-ranking with BERT
4+
<https://arxiv.org/abs/1901.04085>`_.
5+
"""
6+
7+
from typing import Literal, Type
8+
9+
import torch
10+
from transformers import BatchEncoding
11+
12+
from ..base.model import batch_encoding_wrapper
13+
from ..cross_encoder.cross_encoder_config import CrossEncoderConfig
14+
from ..cross_encoder.cross_encoder_model import CrossEncoderModel, CrossEncoderOutput
15+
16+
17+
class ScaleLinear(torch.nn.Linear):
18+
19+
def forward(self, input: torch.Tensor) -> torch.Tensor:
20+
# See https://github.com/tensorflow/mesh/blob/fa19d69eafc9a482aff0b59ddd96b025c0cb207d/mesh_tensorflow/transformer/transformer.py#L586 # noqa
21+
input = input * (input.shape[-1] ** -0.5)
22+
return super().forward(input)
23+
24+
25+
class MonoConfig(CrossEncoderConfig):
26+
"""Configuration class for mono cross-encoder models."""
27+
28+
model_type = "mono"
29+
"""Model type for mono cross-encoder models."""
30+
31+
def __init__(
32+
self,
33+
query_length: int = 32,
34+
doc_length: int = 512,
35+
pooling_strategy: Literal["first", "mean", "max", "sum", "bert_pool"] = "first",
36+
linear_bias: bool = False,
37+
scoring_strategy: Literal["mono", "rank"] = "rank",
38+
tokenizer_pattern: str | None = None,
39+
**kwargs,
40+
):
41+
"""Initialize the configuration for mono cross-encoder models."""
42+
self._bert_pool = False
43+
if pooling_strategy == "bert_pool":
44+
self._bert_pool = True
45+
pooling_strategy = "first"
46+
super().__init__(
47+
query_length=query_length,
48+
doc_length=doc_length,
49+
pooling_strategy=pooling_strategy,
50+
linear_bias=linear_bias,
51+
**kwargs,
52+
)
53+
self.scoring_strategy = scoring_strategy
54+
self.tokenizer_pattern = tokenizer_pattern
55+
56+
57+
class MonoModel(CrossEncoderModel):
58+
config_class: Type[MonoConfig] = MonoConfig
59+
"""Configuration class for mono cross-encoder models."""
60+
61+
def __init__(self, config: MonoConfig, *args, **kwargs):
62+
"""A cross-encoder model that jointly encodes a query and document(s). The contextualized embeddings are
63+
aggragated into a single vector and fed to a linear layer which computes a final relevance score.
64+
65+
:param config: Configuration for the cross-encoder model
66+
:type config: CrossEncoderConfig
67+
"""
68+
super().__init__(config, *args, **kwargs)
69+
70+
if self.config.scoring_strategy == "mono":
71+
output_dim = 2
72+
elif self.config.scoring_strategy == "rank":
73+
output_dim = 1
74+
else:
75+
raise ValueError(
76+
f"Unknown scoring strategy {self.config.scoring_strategy}. Supported strategies are 'mono' and 'rank'."
77+
)
78+
79+
self.bert_pool = torch.nn.Identity()
80+
if self.config._bert_pool:
81+
self.bert_pool = torch.nn.Sequential(
82+
torch.nn.Linear(config.hidden_size, config.hidden_size), torch.nn.Tanh()
83+
)
84+
85+
if self.config.backbone_model_type == "t5":
86+
self.linear = ScaleLinear(config.hidden_size, output_dim, bias=self.config.linear_bias)
87+
else:
88+
self.linear = torch.nn.Linear(config.hidden_size, output_dim, bias=self.config.linear_bias)
89+
90+
@batch_encoding_wrapper
91+
def forward(self, encoding: BatchEncoding) -> CrossEncoderOutput:
92+
"""Computes contextualized embeddings for the joint query-document input sequence and computes a relevance
93+
score.
94+
95+
:param encoding: Tokenizer encoding for the joint query-document input sequence
96+
:type encoding: BatchEncoding
97+
:return: Output of the model
98+
:rtype: CrossEncoderOutput
99+
"""
100+
if hasattr(self, "decoder"):
101+
# NOTE hack to make T5 cross-encoders work. other encoder-decoder models may not have `decoder` as their
102+
# attribute. maybe find a better way to check for this?
103+
decoder_input_ids = torch.zeros(
104+
(encoding["input_ids"].shape[0], 1), device=encoding["input_ids"].device, dtype=torch.long
105+
)
106+
encoding["decoder_input_ids"] = decoder_input_ids
107+
embeddings = self._backbone_forward(**encoding).last_hidden_state
108+
embeddings = self.pooling(
109+
embeddings, encoding.get("attention_mask", None), pooling_strategy=self.config.pooling_strategy
110+
)
111+
embeddings = self.bert_pool(embeddings)
112+
scores = self.linear(embeddings)
113+
114+
if self.config.scoring_strategy == "mono":
115+
scores = torch.nn.functional.log_softmax(scores.view(-1, 2), dim=-1)[:, 1]
116+
117+
return CrossEncoderOutput(scores=scores.view(-1), embeddings=embeddings)

0 commit comments

Comments
 (0)