ramouch/Tunisian-Encoder
0
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)