Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
colbert.py302 linesDownload Raw Back to late_interaction
1import string2from typing import Any, Iterable, Sequence, Type3 4import numpy as np5from tokenizers import Encoding, Tokenizer6 7from fastembed.common.preprocessor_utils import load_tokenizer8from fastembed.common.types import NumpyArray, Device9from fastembed.common import OnnxProvider10from fastembed.common.onnx_model import OnnxOutputContext11from fastembed.common.utils import define_cache_dir, iter_batch12from fastembed.late_interaction.late_interaction_embedding_base import (13    LateInteractionTextEmbeddingBase,14)15from fastembed.text.onnx_text_model import OnnxTextModel, TextEmbeddingWorker16from fastembed.common.model_description import DenseModelDescription, ModelSource17 18supported_colbert_models: list[DenseModelDescription] = [19    DenseModelDescription(20        model="colbert-ir/colbertv2.0",21        dim=128,22        description="Text embeddings, Unimodal (text), English, 512 input tokens truncation, 2023 year",23        license="mit",24        size_in_GB=0.44,25        sources=ModelSource(hf="colbert-ir/colbertv2.0"),26        model_file="model.onnx",27    ),28    DenseModelDescription(29        model="answerdotai/answerai-colbert-small-v1",30        dim=96,31        description="Text embeddings, Unimodal (text), English, 512 input tokens truncation, 2024 year",32        license="apache-2.0",33        size_in_GB=0.13,34        sources=ModelSource(hf="answerdotai/answerai-colbert-small-v1"),35        model_file="vespa_colbert.onnx",36    ),37]38 39 40class Colbert(LateInteractionTextEmbeddingBase, OnnxTextModel[NumpyArray]):41    QUERY_MARKER_TOKEN_ID = 142    DOCUMENT_MARKER_TOKEN_ID = 243    MIN_QUERY_LENGTH = 31  # it's 32, we add one additional special token in the beginning44    MASK_TOKEN = "[MASK]"45 46    def _post_process_onnx_output(47        self, output: OnnxOutputContext, is_doc: bool = True, **kwargs: Any48    ) -> Iterable[NumpyArray]:49        if not is_doc:50            for embedding in output.model_output:51                yield embedding52        else:53            if output.input_ids is None or output.attention_mask is None:54                raise ValueError(55                    "input_ids and attention_mask must be provided for document post-processing"56                )57 58            for i, token_sequence in enumerate(output.input_ids):59                for j, token_id in enumerate(token_sequence):  # type: ignore60                    if token_id in self.skip_list or token_id == self.pad_token_id:61                        output.attention_mask[i, j] = 062 63            output.model_output *= np.expand_dims(output.attention_mask, 2)64            norm = np.linalg.norm(output.model_output, ord=2, axis=2, keepdims=True)65            norm_clamped = np.maximum(norm, 1e-12)66            output.model_output /= norm_clamped67 68            for embedding, attention_mask in zip(output.model_output, output.attention_mask):69                yield embedding[attention_mask == 1]70 71    def _preprocess_onnx_input(72        self, onnx_input: dict[str, NumpyArray], is_doc: bool = True, **kwargs: Any73    ) -> dict[str, NumpyArray]:74        marker_token = self.DOCUMENT_MARKER_TOKEN_ID if is_doc else self.QUERY_MARKER_TOKEN_ID75        onnx_input["input_ids"] = np.insert(76            onnx_input["input_ids"].astype(np.int64), 1, marker_token, axis=177        )78        onnx_input["attention_mask"] = np.insert(79            onnx_input["attention_mask"].astype(np.int64), 1, 1, axis=180        )81        return onnx_input82 83    def tokenize(self, documents: list[str], is_doc: bool = True, **kwargs: Any) -> list[Encoding]:84        return (85            self._tokenize_documents(documents=documents)86            if is_doc87            else self._tokenize_query(query=next(iter(documents)))88        )89 90    def _tokenize_query(self, query: str) -> list[Encoding]:91        assert self.query_tokenizer is not None92        encoded = self.query_tokenizer.encode_batch([query])93        return encoded94 95    def _tokenize_documents(self, documents: list[str]) -> list[Encoding]:96        encoded = self.tokenizer.encode_batch(documents)  # type: ignore[union-attr]97        return encoded98 99    def token_count(100        self,101        texts: str | Iterable[str],102        batch_size: int = 1024,103        is_doc: bool = True,104        include_extension: bool = False,105        **kwargs: Any,106    ) -> int:107        if not hasattr(self, "model") or self.model is None:108            self.load_onnx_model()  # loads the tokenizer as well109        token_num = 0110        texts = [texts] if isinstance(texts, str) else texts111        tokenizer = self.tokenizer if is_doc else self.query_tokenizer112        assert tokenizer is not None113        for batch in iter_batch(texts, batch_size):114            for tokens in tokenizer.encode_batch(batch):115                if is_doc:116                    token_num += sum(tokens.attention_mask)117                else:118                    attend_count = sum(tokens.attention_mask)119                    if include_extension:120                        token_num += max(attend_count, self.MIN_QUERY_LENGTH)121 122                    else:123                        token_num += attend_count124            if include_extension:125                token_num += len(126                    batch127                )  # add 1 for each cls.DOC_MARKER_TOKEN_ID or cls.QUERY_MARKER_TOKEN_ID128 129        return token_num130 131    @classmethod132    def _list_supported_models(cls) -> list[DenseModelDescription]:133        """Lists the supported models.134 135        Returns:136            list[DenseModelDescription]: A list of DenseModelDescription objects containing the model information.137        """138        return supported_colbert_models139 140    def __init__(141        self,142        model_name: str,143        cache_dir: str | None = None,144        threads: int | None = None,145        providers: Sequence[OnnxProvider] | None = None,146        cuda: bool | Device = Device.AUTO,147        device_ids: list[int] | None = None,148        lazy_load: bool = False,149        device_id: int | None = None,150        specific_model_path: str | None = None,151        **kwargs: Any,152    ):153        """154        Args:155            model_name (str): The name of the model to use.156            cache_dir (str, optional): The path to the cache directory.157                                       Can be set using the `FASTEMBED_CACHE_PATH` env variable.158                                       Defaults to `fastembed_cache` in the system's temp directory.159            threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.160            providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.161                Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.162            cuda (Union[bool, Device], optional): Whether to use cuda for inference. Mutually exclusive with `providers`163                Defaults to Device.AUTO.164            device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in165                workers. Should be used with `cuda` equals to `True`, `Device.AUTO` or `Device.CUDA`, mutually exclusive166                with `providers`. Defaults to None.167            lazy_load (bool, optional): Whether to load the model during class initialization or on demand.168                Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.169            device_id (Optional[int], optional): The device id to use for loading the model in the worker process.170            specific_model_path (Optional[str], optional): The specific path to the onnx model dir if it should be imported from somewhere else171 172        Raises:173            ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.174        """175 176        super().__init__(model_name, cache_dir, threads, **kwargs)177        self.providers = providers178        self.lazy_load = lazy_load179        self._extra_session_options = self._select_exposed_session_options(kwargs)180 181        # List of device ids, that can be used for data parallel processing in workers182        self.device_ids = device_ids183        self.cuda = cuda184 185        # This device_id will be used if we need to load model in current process186        self.device_id: int | None = None187        if device_id is not None:188            self.device_id = device_id189        elif self.device_ids is not None:190            self.device_id = self.device_ids[0]191 192        self.model_description = self._get_model_description(model_name)193        self.cache_dir = str(define_cache_dir(cache_dir))194 195        self._specific_model_path = specific_model_path196        self._model_dir = self.download_model(197            self.model_description,198            self.cache_dir,199            local_files_only=self._local_files_only,200            specific_model_path=self._specific_model_path,201        )202        self.mask_token_id: int | None = None203        self.pad_token_id: int | None = None204        self.skip_list: set[int] = set()205 206        self.query_tokenizer: Tokenizer | None = None207 208        if not self.lazy_load:209            self.load_onnx_model()210 211    def load_onnx_model(self) -> None:212        self._load_onnx_model(213            model_dir=self._model_dir,214            model_file=self.model_description.model_file,215            threads=self.threads,216            providers=self.providers,217            cuda=self.cuda,218            device_id=self.device_id,219            extra_session_options=self._extra_session_options,220        )221        self.query_tokenizer, _ = load_tokenizer(model_dir=self._model_dir)222 223        assert self.tokenizer is not None224        self.mask_token_id = self.special_token_to_id[self.MASK_TOKEN]225        self.pad_token_id = self.tokenizer.padding["pad_id"]226        self.skip_list = {227            self.tokenizer.encode(symbol, add_special_tokens=False).ids[0]228            for symbol in string.punctuation229        }230        current_max_length = self.tokenizer.truncation["max_length"]231        # ensure not to overflow after adding document-marker232        self.tokenizer.enable_truncation(max_length=current_max_length - 1)233        self.query_tokenizer.enable_truncation(max_length=current_max_length - 1)234        self.query_tokenizer.enable_padding(235            pad_token=self.MASK_TOKEN,236            pad_id=self.mask_token_id,237            length=self.MIN_QUERY_LENGTH,238        )239 240    def embed(241        self,242        documents: str | Iterable[str],243        batch_size: int = 256,244        parallel: int | None = None,245        **kwargs: Any,246    ) -> Iterable[NumpyArray]:247        """248        Encode a list of documents into list of embeddings.249        We use mean pooling with attention so that the model can handle variable-length inputs.250 251        Args:252            documents: Iterator of documents or single document to embed253            batch_size: Batch size for encoding -- higher values will use more memory, but be faster254            parallel:255                If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.256                If 0, use all available cores.257                If None, don't use data-parallel processing, use default onnxruntime threading instead.258 259        Returns:260            List of embeddings, one per document261        """262        yield from self._embed_documents(263            model_name=self.model_name,264            cache_dir=str(self.cache_dir),265            documents=documents,266            batch_size=batch_size,267            parallel=parallel,268            providers=self.providers,269            cuda=self.cuda,270            device_ids=self.device_ids,271            local_files_only=self._local_files_only,272            specific_model_path=self._specific_model_path,273            extra_session_options=self._extra_session_options,274            **kwargs,275        )276 277    def query_embed(self, query: str | Iterable[str], **kwargs: Any) -> Iterable[NumpyArray]:278        if isinstance(query, str):279            query = [query]280 281        if not hasattr(self, "model") or self.model is None:282            self.load_onnx_model()283 284        for text in query:285            yield from self._post_process_onnx_output(286                self.onnx_embed([text], is_doc=False), is_doc=False287            )288 289    @classmethod290    def _get_worker_class(cls) -> Type[TextEmbeddingWorker[NumpyArray]]:291        return ColbertEmbeddingWorker292 293 294class ColbertEmbeddingWorker(TextEmbeddingWorker[NumpyArray]):295    def init_embedding(self, model_name: str, cache_dir: str, **kwargs: Any) -> Colbert:296        return Colbert(297            model_name=model_name,298            cache_dir=cache_dir,299            threads=1,300            **kwargs,301        )302 
codekingpro/portable-devtools · Team Ai