Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
text_cross_encoder.py179 linesDownload Raw Back to cross_encoder
1from typing import Any, Iterable, Sequence, Type2from dataclasses import asdict3 4from fastembed.common import OnnxProvider5from fastembed.common.types import Device6from fastembed.rerank.cross_encoder.onnx_text_cross_encoder import OnnxTextCrossEncoder7from fastembed.rerank.cross_encoder.custom_text_cross_encoder import CustomTextCrossEncoder8 9from fastembed.rerank.cross_encoder.text_cross_encoder_base import TextCrossEncoderBase10from fastembed.common.model_description import (11    ModelSource,12    BaseModelDescription,13)14 15 16class TextCrossEncoder(TextCrossEncoderBase):17    CROSS_ENCODER_REGISTRY: list[Type[TextCrossEncoderBase]] = [18        OnnxTextCrossEncoder,19        CustomTextCrossEncoder,20    ]21 22    @classmethod23    def list_supported_models(cls) -> list[dict[str, Any]]:24        """Lists the supported models.25 26        Returns:27            list[BaseModelDescription]: A list of dictionaries containing the model information.28 29            Example:30                ```31                [32                    {33                        "model": "Xenova/ms-marco-MiniLM-L-6-v2",34                        "size_in_GB": 0.08,35                        "sources": {36                            "hf": "Xenova/ms-marco-MiniLM-L-6-v2",37                        },38                        "model_file": "onnx/model.onnx",39                        "description": "MiniLM-L-6-v2 model optimized for re-ranking tasks.",40                        "license": "apache-2.0",41                    }42                ]43                ```44        """45        return [asdict(model) for model in cls._list_supported_models()]46 47    @classmethod48    def _list_supported_models(cls) -> list[BaseModelDescription]:49        result: list[BaseModelDescription] = []50        for encoder in cls.CROSS_ENCODER_REGISTRY:51            result.extend(encoder._list_supported_models())52        return result53 54    def __init__(55        self,56        model_name: str,57        cache_dir: str | None = None,58        threads: int | None = None,59        providers: Sequence[OnnxProvider] | None = None,60        cuda: bool | Device = Device.AUTO,61        device_ids: list[int] | None = None,62        lazy_load: bool = False,63        **kwargs: Any,64    ):65        super().__init__(model_name, cache_dir, threads, **kwargs)66 67        for CROSS_ENCODER_TYPE in self.CROSS_ENCODER_REGISTRY:68            supported_models = CROSS_ENCODER_TYPE._list_supported_models()69            if any(model_name.lower() == model.model.lower() for model in supported_models):70                self.model = CROSS_ENCODER_TYPE(71                    model_name=model_name,72                    cache_dir=cache_dir,73                    threads=threads,74                    providers=providers,75                    cuda=cuda,76                    device_ids=device_ids,77                    lazy_load=lazy_load,78                    **kwargs,79                )80                return81 82        raise ValueError(83            f"Model {model_name} is not supported in TextCrossEncoder."84            "Please check the supported models using `TextCrossEncoder.list_supported_models()`"85        )86 87    def rerank(88        self, query: str, documents: Iterable[str], batch_size: int = 64, **kwargs: Any89    ) -> Iterable[float]:90        """Rerank a list of documents based on a query.91 92        Args:93            query: Query to rerank the documents against94            documents: Iterator of documents to rerank95            batch_size: Batch size for reranking96 97        Returns:98            Iterable of scores for each document99        """100        yield from self.model.rerank(query, documents, batch_size=batch_size, **kwargs)101 102    def rerank_pairs(103        self,104        pairs: Iterable[tuple[str, str]],105        batch_size: int = 64,106        parallel: int | None = None,107        **kwargs: Any,108    ) -> Iterable[float]:109        """110        Rerank a list of query-document pairs.111 112        Args:113            pairs (Iterable[tuple[str, str]]): An iterable of tuples, where each tuple contains a query and a document114                to be scored together.115            batch_size (int, optional): The number of query-document pairs to process in a single batch. Defaults to 64.116            parallel (Optional[int], optional): The number of parallel processes to use for reranking.117                If None, parallelization is disabled. Defaults to None.118            **kwargs (Any): Additional arguments to pass to the underlying reranking model.119 120        Returns:121            Iterable[float]: An iterable of scores corresponding to each query-document pair in the input.122            Higher scores indicate a stronger match between the query and the document.123 124        Example:125            >>> encoder = TextCrossEncoder("Xenova/ms-marco-MiniLM-L-6-v2")126            >>> pairs = [("What is AI?", "Artificial intelligence is ..."), ("What is ML?", "Machine learning is ...")]127            >>> scores = list(encoder.rerank_pairs(pairs))128            >>> print(list(map(lambda x: round(x, 2), scores)))129            [-1.24, -10.6]130        """131        yield from self.model.rerank_pairs(132            pairs, batch_size=batch_size, parallel=parallel, **kwargs133        )134 135    @classmethod136    def add_custom_model(137        cls,138        model: str,139        sources: ModelSource,140        model_file: str = "onnx/model.onnx",141        description: str = "",142        license: str = "",143        size_in_gb: float = 0.0,144        additional_files: list[str] | None = None,145    ) -> None:146        registered_models = cls._list_supported_models()147        for registered_model in registered_models:148            if model == registered_model.model:149                raise ValueError(150                    f"Model {model} is already registered in CrossEncoderModel, if you still want to add this model, "151                    f"please use another model name"152                )153 154        CustomTextCrossEncoder.add_model(155            BaseModelDescription(156                model=model,157                sources=sources,158                model_file=model_file,159                description=description,160                license=license,161                size_in_GB=size_in_gb,162                additional_files=additional_files or [],163            )164        )165 166    def token_count(167        self, pairs: Iterable[tuple[str, str]], batch_size: int = 1024, **kwargs: Any168    ) -> int:169        """Returns the number of tokens in the pairs.170 171        Args:172            pairs: Iterable of tuples, where each tuple contains a query and a document to be tokenized173            batch_size: Batch size for tokenizing174 175        Returns:176            token count: overall number of tokens in the pairs177        """178        return self.model.token_count(pairs, batch_size=batch_size, **kwargs)179 
codekingpro/portable-devtools · Team Ai