Skip to content

Commit c5938a2

Browse files
authored
Add Sentence-Transformers Embedding Model and add support with TextEmbedder Task API (#2740)
* Add TextEmbedder task API with BertTextEmbedder for sentence-transformers support * Optimize the code * Add checkpoint conversion script * address gemini comments * Add dynamic config detection and Refactor conversion script to reuse convert_bert.py hooks * Address reviewer comments, make base classes generic, add standardized tests * resolve TextEmbedder training crash by supporting custom compile_kwargs in task tests * Update mean pool test
1 parent 56b2889 commit c5938a2

11 files changed

Lines changed: 1285 additions & 12 deletions

keras_hub/api/models/__init__.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -68,6 +68,12 @@
6868
from keras_hub.src.models.bert.bert_text_classifier_preprocessor import (
6969
BertTextClassifierPreprocessor as BertTextClassifierPreprocessor,
7070
)
71+
from keras_hub.src.models.bert.bert_text_embedder import (
72+
BertTextEmbedder as BertTextEmbedder,
73+
)
74+
from keras_hub.src.models.bert.bert_text_embedder_preprocessor import (
75+
BertTextEmbedderPreprocessor as BertTextEmbedderPreprocessor,
76+
)
7177
from keras_hub.src.models.bert.bert_tokenizer import (
7278
BertTokenizer as BertTokenizer,
7379
)
@@ -846,6 +852,10 @@
846852
from keras_hub.src.models.text_classifier_preprocessor import (
847853
TextClassifierPreprocessor as TextClassifierPreprocessor,
848854
)
855+
from keras_hub.src.models.text_embedder import TextEmbedder as TextEmbedder
856+
from keras_hub.src.models.text_embedder_preprocessor import (
857+
TextEmbedderPreprocessor as TextEmbedderPreprocessor,
858+
)
849859
from keras_hub.src.models.text_to_image import TextToImage as TextToImage
850860
from keras_hub.src.models.text_to_image_preprocessor import (
851861
TextToImagePreprocessor as TextToImagePreprocessor,

keras_hub/src/models/bert/bert_presets.py

Lines changed: 117 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -110,4 +110,121 @@
110110
},
111111
"kaggle_handle": "kaggle://keras/bert/keras/bert_tiny_en_uncased_sst2/5",
112112
},
113+
# Sentence-transformer models (BERT backbone, fine-tuned for embeddings).
114+
# "all-*" family: general-purpose sentence embedding models.
115+
"all_minilm_l6_v2_en": {
116+
"metadata": {
117+
"description": (
118+
"6-layer MiniLM sentence embedding model. Maps sentences "
119+
"to 384-dimensional dense vectors. Trained on 1B+ sentence "
120+
"pairs for semantic similarity, search, and clustering."
121+
),
122+
"params": 22713216,
123+
"path": "bert",
124+
},
125+
"kaggle_handle": "kaggle://keras/sentence-transformers/keras/all_minilm_l6_v2_en/1",
126+
},
127+
"all_minilm_l6_v1_en": {
128+
"metadata": {
129+
"description": (
130+
"6-layer MiniLM sentence embedding model (v1). Maps "
131+
"sentences to 384-dimensional dense vectors."
132+
),
133+
"params": 22713216,
134+
"path": "bert",
135+
},
136+
"kaggle_handle": "kaggle://keras/sentence-transformers/keras/all_minilm_l6_v1_en/1",
137+
},
138+
"all_minilm_l12_v2_en": {
139+
"metadata": {
140+
"description": (
141+
"12-layer MiniLM sentence embedding model. Maps sentences "
142+
"to 384-dimensional dense vectors. Higher accuracy than "
143+
"the L6 variant with moderate speed tradeoff."
144+
),
145+
"params": 33360000,
146+
"path": "bert",
147+
},
148+
"kaggle_handle": "kaggle://keras/sentence-transformers/keras/all_minilm_l12_v2_en/1",
149+
},
150+
# "paraphrase-*" family: optimized for paraphrase detection.
151+
"paraphrase_minilm_l3_v2_en": {
152+
"metadata": {
153+
"description": (
154+
"3-layer MiniLM model for paraphrase detection. Ultra-fast "
155+
"with 384-dimensional sentence embeddings."
156+
),
157+
"params": 17066496,
158+
"path": "bert",
159+
},
160+
"kaggle_handle": "kaggle://keras/sentence-transformers/keras/paraphrase_minilm_l3_v2_en/1",
161+
},
162+
"paraphrase_minilm_l6_v2_en": {
163+
"metadata": {
164+
"description": (
165+
"6-layer MiniLM model for paraphrase detection. Fast "
166+
"with 384-dimensional sentence embeddings."
167+
),
168+
"params": 22713216,
169+
"path": "bert",
170+
},
171+
"kaggle_handle": "kaggle://keras/sentence-transformers/keras/paraphrase_minilm_l6_v2_en/1",
172+
},
173+
"paraphrase_minilm_l12_v2_en": {
174+
"metadata": {
175+
"description": (
176+
"12-layer MiniLM model for paraphrase detection with "
177+
"384-dimensional sentence embeddings."
178+
),
179+
"params": 33360000,
180+
"path": "bert",
181+
},
182+
"kaggle_handle": "kaggle://keras/sentence-transformers/keras/paraphrase_minilm_l12_v2_en/1",
183+
},
184+
# "multi-qa-*" family: optimized for question answering / semantic search.
185+
"multi_qa_minilm_l6_cos_v1_en": {
186+
"metadata": {
187+
"description": (
188+
"6-layer MiniLM model for semantic search with cosine "
189+
"similarity. Trained on 215M QA pairs."
190+
),
191+
"params": 22713216,
192+
"path": "bert",
193+
},
194+
"kaggle_handle": "kaggle://keras/sentence-transformers/keras/multi_qa_minilm_l6_cos_v1_en/1",
195+
},
196+
"multi_qa_minilm_l6_dot_v1_en": {
197+
"metadata": {
198+
"description": (
199+
"6-layer MiniLM model for semantic search with dot-product "
200+
"similarity. Trained on 215M QA pairs."
201+
),
202+
"params": 22713216,
203+
"path": "bert",
204+
},
205+
"kaggle_handle": "kaggle://keras/sentence-transformers/keras/multi_qa_minilm_l6_dot_v1_en/1",
206+
},
207+
# "msmarco-*" family: optimized for information retrieval.
208+
"msmarco_minilm_l6_cos_v5_en": {
209+
"metadata": {
210+
"description": (
211+
"6-layer MiniLM model for information retrieval with "
212+
"cosine similarity. Trained on MS MARCO passage ranking."
213+
),
214+
"params": 22713216,
215+
"path": "bert",
216+
},
217+
"kaggle_handle": "kaggle://keras/sentence-transformers/keras/msmarco_minilm_l6_cos_v5_en/1",
218+
},
219+
"msmarco_minilm_l12_cos_v5_en": {
220+
"metadata": {
221+
"description": (
222+
"12-layer MiniLM model for information retrieval with "
223+
"cosine similarity. Trained on MS MARCO passage ranking."
224+
),
225+
"params": 33360000,
226+
"path": "bert",
227+
},
228+
"kaggle_handle": "kaggle://keras/sentence-transformers/keras/msmarco_minilm_l12_cos_v5_en/1",
229+
},
113230
}
Lines changed: 207 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,207 @@
1+
from keras import ops
2+
3+
from keras_hub.src.api_export import keras_hub_export
4+
from keras_hub.src.models.bert.bert_backbone import BertBackbone
5+
from keras_hub.src.models.bert.bert_text_embedder_preprocessor import (
6+
BertTextEmbedderPreprocessor,
7+
)
8+
from keras_hub.src.models.text_embedder import TextEmbedder
9+
10+
11+
@keras_hub_export("keras_hub.models.BertTextEmbedder")
12+
class BertTextEmbedder(TextEmbedder):
13+
"""An end-to-end BERT model for generating sentence embeddings.
14+
15+
This model attaches a mean pooling and L2 normalization head to a
16+
`keras_hub.models.BertBackbone` instance, mapping from the backbone
17+
outputs to fixed-size sentence embeddings suitable for semantic
18+
similarity, clustering, and retrieval tasks.
19+
20+
This is the architecture used by sentence-transformers models like
21+
`all-MiniLM-L6-v2`. For usage of this model with pre-trained weights,
22+
use the `from_preset()` constructor.
23+
24+
This model can optionally be configured with a `preprocessor` layer, in
25+
which case it will automatically apply preprocessing to raw inputs during
26+
`fit()`, `predict()`, and `evaluate()`. This is done by default when
27+
creating the model with `from_preset()`.
28+
29+
Disclaimer: Pre-trained models are provided on an "as is" basis, without
30+
warranties or conditions of any kind.
31+
32+
Args:
33+
backbone: A `keras_hub.models.BertBackbone` instance.
34+
preprocessor: A `keras_hub.models.BertTextEmbedderPreprocessor` or
35+
`None`. If `None`, this model will not apply preprocessing, and
36+
inputs should be preprocessed before calling the model.
37+
pooling_mode: str. The pooling strategy to use. One of `"mean"`,
38+
`"cls"`, or `"max"`. Defaults to `"mean"`.
39+
- `"mean"`: Attention-mask-aware mean pooling over all tokens.
40+
- `"cls"`: Use the `[CLS]` token representation.
41+
- `"max"`: Max pooling over all tokens.
42+
normalize: bool. Whether to L2 normalize the output embeddings.
43+
Defaults to `True`.
44+
45+
Examples:
46+
47+
Raw string data.
48+
```python
49+
embedder = keras_hub.models.BertTextEmbedder.from_preset(
50+
"all_minilm_l6_v2_en",
51+
)
52+
53+
# Semantic search.
54+
query = "Which planet is known as the Red Planet?"
55+
documents = [
56+
"Mars is often referred to as the Red Planet.",
57+
"Venus is often called Earth's twin.",
58+
]
59+
q_emb = embedder.encode_text(query)
60+
d_embs = embedder.encode_documents(documents)
61+
sims = embedder.similarity(q_emb, d_embs)
62+
print("Best match:", documents[sims.argmax()])
63+
```
64+
65+
Preprocessed integer data.
66+
```python
67+
features = {
68+
"token_ids": np.ones(shape=(2, 12), dtype="int32"),
69+
"segment_ids": np.array([[0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 0, 0]] * 2),
70+
"padding_mask": np.array([[1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 0, 0]] * 2),
71+
}
72+
73+
embedder = keras_hub.models.BertTextEmbedder.from_preset(
74+
"all_minilm_l6_v2_en",
75+
preprocessor=None,
76+
)
77+
embeddings = embedder.predict(features)
78+
```
79+
"""
80+
81+
backbone_cls = BertBackbone
82+
preprocessor_cls = BertTextEmbedderPreprocessor
83+
84+
def __init__(
85+
self,
86+
backbone,
87+
preprocessor=None,
88+
pooling_mode="mean",
89+
normalize=True,
90+
**kwargs,
91+
):
92+
# === Layers ===
93+
self.backbone = backbone
94+
self.preprocessor = preprocessor
95+
96+
# === Functional Model ===
97+
inputs = backbone.input
98+
backbone_outputs = backbone(inputs)
99+
sequence_output = backbone_outputs["sequence_output"]
100+
padding_mask = inputs["padding_mask"]
101+
102+
# Apply pooling.
103+
if pooling_mode == "mean":
104+
pooled = self._mean_pooling(sequence_output, padding_mask)
105+
elif pooling_mode == "cls":
106+
pooled = sequence_output[:, 0, :]
107+
elif pooling_mode == "max":
108+
pooled = self._max_pooling(sequence_output, padding_mask)
109+
else:
110+
raise ValueError(
111+
f"Invalid pooling_mode: '{pooling_mode}'. "
112+
"Expected one of 'mean', 'cls', or 'max'."
113+
)
114+
115+
# Apply L2 normalization.
116+
if normalize:
117+
pooled = self._l2_normalize(pooled)
118+
119+
super().__init__(
120+
inputs=inputs,
121+
outputs=pooled,
122+
**kwargs,
123+
)
124+
125+
# === Config ===
126+
self.pooling_mode = pooling_mode
127+
self.normalize = normalize
128+
129+
@staticmethod
130+
def _mean_pooling(sequence_output, padding_mask):
131+
"""Attention-mask-aware mean pooling over token embeddings."""
132+
# Expand mask: [batch, seq_len] -> [batch, seq_len, 1]
133+
mask = ops.cast(
134+
ops.expand_dims(padding_mask, axis=-1), sequence_output.dtype
135+
)
136+
# Sum token embeddings, masked.
137+
sum_embeddings = ops.sum(sequence_output * mask, axis=1)
138+
# Sum mask for normalization.
139+
sum_mask = ops.maximum(ops.sum(mask, axis=1), 1e-9)
140+
return sum_embeddings / sum_mask
141+
142+
@staticmethod
143+
def _max_pooling(sequence_output, padding_mask):
144+
"""Max pooling over token embeddings, ignoring padding."""
145+
mask = ops.cast(ops.expand_dims(padding_mask, axis=-1), dtype="bool")
146+
# Set padding positions to -inf so they don't affect max.
147+
fill_value = ops.cast(
148+
ops.convert_to_tensor(float("-inf")), sequence_output.dtype
149+
)
150+
masked_output = ops.where(mask, sequence_output, fill_value)
151+
return ops.max(masked_output, axis=1)
152+
153+
@staticmethod
154+
def _l2_normalize(embeddings):
155+
"""L2 normalize embeddings to unit length."""
156+
return ops.nn.normalize(embeddings, axis=-1, order=2)
157+
158+
def get_config(self):
159+
config = super().get_config()
160+
config.update(
161+
{
162+
"pooling_mode": self.pooling_mode,
163+
"normalize": self.normalize,
164+
}
165+
)
166+
return config
167+
168+
def encode_documents(self, documents, **kwargs):
169+
"""Encode a string or list of documents into embeddings.
170+
171+
This is a convenience method that wraps `predict()` for a
172+
single document or batch of documents. The output embeddings
173+
are suitable for computing similarity against query embeddings.
174+
175+
Args:
176+
documents: A string or list of strings to encode.
177+
**kwargs: Additional keyword arguments passed to
178+
`predict()`.
179+
180+
Returns:
181+
A tensor of shape `(batch_size, embedding_dim)`.
182+
"""
183+
if isinstance(documents, str):
184+
documents = [documents]
185+
return self.predict(documents, **kwargs)
186+
187+
def similarity(self, query_embeddings, document_embeddings):
188+
"""Compute similarity between query and document embeddings.
189+
190+
Computes the dot product between query and document embeddings.
191+
When embeddings are L2 normalized (the default), this is
192+
equivalent to cosine similarity. Returns a similarity matrix
193+
of shape `(num_queries, num_documents)`.
194+
195+
Args:
196+
query_embeddings: An array or tensor of shape
197+
`(num_queries, embedding_dim)`.
198+
document_embeddings: An array or tensor of shape
199+
`(num_documents, embedding_dim)`.
200+
201+
Returns:
202+
A tensor of shape `(num_queries, num_documents)` containing
203+
similarity scores.
204+
"""
205+
query_tensor = ops.convert_to_tensor(query_embeddings)
206+
document_tensor = ops.convert_to_tensor(document_embeddings)
207+
return ops.matmul(query_tensor, ops.transpose(document_tensor))

0 commit comments

Comments
 (0)