joshcd/MNLP_M3_document_encoder
082
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 