CoolFace
Apppublic

BridgeAI-Lab/Sem-nCG

sourceHugging Faceupdated 2y agoView on Hugging Face
1likes
encoder_models.py127 linesDownload Raw Back to root
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