Kashish
Initial commit
345991e
Raw
History Blame Contribute Delete
1.8 kB
"""
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