CoolFace
Apppublic

MMo4/csit-ned-chatbot

sourceHugging Faceupdated 10mo agoView on Hugging Face
0likes
embeddings.py91 linesDownload Raw Back to rag
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)