MMo4/csit-ned-chatbot
0
1from sentence_transformers import SentenceTransformer2from typing import List, Union3import numpy as np4import logging5from src.config import settings6 7logger = logging.getLogger(__name__)8 9class EmbeddingModel:10 """Wrapper for Sentence Transformers embedding model"""11 12 def __init__(self, model_name: str = None):13 self.model_name = model_name or settings.embedding_model14 self.model = None15 self._load_model()16 17 def _load_model(self):18 """Load the sentence transformer model"""19 try:20 logger.info(f"Loading embedding model: {self.model_name}")21 self.model = SentenceTransformer(self.model_name)22 logger.info("Embedding model loaded successfully")23 except Exception as e:24 logger.error(f"Failed to load embedding model: {e}")25 raise26 27 def encode(self, texts: Union[str, List[str]], **kwargs) -> np.ndarray:28 """Encode text(s) into embeddings"""29 if not self.model:30 raise RuntimeError("Embedding model not loaded")31 32 if isinstance(texts, str):33 texts = [texts]34 35 try:36 embeddings = self.model.encode(texts, **kwargs)37 logger.debug(f"Generated embeddings for {len(texts)} texts")38 return embeddings39 except Exception as e:40 logger.error(f"Failed to generate embeddings: {e}")41 raise42 43 def encode_single(self, text: str, **kwargs) -> np.ndarray:44 """Encode a single text into embedding"""45 embedding = self.encode([text], **kwargs)46 return embedding[0] if len(embedding) > 0 else np.array([])47 48 def get_embedding_dimension(self) -> int:49 """Get the dimension of embeddings produced by this model"""50 if not self.model:51 raise RuntimeError("Embedding model not loaded")52 return self.model.get_sentence_embedding_dimension()53 54 def similarity(self, embedding1: np.ndarray, embedding2: np.ndarray) -> float:55 """Calculate cosine similarity between two embeddings"""56 # Normalize embeddings57 embedding1_norm = embedding1 / np.linalg.norm(embedding1)58 embedding2_norm = embedding2 / np.linalg.norm(embedding2)59 60 # Calculate cosine similarity61 similarity = np.dot(embedding1_norm, embedding2_norm)62 return float(similarity)63 64 def batch_similarity(self, query_embedding: np.ndarray, 65 document_embeddings: List[np.ndarray]) -> List[float]:66 """Calculate similarities between query and multiple documents"""67 similarities = []68 for doc_embedding in document_embeddings:69 sim = self.similarity(query_embedding, doc_embedding)70 similarities.append(sim)71 return similarities72 73# Global instance74embedding_model = None75 76def get_embedding_model() -> EmbeddingModel:77 """Get or create the global embedding model instance"""78 global embedding_model79 if embedding_model is None:80 embedding_model = EmbeddingModel()81 return embedding_model82 83def encode_text(text: str) -> np.ndarray:84 """Convenience function to encode text using global model"""85 model = get_embedding_model()86 return model.encode_single(text)87 88def encode_texts(texts: List[str]) -> np.ndarray:89 """Convenience function to encode multiple texts using global model"""90 model = get_embedding_model()91 return model.encode(texts)