Team Ai
Modelpublic

joshcd/MNLP_M3_document_encoder

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
0likes82downloads
eval_mteb.py846 linesDownload Raw Back to scripts
1from __future__ import annotations2 3import argparse4import logging5import math6import queue7from typing import Dict, List, Optional, Union8 9import numpy as np10import torch11import torch.multiprocessing as mp12from tqdm.autonotebook import trange13from transformers import AutoModel, AutoTokenizer14 15from mteb import MTEB16 17TASK_LIST_CLASSIFICATION = [18    "AmazonCounterfactualClassification",19    "AmazonPolarityClassification",20    "AmazonReviewsClassification",21    "Banking77Classification",22    "EmotionClassification",23    "ImdbClassification",24    "MassiveIntentClassification",25    "MassiveScenarioClassification",26    "MTOPDomainClassification",27    "MTOPIntentClassification",28    "ToxicConversationsClassification",29    "TweetSentimentExtractionClassification",30]31 32TASK_LIST_CLUSTERING = [33    "ArxivClusteringP2P",34    "ArxivClusteringS2S",35    "BiorxivClusteringP2P",36    "BiorxivClusteringS2S",37    "MedrxivClusteringP2P",38    "MedrxivClusteringS2S",39    "RedditClustering",40    "RedditClusteringP2P",41    "StackExchangeClustering",42    "StackExchangeClusteringP2P",43    "TwentyNewsgroupsClustering",44]45 46TASK_LIST_PAIR_CLASSIFICATION = [47    "SprintDuplicateQuestions",48    "TwitterSemEval2015",49    "TwitterURLCorpus",50]51 52TASK_LIST_RERANKING = [53    "AskUbuntuDupQuestions",54    "MindSmallReranking",55    "SciDocsRR",56    "StackOverflowDupQuestions",57]58 59TASK_LIST_RETRIEVAL = [60    "ArguAna",61    "ClimateFEVER",62    "CQADupstackAndroidRetrieval",63    "CQADupstackEnglishRetrieval",64    "CQADupstackGamingRetrieval",65    "CQADupstackGisRetrieval",66    "CQADupstackMathematicaRetrieval",67    "CQADupstackPhysicsRetrieval",68    "CQADupstackProgrammersRetrieval",69    "CQADupstackStatsRetrieval",70    "CQADupstackTexRetrieval",71    "CQADupstackUnixRetrieval",72    "CQADupstackWebmastersRetrieval",73    "CQADupstackWordpressRetrieval",74    "DBPedia",75    "FEVER",76    "FiQA2018",77    "HotpotQA",78    "MSMARCO",79    "NFCorpus",80    "NQ",81    "QuoraRetrieval",82    "SCIDOCS",83    "SciFact",84    "Touche2020",85    "TRECCOVID",86]87 88TASK_LIST_STS = [89    "BIOSSES",90    "SICK-R",91    "STS12",92    "STS13",93    "STS14",94    "STS15",95    "STS16",96    "STS17",97    "STS22",98    "STSBenchmark",99    "SummEval",100]101 102MTEB_TASK_LIST = (103    TASK_LIST_CLASSIFICATION104    + TASK_LIST_CLUSTERING105    + TASK_LIST_PAIR_CLASSIFICATION106    + TASK_LIST_RERANKING107    + TASK_LIST_RETRIEVAL108    + TASK_LIST_STS109)110 111 112CMTEB_TASK_LIST = [113    "TNews",114    "IFlyTek",115    "MultilingualSentiment",116    "JDReview",117    "OnlineShopping",118    "Waimai",119    "AmazonReviewsClassification",120    "MassiveIntentClassification",121    "MassiveScenarioClassification",122    "MultilingualSentiment",123    "CLSClusteringS2S",124    "CLSClusteringP2P",125    "ThuNewsClusteringS2S",126    "ThuNewsClusteringP2P",127    "Ocnli",128    "Cmnli",129    "T2Reranking",130    "MmarcoReranking",131    "CMedQAv1",132    "CMedQAv2",133    "T2Retrieval",134    "MMarcoRetrieval",135    "DuRetrieval",136    "CovidRetrieval",137    "CmedqaRetrieval",138    "EcomRetrieval",139    "MedicalRetrieval",140    "VideoRetrieval",141    "ATEC",142    "BQ",143    "LCQMC",144    "PAWSX",145    "STSB",146    "AFQMC",147    "QBQTC",148    "STS22",149]150 151MTEB_PL = [152    "CBD",153    "PolEmo2.0-IN",154    "PolEmo2.0-OUT",155    "AllegroReviews",156    "PAC",157    "MassiveIntentClassification",158    "MassiveScenarioClassification",159    "SICK-E-PL",160    "PPC",161    "CDSC-E",162    "PSC",163    "8TagsClustering",164    "SICK-R-PL",165    "CDSC-R",166    "STS22",167    "ArguAna-PL",168    "DBPedia-PL",169    "FiQA-PL",170    "HotpotQA-PL",171    "MSMARCO-PL",172    "NFCorpus-PL",173    "NQ-PL",174    "Quora-PL",175    "SCIDOCS-PL",176    "SciFact-PL",177    "TRECCOVID-PL",178]179 180MTEB_FR = [181    "AmazonReviewsClassification",182    "MasakhaNEWSClassification",183    "MassiveIntentClassification",184    "MassiveScenarioClassification",185    "MTOPDomainClassification",186    "MTOPIntentClassification",187    "OpusparcusPC",188    "PawsX",189    "AlloProfClusteringP2P",190    "AlloProfClusteringS2S",191    "HALClusteringS2S",192    "MasakhaNEWSClusteringP2P",193    "MasakhaNEWSClusteringS2S",194    "MLSUMClusteringP2P",195    "MLSUMClusteringS2S",196    "SyntecReranking",197    "AlloprofReranking",198    "AlloprofRetrieval",199    "BSARDRetrieval",200    "SyntecRetrieval",201    "XPQARetrieval",202    "MintakaRetrieval",203    "SummEvalFr",204    "STSBenchmarkMultilingualSTS",205    "STS22",206    "SICKFr",207]208 209logging.basicConfig(210    level=logging.INFO, format="%(asctime)s - %(levelname)s - %(name)s : %(message)s"211)212 213logger = logging.getLogger("eval_mteb_qwen.py")214 215 216def get_detailed_instruct(task_description: str) -> str:217    if not task_description:218        return ""219 220    return "Instruct: {}\nQuery: ".format(task_description)221 222 223def get_task_def_by_task_name_and_type(224    task_name: str,225    task_type: str,226    default_instruct="Given a web search query, retrieve relevant passages that answer the query",227) -> str:228    if task_type in ["STS"]:229        return "Retrieve semantically similar text"230 231    if task_type in ["Summarization"]:232        return "Given a news summary, retrieve other semantically similar summaries"233 234    if task_type in ["BitextMining"]:235        return "Retrieve parallel sentences"236 237    if task_type in ["Classification"]:238        task_name_to_instruct: Dict[str, str] = {239            "AmazonCounterfactualClassification": "Classify a given Amazon customer review text as either counterfactual or not-counterfactual",240            "AmazonPolarityClassification": "Classify Amazon reviews into positive or negative sentiment",241            "AmazonReviewsClassification": "Classify the given Amazon review into its appropriate rating category",242            "Banking77Classification": "Given a online banking query, find the corresponding intents",243            "EmotionClassification": "Classify the emotion expressed in the given Twitter message into one of the six emotions: anger, fear, joy, love, sadness, and surprise",244            "ImdbClassification": "Classify the sentiment expressed in the given movie review text from the IMDB dataset",245            "MassiveIntentClassification": "Given a user utterance as query, find the user intents",246            "MassiveScenarioClassification": "Given a user utterance as query, find the user scenarios",247            "MTOPDomainClassification": "Classify the intent domain of the given utterance in task-oriented conversation",248            "MTOPIntentClassification": "Classify the intent of the given utterance in task-oriented conversation",249            "ToxicConversationsClassification": "Classify the given comments as either toxic or not toxic",250            "TweetSentimentExtractionClassification": "Classify the sentiment of a given tweet as either positive, negative, or neutral",251            # C-MTEB eval instructions252            "TNews": "Classify the fine-grained category of the given news title",253            "IFlyTek": "Given an App description text, find the appropriate fine-grained category",254            "MultilingualSentiment": "Classify sentiment of the customer review into positive, neutral, or negative",255            "JDReview": "Classify the customer review for iPhone on e-commerce platform into positive or negative",256            "OnlineShopping": "Classify the customer review for online shopping into positive or negative",257            "Waimai": "Classify the customer review from a food takeaway platform into positive or negative",258            # MTEB-pl eval instructions259            "CBD": "Classify the sentiment of polish tweet reviews",260            "PolEmo2.0-IN": "Classify the sentiment of in-domain (medicine and hotels) online reviews",261            "PolEmo2.0-OUT": "Classify the sentiment of out-of-domain (products and school) online reviews",262            "AllegroReviews": "Classify the sentiment of reviews from e-commerce marketplace Allegro",263            "PAC": 'Classify the sentence into one of the two types: "BEZPIECZNE_POSTANOWIENIE_UMOWNE" and "KLAUZULA_ABUZYWNA"',264        }265        return task_name_to_instruct[task_name]266 267    if task_type in ["Clustering"]:268        task_name_to_instruct: Dict[str, str] = {269            "ArxivClusteringP2P": "Identify the main and secondary category of Arxiv papers based on the titles and abstracts",270            "ArxivClusteringS2S": "Identify the main and secondary category of Arxiv papers based on the titles",271            "BiorxivClusteringP2P": "Identify the main category of Biorxiv papers based on the titles and abstracts",272            "BiorxivClusteringS2S": "Identify the main category of Biorxiv papers based on the titles",273            "MedrxivClusteringP2P": "Identify the main category of Medrxiv papers based on the titles and abstracts",274            "MedrxivClusteringS2S": "Identify the main category of Medrxiv papers based on the titles",275            "RedditClustering": "Identify the topic or theme of Reddit posts based on the titles",276            "RedditClusteringP2P": "Identify the topic or theme of Reddit posts based on the titles and posts",277            "StackExchangeClustering": "Identify the topic or theme of StackExchange posts based on the titles",278            "StackExchangeClusteringP2P": "Identify the topic or theme of StackExchange posts based on the given paragraphs",279            "TwentyNewsgroupsClustering": "Identify the topic or theme of the given news articles",280            # C-MTEB eval instructions281            "CLSClusteringS2S": "Identify the main category of scholar papers based on the titles",282            "CLSClusteringP2P": "Identify the main category of scholar papers based on the titles and abstracts",283            "ThuNewsClusteringS2S": "Identify the topic or theme of the given news articles based on the titles",284            "ThuNewsClusteringP2P": "Identify the topic or theme of the given news articles based on the titles and contents",285            # MTEB-fr eval instructions286            "AlloProfClusteringP2P": "Identify the main category of Allo Prof document based on the titles and descriptions",287            "AlloProfClusteringS2S": "Identify the main category of Allo Prof document based on the titles",288            "HALClusteringS2S": "Identify the main category of academic passage based on the titles and contents",289            "MasakhaNEWSClusteringP2P": "Identify the topic or theme of the given news articles based on the titles and contents",290            "MasakhaNEWSClusteringS2S": "Identify the topic or theme of the given news articles based on the titles",291            "MLSUMClusteringP2P": "Identify the topic or theme of the given articles based on the titles and contents",292            "MLSUMClusteringS2S": "Identify the topic or theme of the given articles based on the titles",293            # MTEB-pl eval instructions294            "8TagsClustering": "Identify of headlines from social media posts in Polish  into 8 categories: film, history, food, medicine, motorization, work, sport and technology",295        }296        return task_name_to_instruct[task_name]297 298    if task_type in ["Reranking", "PairClassification"]:299        task_name_to_instruct: Dict[str, str] = {300            "AskUbuntuDupQuestions": "Retrieve duplicate questions from AskUbuntu forum",301            "MindSmallReranking": "Retrieve relevant news articles based on user browsing history",302            "SciDocsRR": "Given a title of a scientific paper, retrieve the titles of other relevant papers",303            "StackOverflowDupQuestions": "Retrieve duplicate questions from StackOverflow forum",304            "SprintDuplicateQuestions": "Retrieve duplicate questions from Sprint forum",305            "TwitterSemEval2015": "Retrieve tweets that are semantically similar to the given tweet",306            "TwitterURLCorpus": "Retrieve tweets that are semantically similar to the given tweet",307            # C-MTEB eval instructions308            "T2Reranking": "Given a Chinese search query, retrieve web passages that answer the question",309            "MmarcoReranking": "Given a Chinese search query, retrieve web passages that answer the question",310            "CMedQAv1": "Given a Chinese community medical question, retrieve replies that best answer the question",311            "CMedQAv2": "Given a Chinese community medical question, retrieve replies that best answer the question",312            "Ocnli": "Retrieve semantically similar text.",313            "Cmnli": "Retrieve semantically similar text.",314            # MTEB-fr eval instructions315            "AlloprofReranking": "Given a question, retrieve passages that answer the question",316            "OpusparcusPC": "Retrieve semantically similar text",317            "PawsX": "Retrieve semantically similar text",318            "SyntecReranking": "Given a question, retrieve passages that answer the question",319            # MTEB-pl eval instructions320            "SICK-E-PL": "Retrieve semantically similar text",321            "PPC": "Retrieve semantically similar text",322            "CDSC-E": "Retrieve semantically similar text",323            "PSC": "Retrieve semantically similar text",324        }325        return task_name_to_instruct[task_name]326 327    if task_type in ["Retrieval"]:328        if task_name.lower().startswith("cqadupstack"):329            return "Given a question, retrieve detailed question descriptions from Stackexchange that are duplicates to the given question"330 331        task_name_to_instruct: Dict[str, str] = {332            "ArguAna": "Given a claim, find documents that refute the claim",333            "ClimateFEVER": "Given a claim about climate change, retrieve documents that support or refute the claim",334            "DBPedia": "Given a query, retrieve relevant entity descriptions from DBPedia",335            "FEVER": "Given a claim, retrieve documents that support or refute the claim",336            "FiQA2018": "Given a financial question, retrieve user replies that best answer the question",337            "HotpotQA": "Given a multi-hop question, retrieve documents that can help answer the question",338            "MSMARCO": "Given a web search query, retrieve relevant passages that answer the query",339            "NFCorpus": "Given a question, retrieve relevant documents that best answer the question",340            "NQ": "Given a question, retrieve Wikipedia passages that answer the question",341            "QuoraRetrieval": "Given a question, retrieve questions that are semantically equivalent to the given question",342            "SCIDOCS": "Given a scientific paper title, retrieve paper abstracts that are cited by the given paper",343            "SciFact": "Given a scientific claim, retrieve documents that support or refute the claim",344            "Touche2020": "Given a question, retrieve detailed and persuasive arguments that answer the question",345            "TRECCOVID": "Given a query on COVID-19, retrieve documents that answer the query",346            # C-MTEB eval instructions347            "T2Retrieval": "Given a Chinese search query, retrieve web passages that answer the question",348            "MMarcoRetrieval": "Given a web search query, retrieve relevant passages that answer the query",349            "DuRetrieval": "Given a Chinese search query, retrieve web passages that answer the question",350            "CovidRetrieval": "Given a question on COVID-19, retrieve news articles that answer the question",351            "CmedqaRetrieval": "Given a Chinese community medical question, retrieve replies that best answer the question",352            "EcomRetrieval": "Given a user query from an e-commerce website, retrieve description sentences of relevant products",353            "MedicalRetrieval": "Given a medical question, retrieve user replies that best answer the question",354            "VideoRetrieval": "Given a video search query, retrieve the titles of relevant videos",355            # MTEB-fr eval instructions356            "AlloprofRetrieval": "Given a question, retrieve passages that answer the question",357            "BSARDRetrieval": "Given a question, retrieve passages that answer the question",358            "SyntecRetrieval": "Given a question, retrieve passages that answer the question",359            "XPQARetrieval": "Given a question, retrieve passages that answer the question",360            "MintakaRetrieval": "Given a question, retrieve passages that answer the question",361            # MTEB-pl eval instructions362            "ArguAna-PL": "Given a claim, find documents that refute the claim",363            "DBPedia-PL": "Given a query, retrieve relevant entity descriptions from DBPedia",364            "FiQA-PL": "Given a financial question, retrieve user replies that best answer the question",365            "HotpotQA-PL": "Given a multi-hop question, retrieve documents that can help answer the question",366            "MSMARCO-PL": "Given a web search query, retrieve relevant passages that answer the query",367            "NFCorpus-PL": "Given a question, retrieve relevant documents that best answer the question",368            "NQ-PL": "Given a question, retrieve Wikipedia passages that answer the question",369            "Quora-PL": "Given a question, retrieve questions that are semantically equivalent to the given question",370            "SCIDOCS-PL": "Given a scientific paper title, retrieve paper abstracts that are cited by the given paper",371            "SciFact-PL": "Given a scientific claim, retrieve documents that support or refute the claim",372            "TRECCOVID-PL": "Given a query on COVID-19, retrieve documents that answer the query",373        }374 375        # add lower case keys to match some beir names376        task_name_to_instruct.update({k.lower(): v for k, v in task_name_to_instruct.items()})377        # other cases where lower case match still doesn't work378        task_name_to_instruct["trec-covid"] = task_name_to_instruct["TRECCOVID"]379        task_name_to_instruct["climate-fever"] = task_name_to_instruct["ClimateFEVER"]380        task_name_to_instruct["dbpedia-entity"] = task_name_to_instruct["DBPedia"]381        task_name_to_instruct["webis-touche2020"] = task_name_to_instruct["Touche2020"]382        task_name_to_instruct["fiqa"] = task_name_to_instruct["FiQA2018"]383        task_name_to_instruct["quora"] = task_name_to_instruct["QuoraRetrieval"]384 385        # for miracl evaluation386        task_name_to_instruct["miracl"] = (387            "Given a question, retrieve Wikipedia passages that answer the question"388        )389 390        return task_name_to_instruct[task_name]391    logging.warning(392        f"No instruction config for task {task_name} with type {task_type}, use default instruction."393    )394    return default_instruct395 396 397class Encoder(torch.nn.Module):398    def __init__(self, name_or_path: str, pooling: str):399        super().__init__()400        self.model = AutoModel.from_pretrained(name_or_path, trust_remote_code=True)401        self.model = self.model.half()402        self.model.eval()403        self.pooling = pooling404 405    def forward(self, **features) -> torch.Tensor:406        output = self.model(**features, output_hidden_states=True, return_dict=True)407        hidden_state = output.hidden_states[-1]408        embeddings = self.pooler(hidden_state, **features)409        return embeddings410 411    def pooler(412        self, hidden_state: torch.Tensor, attention_mask: torch.Tensor, **kwargs413    ) -> torch.Tensor:414        if attention_mask.ndim == 2:415            mask_expanded = attention_mask.unsqueeze(-1).expand(hidden_state.size())416        elif attention_mask.ndim == 3:417            mask_expanded = attention_mask418        else:419            raise RuntimeError(f"Unexpected {attention_mask.ndim=}")420 421        hidden_state = hidden_state * mask_expanded422 423        if self.pooling == "first":424            pooled_output = hidden_state[:, 0]425 426        elif self.pooling == "last":427            left_padding = attention_mask[:, -1].sum() == attention_mask.shape[0]428            if left_padding:429                return hidden_state[:, -1]430            else:431                sequence_lengths = attention_mask.sum(dim=1) - 1432                batch_size = hidden_state.shape[0]433            return hidden_state[434                torch.arange(batch_size, device=hidden_state.device), sequence_lengths435            ]436        elif self.pooling == "mean":437            # TODO: weight438            lengths = mask_expanded.sum(1).clamp(min=1e-9)439            pooled_output = hidden_state.sum(dim=1) / lengths440 441        elif self.pooling == "weightedmean":442            input_mask_expanded = attention_mask.unsqueeze(-1).expand(hidden_state.size()).float()443            # hidden_state shape: bs, seq, hidden_dim444            weights = (445                torch.arange(start=1, end=hidden_state.shape[1] + 1)446                .unsqueeze(0)447                .unsqueeze(-1)448                .expand(hidden_state.size())449                .float()450                .to(hidden_state.device)451            )452            assert weights.shape == hidden_state.shape == input_mask_expanded.shape453            input_mask_expanded = input_mask_expanded * weights454 455            sum_embeddings = torch.sum(hidden_state * input_mask_expanded, 1)456            sum_mask = input_mask_expanded.sum(1)457            sum_mask = torch.clamp(sum_mask, min=1e-9)458            pooled_output = sum_embeddings / sum_mask459 460        else:461            raise ValueError(f"Wrong pooler mode : {self.pooling}")462        return pooled_output463 464 465class Wrapper:466    def __init__(467        self,468        tokenizer,469        encoder: Encoder,470        batch_size: int,471        max_seq_len: int = 512,472        normalize_embeddings: bool = False,473        default_query: bool = False,474        force_default: bool = False,475        sep: str = " ",476        mp_tensor_to_cuda: bool = False,477        instruction: Optional[str] = None,478    ):479        self.tokenizer = tokenizer480        self.model = encoder481        self.batch_size = batch_size482        self.max_seq_len = max_seq_len483        self.pool: Optional[dict] = None484        self.normalize_embeddings = normalize_embeddings485        self.mp_tensor_to_cuda = mp_tensor_to_cuda486        self._target_device = torch.device("cuda" if torch.cuda.is_available() else "cpu")487        self.eod_id = self.tokenizer.convert_tokens_to_ids("<|endoftext|>")488        self.instruction = instruction489        self.default_query = default_query490        self.sep = sep491        self.force_default = force_default492        if self.tokenizer.padding_side != "right":493            logger.warning(494                f"Change tokenizer.padding_side from {self.tokenizer.padding_side} to right"495            )496            self.tokenizer.padding_side = "right"497        if self.tokenizer.pad_token is None:498            logger.warning(f"Set tokenizer.pad_token as eos_token {self.tokenizer.eos_token}")499            self.tokenizer.pad_token = "<|endoftext|>"500 501    def start(self, target_devices: Optional[List[str]] = None):502        """503        Starts multi process to process the encoding with several, independent processes.504        This method is recommended if you want to encode on multiple GPUs. It is advised505        to start only one process per GPU. This method works together with encode_multi_process506 507        :param target_devices: PyTorch target devices, e.g. cuda:0, cuda:1... If None, all available CUDA devices will be used508        :return: Returns a dict with the target processes, an input queue and and output queue.509        """510        if target_devices is None:511            if torch.cuda.is_available():512                target_devices = ["cuda:{}".format(i) for i in range(torch.cuda.device_count())]513            else:514                logger.info("CUDA is not available. Start 4 CPU worker")515                target_devices = ["cpu"] * 4516 517        logger.info(518            "Start multi-process pool on devices: {}".format(", ".join(map(str, target_devices)))519        )520        print("multi instruction", self.instruction)521        ctx = mp.get_context("spawn")522        input_queue = ctx.Queue()523        output_queue = ctx.Queue()524        processes = []525 526        for cuda_id in target_devices:527            p = ctx.Process(528                target=self._encode_multi_process_worker,529                args=(cuda_id, self, input_queue, output_queue),530                daemon=True,531            )532            p.start()533            processes.append(p)534 535        self.pool = {"input": input_queue, "output": output_queue, "processes": processes}536 537    def stop(self):538        """539        Stops all processes started with start_multi_process_pool540        """541        for p in self.pool["processes"]:542            p.terminate()543 544        for p in self.pool["processes"]:545            p.join()546            p.close()547 548        self.pool["input"].close()549        self.pool["output"].close()550 551    @staticmethod552    def _encode_multi_process_worker(target_device: str, model, input_queue, results_queue):553        """554        Internal working process to encode sentences in multi-process setup555        """556        while True:557            try:558                id, sentences, kwargs = input_queue.get()559                kwargs.update(device=target_device, show_progress_bar=False, convert_to_numpy=True)560                embeddings = model._encode(sentences, **kwargs)561                results_queue.put([id, embeddings])562            except queue.Empty:563                break564 565    def encode_multi_process(self, sentences: List[str], **kwargs):566        """567        This method allows to run encode() on multiple GPUs. The sentences are chunked into smaller packages568        and sent to individual processes, which encode these on the different GPUs. This method is only suitable569        for encoding large sets of sentences570 571        :param sentences: List of sentences572        :param pool: A pool of workers started with SentenceTransformer.start_multi_process_pool573        :param chunk_size: Sentences are chunked and sent to the individual processes. If none, it determine a sensible size.574        :param kwargs: other keyword arguments for model.encode() such as batch_size575        :return: Numpy matrix with all embeddings576        """577        part_size = math.ceil(len(sentences) / len(self.pool["processes"]))578        chunk_size = part_size if part_size < 3200 else 3200  # for retrieval chunk 50000579 580        logger.debug(581            f"Chunk data into {math.ceil(len(sentences) / chunk_size)} packages of size {chunk_size}"582        )583 584        input_queue = self.pool["input"]585        last_chunk_id = 0586        chunk = []587 588        for sentence in sentences:589            chunk.append(sentence)590            if len(chunk) >= chunk_size:591                input_queue.put([last_chunk_id, chunk, kwargs])592                last_chunk_id += 1593                chunk = []594 595        if len(chunk) > 0:596            input_queue.put([last_chunk_id, chunk, kwargs])597            last_chunk_id += 1598 599        output_queue = self.pool["output"]600        results_list = sorted(601            [output_queue.get() for _ in range(last_chunk_id)], key=lambda x: x[0]602        )603        embeddings = np.concatenate([result[1] for result in results_list])604        return embeddings605 606    @staticmethod607    def batch_to_device(batch, target_device):608        """609        send a pytorch batch to a device (CPU/GPU)610        """611        for key in batch:612            if isinstance(batch[key], torch.Tensor):613                batch[key] = batch[key].to(target_device)614        return batch615 616    def _text_length(self, text: Union[List[int], List[List[int]]]):617        """618        Help function to get the length for the input text. Text can be either619        a list of ints (which means a single text as input), or a tuple of list of ints620        (representing several text inputs to the model).621        """622 623        if isinstance(text, dict):  # {key: value} case624            return len(next(iter(text.values())))625        elif not hasattr(text, "__len__"):  # Object has no len() method626            return 1627        elif len(text) == 0 or isinstance(text[0], int):  # Empty string or list of ints628            return len(text)629        else:630            return sum([len(t) for t in text])  # Sum of length of individual strings631 632    def _tokenize(self, sentences: List[str], is_query: bool):633        batch_dict = self.tokenizer(634            sentences,635            max_length=self.max_seq_len - 1,636            return_attention_mask=False,637            padding=False,638            truncation=True,639        )640        batch_dict["is_causal"] = False641        return batch_dict642 643    def _encode(644        self,645        sentences: List[str],646        is_query: bool,647        convert_to_numpy: bool = True,648        convert_to_tensor: bool = False,649        device: Optional[str] = None,650        show_progress_bar: bool = True,651        **kwargs,652    ):653        """654        Computes sentence embeddings655 656        :param sentences: the sentences to embed657        :param batch_size: the batch size used for the computation658        :param show_progress_bar: Output a progress bar when encode sentences659        :param output_value:  Default sentence_embedding, to get sentence embeddings. Can be set to token_embeddings to get wordpiece token embeddings. Set to None, to get all output values660        :param convert_to_numpy: If true, the output is a list of numpy vectors. Else, it is a list of pytorch tensors.661        :param convert_to_tensor: If true, you get one large tensor as return. Overwrites any setting from convert_to_numpy662        :param device: Which torch.device to use for the computation663        :param normalize_embeddings: If set to true, returned vectors will have length 1. In that case, the faster dot-product (util.dot_score) instead of cosine similarity can be used.664 665        :return:666           By default, a list of tensors is returned. If convert_to_tensor, a stacked tensor is returned. If convert_to_numpy, a numpy matrix is returned.667        """668        self.model.eval()669 670        if convert_to_tensor:671            convert_to_numpy = False672 673        input_was_string = False674        if isinstance(sentences, str) or not hasattr(675            sentences, "__len__"676        ):  # Cast an individual sentence to a list with length 1677            sentences = [sentences]678            input_was_string = True679 680        if device is None:681            device = self._target_device682 683        self.model.to(device)684 685        all_embeddings = []686        length_sorted_idx = np.argsort([-self._text_length(s) for s in sentences])687        sentences_sorted = [sentences[idx] for idx in length_sorted_idx]688 689        for start_index in trange(690            0, len(sentences), self.batch_size, desc="Batches", disable=not show_progress_bar691        ):692            sentences_batch = sentences_sorted[start_index : start_index + self.batch_size]693            features = self._tokenize(sentences_batch, is_query)694            features = self.batch_to_device(features, device)695 696            with torch.no_grad():697                embeddings = self.model(**features)698 699                if self.normalize_embeddings:700                    embeddings = torch.nn.functional.normalize(embeddings, p=2, dim=1)701 702                # fixes for #522 and #487 to avoid oom problems on gpu with large datasets703                if convert_to_numpy:704                    embeddings = embeddings.cpu()705 706                all_embeddings.extend(embeddings)707 708        all_embeddings = [all_embeddings[idx] for idx in np.argsort(length_sorted_idx)]709 710        if convert_to_tensor:711            all_embeddings = torch.stack(all_embeddings)712        elif convert_to_numpy:713            # all_embeddings = np.asarray([emb.numpy() for emb in all_embeddings])714            all_embeddings = np.asarray([emb.to(torch.float).numpy() for emb in all_embeddings])715        if input_was_string:716            all_embeddings = all_embeddings[0]717 718        return all_embeddings719 720    def encode(721        self,722        sentences: List[str],723        is_query: Optional[bool] = None,724        convert_to_tensor: bool = False,725        **kwargs,726    ):727        is_query = self.default_query if is_query is None else is_query728        if is_query and self.instruction:729            sentences = [self.instruction + sent for sent in sentences]730        kwargs.update(is_query=is_query)731        if self.pool is not None:732            kwargs.update(show_progress_bar=False)733            embeddings = self.encode_multi_process(sentences, **kwargs)734            if convert_to_tensor:735                embeddings = torch.from_numpy(embeddings)736                if self.mp_tensor_to_cuda and torch.cuda.is_available():737                    embeddings = embeddings.to(torch.device("cuda"))  # default 0-th gpu738            return embeddings739 740        return self._encode(sentences, convert_to_tensor=convert_to_tensor, **kwargs)741 742    def encode_queries(self, queries: List[str], **kwargs):743        is_query = self.default_query if self.force_default else True744        return self.encode(queries, is_query=is_query, **kwargs)745 746    def encode_corpus(self, corpus: List[Dict[str, str]], **kwargs):747        # borrowed from mteb.abstasks.AbsTaskRetrieval.DRESModel748        if type(corpus) is dict:749            sentences = [750                (corpus["title"][i] + self.sep + corpus["text"][i]).strip()751                if "title" in corpus752                else corpus["text"][i].strip()753                for i in range(len(corpus["text"]))754            ]755        elif isinstance(corpus[0], dict):756            sentences = [757                (doc["title"] + self.sep + doc["text"]).strip()758                if "title" in doc759                else doc["text"].strip()760                for doc in corpus761            ]762        else:763            sentences = corpus764        is_query = self.default_query if self.force_default else False765        return self.encode(sentences, is_query=is_query, **kwargs)766 767 768def main(args):769    tokenizer = AutoTokenizer.from_pretrained(args.model, trust_remote_code=True)770    encoder = Encoder(args.model, args.pooling)771    default_query = args.default_type == "query"772    model = Wrapper(773        tokenizer,774        encoder,775        batch_size=args.batch_size,776        max_seq_len=args.max_seq_len,777        normalize_embeddings=args.norm,778        default_query=default_query,779    )780    sym_retrievals = ["QuoraRetrieval", "ArguAna", "CQADupstack"]781    if args.task == "mteb":782        task_names = MTEB_TASK_LIST783        lang = ["en"]784    elif args.task == "cmteb":785        task_names = CMTEB_TASK_LIST786        lang = ["zh", "zh-CN"]787    elif args.task == "mteb-fr":788        task_names = MTEB_FR789        lang = ["fr"]790    elif args.task == "mteb-pl":791        task_names = MTEB_PL792        lang = ["pl"]793    else:794        task_names = [args.task]795        lang = ["en", "zh", "zh-CN", "pl", "fr"]796    for task in task_names:797        evaluation = MTEB(tasks=[task], task_langs=lang)798        task_cls = evaluation.tasks[0]799        task_name: str = task_cls.metadata_dict["name"]800        task_type: str = task_cls.metadata_dict["type"]801        instruction = get_task_def_by_task_name_and_type(task_name, task_type)802        model.instruction = get_detailed_instruct(instruction)803        if task == "MSMARCO":804            eval_splits = ["dev"]805        elif task in CMTEB_TASK_LIST:806            eval_splits = task_cls.metadata_dict["eval_splits"]807        else:808            eval_splits = ["test"]809        sym = False810        for name in sym_retrievals:811            if task.startswith(name):812                sym = True813                break814            else:815                sym = False816        if sym:817            logger.info(818                f"Switch to symmetric mode for {task}, all as {'query' if default_query else 'doc'}."819            )820            model.force_default = True821        evaluation.run(model, output_folder=args.output_dir, eval_splits=eval_splits)822 823        if sym:824            logger.info(f"Switch back.")825            model.force_default = force_default_ori826        print("\n")827 828 829if __name__ == "__main__":830    _PARSER = argparse.ArgumentParser()831    _PARSER.add_argument("-m", "--model", type=str, default=None)832    _PARSER.add_argument("--pooling", type=str, default="last")833    _PARSER.add_argument("--output_dir", type=str, default=None)834    _PARSER.add_argument("--default_type", type=str, default="query")835    _PARSER.add_argument("--max_seq_len", type=int, default=512)836    _PARSER.add_argument("-b", "--batch_size", type=int, default=32)837    _PARSER.add_argument(838        "-t",839        "--task",840        type=str,841        default=None,  # None for running default tasks842    )843    _PARSER.add_argument("--norm", action="store_true")844    _ARGS = _PARSER.parse_args()845    main(_ARGS)846