BridgeAI-Lab/Sem-nCG
1
1import abc2from typing import List, Union3 4from numpy.typing import NDArray5from sentence_transformers import SentenceTransformer6 7from .type_aliases import ENCODER_DEVICE_TYPE8 9 10class Encoder(abc.ABC):11 @abc.abstractmethod12 def encode(self, prediction: List[str]) -> NDArray:13 """14 Abstract method to encode a list of sentences into sentence embeddings.15 16 Args:17 prediction (List[str]): List of sentences to encode.18 19 Returns:20 NDArray: Array of sentence embeddings with shape (num_sentences, embedding_dim).21 22 Raises:23 NotImplementedError: If the method is not implemented in the subclass.24 """25 raise NotImplementedError("Method 'encode' must be implemented in subclass.")26 27 28class SBertEncoder(Encoder):29 def __init__(self, model: SentenceTransformer, device: ENCODER_DEVICE_TYPE, batch_size: int, verbose: bool):30 """31 Initialize SBertEncoder instance.32 33 Args:34 model (SentenceTransformer): The Sentence Transformer model instance to use for encoding.35 device (Union[str, int, List[Union[str, int]]]): Device specification for encoding36 batch_size (int): Batch size for encoding.37 verbose (bool): Whether to print verbose information during encoding.38 """39 self.model = model40 self.device = device41 self.batch_size = batch_size42 self.verbose = verbose43 44 def encode(self, prediction: List[str]) -> NDArray:45 """46 Encode a list of sentences into sentence embeddings.47 48 Args:49 prediction (List[str]): List of sentences to encode.50 51 Returns:52 NDArray: Array of sentence embeddings with shape (num_sentences, embedding_dim).53 """54 55 # SBert output is always Batch x Dim56 if isinstance(self.device, list):57 # Use multiprocess encoding for list of devices58 pool = self.model.start_multi_process_pool(target_devices=self.device)59 embeddings = self.model.encode_multi_process(prediction, pool=pool, batch_size=self.batch_size)60 self.model.stop_multi_process_pool(pool)61 else:62 # Single device encoding63 embeddings = self.model.encode(64 prediction,65 device=self.device,66 batch_size=self.batch_size,67 )68 69 return embeddings70 71 72def get_encoder(73 sbert_model: SentenceTransformer,74 device: ENCODER_DEVICE_TYPE,75 batch_size: int,76 verbose: bool,77) -> Encoder:78 """79 Get an instance of SBertEncoder using the provided parameters.80 81 Args:82 sbert_model (SentenceTransformer): An instance of SentenceTransformer model to use for encoding.83 device (Union[str, int, List[Union[str, int]]): Device specification for the encoder84 (e.g., "cuda", 0 for GPU, "cpu").85 batch_size (int): Batch size to use for encoding.86 verbose (bool): Whether to print verbose information during encoding.87 88 Returns:89 SBertEncoder: Instance of the selected encoder based on the model_name.90 91 Example:92 >>> model_name = "paraphrase-distilroberta-base-v1"93 >>> sbert_model = get_sbert_encoder(model_name)94 >>> device = get_gpu("cuda")95 >>> batch_size = 3296 >>> verbose = True97 >>> encoder = get_encoder(sbert_model, device, batch_size, verbose)98 """99 encoder = SBertEncoder(sbert_model, device, batch_size, verbose)100 return encoder101 102 103def get_sbert_encoder(model_name: str) -> SentenceTransformer:104 """105 Get an instance of SentenceTransformer encoder based on the specified model name.106 107 Args:108 model_name (str): Name of the model to instantiate. You can use any model on Huggingface/SentenceTransformer109 that is supported by SentenceTransformer.110 111 Returns:112 SentenceTransformer: Instance of the selected encoder based on the model_name.113 114 Raises:115 EnvironmentError: If an unsupported model_name is provided.116 RuntimeError: If there's an issue during instantiation of the encoder.117 """118 119 try:120 encoder = SentenceTransformer(model_name, trust_remote_code=True)121 except EnvironmentError as err:122 raise EnvironmentError(str(err)) from None123 except Exception as err:124 raise RuntimeError(str(err)) from None125 126 return encoder127 