| """ |
| Embedding layer. |
| """ |
| import logging |
| from typing import List |
| import numpy as np |
|
|
| from sentence_transformers import SentenceTransformer |
| from config import settings |
|
|
| logger = logging.getLogger(__name__) |
|
|
| class EmbeddingLayer: |
|
|
| @staticmethod |
| def create_embedder() -> 'EmbeddingLayer': |
| return EmbeddingLayer() |
|
|
| def __init__(self): |
| self.model_name = settings.EMBEDDING_MODEL |
| self.dimension = settings.EMBEDDING_DIM |
| self.model = None |
|
|
| def _load_model(self): |
| if self.model is None: |
| logger.info(f"Loading embedding model: {self.model_name}") |
| self.model = SentenceTransformer( |
| self.model_name, |
| device="cpu" |
| ) |
|
|
| def embed_texts(self, texts: List[str]) -> np.ndarray: |
| |
| if not texts: |
| raise ValueError("No texts provided for embedding") |
|
|
| self._load_model() |
|
|
| embeddings = self.model.encode( |
| texts, |
| batch_size=settings.BATCH_SIZE, |
| normalize_embeddings=True, |
| show_progress_bar=True, |
| ) |
|
|
| embeddings = np.asarray(embeddings, dtype=np.float32) |
|
|
| if embeddings.shape[1] != self.dimension: |
| raise ValueError( |
| f"Embedding dimension mismatch: got {embeddings.shape[1]}, expected {self.dimension}" |
| ) |
|
|
| logger.info(f"Generated embeddings: shape={embeddings.shape}") |
| return embeddings |
|
|
| def embed_query(self, query: str) -> np.ndarray: |
| |
| self._load_model() |
|
|
| vec = self.model.encode( |
| [query], |
| normalize_embeddings=True, |
| )[0].astype(np.float32) |
|
|
| if vec.shape[0] != self.dimension: |
| raise ValueError("Query embedding dimension mismatch") |
|
|
| return vec |
|
|