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
18 changes: 17 additions & 1 deletion fastembed/text/custom_text_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
class PostprocessingConfig:
pooling: PoolingType
normalization: bool
output_name: str | None = None


class CustomTextEmbedding(OnnxTextEmbedding):
Expand Down Expand Up @@ -54,6 +55,19 @@ def __init__(
postprocessing_config = self.POSTPROCESSING_MAPPING[self.model_description.model]
self._pooling = postprocessing_config.pooling
self._normalization = postprocessing_config.normalization
if postprocessing_config.output_name is not None:
self.ONNX_OUTPUT_NAMES = [postprocessing_config.output_name]

def load_onnx_model(self) -> None:
super().load_onnx_model()
# with eager loading this runs inside super().__init__(), before ONNX_OUTPUT_NAMES is set,
# so the output name is taken from the registered config
output_name = self.POSTPROCESSING_MAPPING[self.model_description.model].output_name
output_names = [output.name for output in self.model.get_outputs()] # type: ignore[union-attr]
if output_name is not None and output_name not in output_names:
raise ValueError(
f"Output {output_name!r} not found in the model, available outputs: {output_names}"
)

@classmethod
def _list_supported_models(cls) -> list[DenseModelDescription]:
Expand Down Expand Up @@ -108,10 +122,11 @@ def add_model(
model_description: DenseModelDescription,
pooling: PoolingType,
normalization: bool,
output_name: str | None = None,
) -> None:
cls.SUPPORTED_MODELS.append(model_description)
cls.POSTPROCESSING_MAPPING[model_description.model] = PostprocessingConfig(
pooling=pooling, normalization=normalization
pooling=pooling, normalization=normalization, output_name=output_name
)


Expand All @@ -135,6 +150,7 @@ def init_embedding(
model_description,
pooling=postprocessing_config.pooling,
normalization=postprocessing_config.normalization,
output_name=postprocessing_config.output_name,
)
return CustomTextEmbedding(
model_name=model_name,
Expand Down
2 changes: 2 additions & 0 deletions fastembed/text/text_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,7 @@ def add_custom_model(
license: str = "",
size_in_gb: float = 0.0,
additional_files: list[str] | None = None,
output_name: str | None = None,
) -> None:
registered_models = cls._list_supported_models()
for registered_model in registered_models:
Expand All @@ -80,6 +81,7 @@ def add_custom_model(
),
pooling=pooling,
normalization=normalization,
output_name=output_name,
)

def __init__(
Expand Down
37 changes: 37 additions & 0 deletions tests/test_custom_models.py
Original file line number Diff line number Diff line change
Expand Up @@ -101,6 +101,43 @@ def test_text_custom_model_parallel_processing():
delete_model_cache(model.model._model_dir)


def test_text_custom_model_output_name():
is_ci = os.getenv("CI")
custom_model_name = "custom/granite-embedding-small-english-r2"

# onnx outputs are `last_hidden_state` and `sentence_embedding`, the builtin model uses the latter
TextEmbedding.add_custom_model(
custom_model_name,
pooling=PoolingType.DISABLED,
normalization=False,
sources=ModelSource(hf="onnx-community/granite-embedding-small-english-r2-ONNX"),
dim=384,
additional_files=["onnx/model.onnx_data"],
output_name="sentence_embedding",
)

model = TextEmbedding(custom_model_name)
builtin_model = TextEmbedding("ibm-granite/granite-embedding-small-english-r2")
docs = ["hello world", "flag embedding"]
expected = np.stack(list(builtin_model.embed(docs)), axis=0)

embeddings = np.stack(list(model.embed(docs)), axis=0)
assert np.allclose(embeddings, expected, atol=1e-5)

# workers re-register the model, so they have to receive the output name too
embeddings = np.stack(list(model.embed(docs, batch_size=1, parallel=2)), axis=0)
assert np.allclose(embeddings, expected, atol=1e-5)

CustomTextEmbedding.POSTPROCESSING_MAPPING[custom_model_name] = PostprocessingConfig(
pooling=PoolingType.DISABLED, normalization=False, output_name="sentence_embeddings"
)
with pytest.raises(ValueError, match="available outputs"):
TextEmbedding(custom_model_name)

if is_ci:
delete_model_cache(model.model._model_dir)


def test_cross_encoder_custom_model():
is_ci = os.getenv("CI")
custom_model_name = "Xenova/ms-marco-MiniLM-L-4-v2"
Expand Down
Loading