Team Ai
Apppublic

ramouch/Tunisian-Encoder

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
encoder.py67 linesDownload Raw Back to root
1import os2import torch3import numpy as np4from sentence_transformers import SentenceTransformer5from typing import List, Union6 7# Paths to the models8ORIGINAL_MODEL_PATH = "original_model.pt"9OPTIMIZED_MODEL_PATH = "optimized_model.pt"10 11# Check quantization engine12if 'qnnpack' in torch.backends.quantized.supported_engines:13    torch.backends.quantized.engine = 'qnnpack'14elif 'fbgemm' in torch.backends.quantized.supported_engines:15    torch.backends.quantized.engine = 'fbgemm'16else:17    print("โŒ No quantized engine found. Quantized models cannot be loaded.")18 19class FineTunedArabicBERTEncoder:20    def __init__(self, use_optimized: bool = False):21        """22        Initialize the encoder with either the original or the optimized model.23        24        Parameters:25        - use_optimized (bool): If True, load the optimized model, else load the original.26        """27        model_path = OPTIMIZED_MODEL_PATH if use_optimized else ORIGINAL_MODEL_PATH28        29        if not os.path.exists(model_path):30            raise FileNotFoundError(f"โŒ Model not found at path: {model_path}")31        32        print(f"๐Ÿš€ Loading {'optimized' if use_optimized else 'original'} model from {model_path}")33        34        # โœ… **Force to CPU**35        self.device = torch.device("cpu")36        try:37            # Explicitly load on CPU38            self.model = torch.load(model_path, map_location=self.device, weights_only=False)39        except Exception as e:40            print(f"โš ๏ธ Failed to load with `map_location=cpu`: {e}")41            print("๐Ÿ”„ Attempting safe globals registration...")42            torch.serialization.add_safe_globals([SentenceTransformer])43            self.model = torch.load(model_path, map_location=self.device)44        45        self.model.eval()  # Set to evaluation mode46        print("โœ… Model loaded successfully on CPU.")47 48    def encode(self, sentences: Union[str, List[str]]) -> np.ndarray:49        """50        Encode a sentence or a list of sentences into embeddings.51 52        Parameters:53        - sentences (str or List[str]): The sentences to encode.54 55        Returns:56        - np.ndarray: The sentence embeddings.57        """58        if isinstance(sentences, str):59            sentences = [sentences]60        61        with torch.no_grad():62            embeddings = self.model.encode(sentences, convert_to_numpy=True)63        64        return embeddings65 66    def __call__(self, sentences: Union[str, List[str]]) -> np.ndarray:67        return self.encode(sentences)