diff --git a/fastembed/text/custom_text_embedding.py b/fastembed/text/custom_text_embedding.py index c3d5cf5f..22482a8a 100644 --- a/fastembed/text/custom_text_embedding.py +++ b/fastembed/text/custom_text_embedding.py @@ -20,6 +20,7 @@ class PostprocessingConfig: pooling: PoolingType normalization: bool + output_name: str | None = None class CustomTextEmbedding(OnnxTextEmbedding): @@ -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]: @@ -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 ) @@ -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, diff --git a/fastembed/text/text_embedding.py b/fastembed/text/text_embedding.py index 8848389b..14966383 100644 --- a/fastembed/text/text_embedding.py +++ b/fastembed/text/text_embedding.py @@ -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: @@ -80,6 +81,7 @@ def add_custom_model( ), pooling=pooling, normalization=normalization, + output_name=output_name, ) def __init__( diff --git a/tests/test_custom_models.py b/tests/test_custom_models.py index 1cf47f7f..c20d63f0 100644 --- a/tests/test_custom_models.py +++ b/tests/test_custom_models.py @@ -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"