Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
image_embedding.py136 linesDownload Raw Back to image
1from typing import Any, Iterable, Sequence, Type2from dataclasses import asdict3 4from fastembed.common.types import NumpyArray, Device5from fastembed.common import ImageInput, OnnxProvider6from fastembed.image.image_embedding_base import ImageEmbeddingBase7from fastembed.image.onnx_embedding import OnnxImageEmbedding8from fastembed.common.model_description import DenseModelDescription9 10 11class ImageEmbedding(ImageEmbeddingBase):12    EMBEDDINGS_REGISTRY: list[Type[ImageEmbeddingBase]] = [OnnxImageEmbedding]13 14    @classmethod15    def list_supported_models(cls) -> list[dict[str, Any]]:16        """17        Lists the supported models.18 19        Returns:20            list[dict[str, Any]]: A list of dictionaries containing the model information.21 22            Example:23                ```24                [25                    {26                        "model": "Qdrant/clip-ViT-B-32-vision",27                        "dim": 512,28                        "description": "CLIP vision encoder based on ViT-B/32",29                        "license": "mit",30                        "size_in_GB": 0.33,31                        "sources": {32                            "hf": "Qdrant/clip-ViT-B-32-vision",33                        },34                        "model_file": "model.onnx",35                    }36                ]37                ```38        """39        return [asdict(model) for model in cls._list_supported_models()]40 41    @classmethod42    def _list_supported_models(cls) -> list[DenseModelDescription]:43        result: list[DenseModelDescription] = []44        for embedding in cls.EMBEDDINGS_REGISTRY:45            result.extend(embedding._list_supported_models())46        return result47 48    def __init__(49        self,50        model_name: str,51        cache_dir: str | None = None,52        threads: int | None = None,53        providers: Sequence[OnnxProvider] | None = None,54        cuda: bool | Device = Device.AUTO,55        device_ids: list[int] | None = None,56        lazy_load: bool = False,57        **kwargs: Any,58    ):59        super().__init__(model_name, cache_dir, threads, **kwargs)60        for EMBEDDING_MODEL_TYPE in self.EMBEDDINGS_REGISTRY:61            supported_models = EMBEDDING_MODEL_TYPE._list_supported_models()62            if any(model_name.lower() == model.model.lower() for model in supported_models):63                self.model = EMBEDDING_MODEL_TYPE(64                    model_name,65                    cache_dir,66                    threads=threads,67                    providers=providers,68                    cuda=cuda,69                    device_ids=device_ids,70                    lazy_load=lazy_load,71                    **kwargs,72                )73                return74 75        raise ValueError(76            f"Model {model_name} is not supported in ImageEmbedding."77            "Please check the supported models using `ImageEmbedding.list_supported_models()`"78        )79 80    @property81    def embedding_size(self) -> int:82        """Get the embedding size of the current model"""83        if self._embedding_size is None:84            self._embedding_size = self.get_embedding_size(self.model_name)85        return self._embedding_size86 87    @classmethod88    def get_embedding_size(cls, model_name: str) -> int:89        """Get the embedding size of the passed model90 91        Args:92            model_name (str): The name of the model to get embedding size for.93 94        Returns:95            int: The size of the embedding.96 97        Raises:98            ValueError: If the model name is not found in the supported models.99        """100        descriptions = cls._list_supported_models()101        embedding_size: int | None = None102        for description in descriptions:103            if description.model.lower() == model_name.lower():104                embedding_size = description.dim105                break106        if embedding_size is None:107            model_names = [description.model for description in descriptions]108            raise ValueError(109                f"Embedding size for model {model_name} was None. "110                f"Available model names: {model_names}"111            )112        return embedding_size113 114    def embed(115        self,116        images: ImageInput | Iterable[ImageInput],117        batch_size: int = 16,118        parallel: int | None = None,119        **kwargs: Any,120    ) -> Iterable[NumpyArray]:121        """122        Encode a list of images into list of embeddings.123 124        Args:125            images: Iterator of image paths or single image path to embed126            batch_size: Batch size for encoding -- higher values will use more memory, but be faster127            parallel:128                If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.129                If 0, use all available cores.130                If None, don't use data-parallel processing, use default onnxruntime threading instead.131 132        Returns:133            List of embeddings, one per document134        """135        yield from self.model.embed(images, batch_size, parallel, **kwargs)136 
codekingpro/portable-devtools · Team Ai