-
Notifications
You must be signed in to change notification settings - Fork 5
feat(vectorisers): Add FastEmbedVectoriser implementation of VectoriserBase and add fast embed dependency group #224
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
a19bcb3
3f8f2cc
6270bea
9d046da
9b027c4
f94af1f
e0172fe
9f24eed
3ce73b2
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,149 @@ | ||
| """A module that provides a wrapper for FastEmbed models to generate text embeddings.""" | ||
|
|
||
| import numpy as np | ||
|
|
||
| from classifai._optional import check_deps | ||
| from classifai.exceptions import ExternalServiceError, VectorisationError | ||
|
|
||
| from .base import VectoriserBase | ||
|
|
||
|
|
||
| class FastEmbedVectoriser(VectoriserBase): | ||
| """A lightweight wrapper class for generating embeddings with FastEmbed. | ||
|
|
||
| The `FastEmbedVectoriser` uses FastEmbed's ONNX backend to generate | ||
| embeddings from FastEmbed-compatible sentence embedding models without | ||
| requiring `torch` or `transformers` as runtime dependencies. The | ||
| `model_name` must be a name recognised by FastEmbed. To see all | ||
| supported models, you can run: | ||
| `FastEmbedVectoriser.list_supported_models()` | ||
|
|
||
| To use a pre-downloaded model in an air-gapped environment, provide both | ||
| its official FastEmbed `model_name` and its local directory through | ||
| `specific_model_path`. FastEmbed uses `model_name` to identify the model | ||
| configuration and `specific_model_path` to locate its local ONNX files. | ||
|
|
||
| Attributes: | ||
| model_name (str): The official FastEmbed name of the embedding model. | ||
| model (fastembed.TextEmbedding): The FastEmbed model instance. | ||
| specific_model_path (str | None): The path of the local FastEmbed model. | ||
| """ | ||
|
|
||
| def __init__( | ||
| self, | ||
| model_name: str, | ||
| specific_model_path: str | None = None, | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Consider removing this parameter and then user can pass it as part of the kwarg to the constructor? Looking at how this is handled in the HuggingFaceVectoriser class (HuggingFace also has its own way of local cache checking for models under the hood), we don't have a special parameter in that constructor. So mirroring that might be good for consistency
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. @jamie-ons added this so not sure of the rationale, don't mind either way.
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Will confirm when @jamie-ons returns from leave, but I think this was to allow someone to specify whether to use ONNX format weights on a model card if there's multiple options available
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. The HuggingFaceVectoriser allows a user to put in a path to a local model as the I therefore decided to add in this specific model path as a key use of the FastEmbedVectoriser is to run models in enviroments where resources may be lower. This could often mean that ability to download large packages (torch) or files (model weights) may be restricted. I also don't think the FastEmbed documentation is that clear about how to run models from a local download and so thought adding the argument would save them time in researching how to do it. I will update the docstrings to make this clearer. |
||
| model_kwargs: dict | None = None, | ||
| ): | ||
| """Initialises the FastEmbedVectoriser with the specified model name. | ||
|
|
||
| Args: | ||
| model_name (str): The official name of the embedding model for FastEmbed | ||
| (e.g., "sentence-transformers/all-MiniLM-L6-v2"). | ||
| specific_model_path (str | None): The local directory containing | ||
| the pre-downloaded ONNX model. To run offline, provide this | ||
| value together with the model's official `model_name`. | ||
| Defaults to None. | ||
| model_kwargs (dict): [optional] Additional keyword arguments to | ||
| pass to the model (e.g., `cache_dir`). Defaults to None. | ||
|
|
||
| Raises: | ||
| `ExternalServiceError`: If the FastEmbed model cannot be loaded. | ||
| """ | ||
| check_deps(["fastembed"], extra="fastembed") | ||
| from fastembed import TextEmbedding # type: ignore | ||
|
|
||
| self.model_name = model_name | ||
| self.specific_model_path = specific_model_path | ||
| model_kwargs = dict(model_kwargs or {}) | ||
|
|
||
| if self.specific_model_path is not None: | ||
| model_kwargs["specific_model_path"] = str(self.specific_model_path) | ||
|
|
||
| try: | ||
| self.model = TextEmbedding(model_name=self.model_name, **model_kwargs) | ||
| except Exception as e: | ||
| raise ExternalServiceError( | ||
| "Failed to load FastEmbed model.", | ||
| context={ | ||
| "vectoriser": "fastembed", | ||
| "model": self.model_name, | ||
| "cause": str(e), | ||
| "cause_type": type(e).__name__, | ||
| }, | ||
| ) from e | ||
|
|
||
| def transform(self, texts: str | list[str]) -> np.ndarray: | ||
| """Transforms input text(s) into embeddings using FastEmbed. | ||
|
|
||
| Args: | ||
| texts (str | list[str]): The input text(s) to embed. Can be a | ||
| single string or a list of strings. | ||
|
|
||
| Returns: | ||
| numpy.ndarray: A 2D array of embeddings, where each row | ||
| corresponds to an input text. | ||
|
|
||
| Raises: | ||
| `VectorisationError`: If FastEmbed fails to generate or parse | ||
| embeddings. | ||
| """ | ||
| # If a single string is passed as arg to texts, convert to list | ||
| if isinstance(texts, str): | ||
| texts = [texts] | ||
|
|
||
| try: | ||
| raw_embeddings = list(self.model.embed(texts)) | ||
| except Exception as e: | ||
| raise VectorisationError( | ||
| "Failed to generate embeddings using FastEmbed.", | ||
| context={ | ||
| "vectoriser": "fastembed", | ||
| "model": self.model_name, | ||
| "n_texts": len(texts), | ||
| "cause": str(e), | ||
| "cause_type": type(e).__name__, | ||
| }, | ||
| ) from e | ||
|
|
||
| try: | ||
| embeddings = np.asarray(raw_embeddings, dtype=np.float32) | ||
| except Exception as e: | ||
| raise VectorisationError( | ||
| "Failed to convert FastEmbed embeddings to a numpy array.", | ||
| context={ | ||
| "vectoriser": "fastembed", | ||
| "model": self.model_name, | ||
| "n_texts": len(texts), | ||
| "cause": str(e), | ||
| "cause_type": type(e).__name__, | ||
| }, | ||
| ) from e | ||
|
|
||
| if embeddings.ndim == 1: | ||
| embeddings = embeddings.reshape(1, -1) | ||
|
|
||
| if embeddings.ndim != 2: # noqa: PLR2004 | ||
| raise VectorisationError( | ||
| "FastEmbed returned embeddings with an unexpected shape.", | ||
| context={ | ||
| "vectoriser": "fastembed", | ||
| "model": self.model_name, | ||
| "n_texts": len(texts), | ||
| "shape": list(embeddings.shape), | ||
| }, | ||
| ) | ||
|
|
||
| return embeddings | ||
|
|
||
| @staticmethod | ||
| def list_supported_models() -> list[dict[str, any]]: | ||
| """Wrapper to list the supported models. | ||
|
|
||
| Returns: | ||
| list[dict[str, Any]]: A list of dictionaries containing the model information. | ||
| """ | ||
| check_deps(["fastembed"], extra="fastembed") | ||
| from fastembed import TextEmbedding # type: ignore | ||
|
|
||
| return TextEmbedding.list_supported_models() | ||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Do we update the changelog with every PR on the package? I thought we just did one changelog update per release, and looked back at the merged PRs when writing it?
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Soz haven't had to think about CHANGELOGs in a while as that's done for me on scanner ;)
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Will have a look later today