Coverage for src/semware/services/embedding.py: 64%
69 statements
« prev ^ index » next coverage.py v7.10.6, created at 2025-09-09 02:16 -0700
« prev ^ index » next coverage.py v7.10.6, created at 2025-09-09 02:16 -0700
1"""Embedding generation service using Google's EmbeddingGemma model."""
4import numpy as np
5from loguru import logger
6from sentence_transformers import SentenceTransformer
8from ..config import settings
9from ..utils.tokenizer import tokenizer
12class EmbeddingService:
13 """Service for generating text embeddings using EmbeddingGemma."""
15 def __init__(self, model_name: str | None = None):
16 """Initialize the embedding service.
18 Args:
19 model_name: Name of the embedding model to use
20 """
21 self.model_name = model_name or settings.embedding_model_name
22 self.max_tokens_per_batch = settings.max_tokens_per_batch
23 self.embedding_dimension = settings.embedding_dimension
25 logger.info(f"Loading embedding model: {self.model_name}")
26 try:
27 self.model = SentenceTransformer(self.model_name)
28 logger.info("Embedding model loaded successfully")
29 except Exception as e:
30 logger.error(f"Failed to load embedding model: {e}")
31 raise
33 def generate_embedding(self, text: str) -> np.ndarray:
34 """Generate embedding for a single text.
36 Args:
37 text: Input text
39 Returns:
40 Normalized embedding vector
41 """
42 if not text.strip():
43 logger.warning("Empty text provided for embedding")
44 return np.zeros(self.embedding_dimension, dtype=np.float32)
46 # Tokenize and batch the text if necessary
47 batches = tokenizer.batch_text(text, self.max_tokens_per_batch)
49 if not batches:
50 logger.warning("No valid batches created from text")
51 return np.zeros(self.embedding_dimension, dtype=np.float32)
53 try:
54 # Generate embeddings for each batch
55 batch_embeddings = []
56 for batch in batches:
57 # Use encode_document for document embedding
58 embedding = self.model.encode(
59 batch, convert_to_numpy=True, normalize_embeddings=False
60 )
61 batch_embeddings.append(embedding)
63 # Combine embeddings if multiple batches
64 if len(batch_embeddings) == 1:
65 combined_embedding = batch_embeddings[0]
66 else:
67 # Average pooling for combining embeddings
68 combined_embedding = np.mean(batch_embeddings, axis=0)
70 # Normalize the final embedding for cosine similarity
71 norm = np.linalg.norm(combined_embedding)
72 if norm > 0:
73 combined_embedding = combined_embedding / norm
74 else:
75 logger.warning("Zero norm embedding, returning zeros")
76 combined_embedding = np.zeros_like(combined_embedding)
78 logger.debug(f"Generated embedding with {len(batches)} batches")
79 return combined_embedding.astype(np.float32)
81 except Exception as e:
82 logger.error(f"Failed to generate embedding: {e}")
83 return np.zeros(self.embedding_dimension, dtype=np.float32)
85 def generate_embeddings(self, texts: list[str]) -> list[np.ndarray]:
86 """Generate embeddings for multiple texts.
88 Args:
89 texts: List of input texts
91 Returns:
92 List of normalized embedding vectors
93 """
94 embeddings = []
95 for text in texts:
96 embedding = self.generate_embedding(text)
97 embeddings.append(embedding)
99 logger.info(f"Generated embeddings for {len(texts)} texts")
100 return embeddings
102 def generate_query_embedding(self, query: str) -> np.ndarray:
103 """Generate embedding for a query text.
105 This method is optimized for query embeddings and uses the same
106 batching strategy as document embeddings for consistency.
108 Args:
109 query: Query text
111 Returns:
112 Normalized query embedding vector
113 """
114 # Use the same logic as generate_embedding for consistency
115 return self.generate_embedding(query)
117 def compute_similarity(
118 self, embedding1: np.ndarray, embedding2: np.ndarray
119 ) -> float:
120 """Compute cosine similarity between two embeddings.
122 Args:
123 embedding1: First embedding vector
124 embedding2: Second embedding vector
126 Returns:
127 Similarity score between 0 and 1
128 """
129 try:
130 # Normalize vectors if not already normalized
131 norm1 = np.linalg.norm(embedding1)
132 norm2 = np.linalg.norm(embedding2)
134 if norm1 == 0 or norm2 == 0:
135 return 0.0
137 normalized1 = embedding1 / norm1
138 normalized2 = embedding2 / norm2
140 # Compute cosine similarity
141 similarity = np.dot(normalized1, normalized2)
143 # Clamp to [0, 1] range (cosine can be negative)
144 similarity = max(0.0, min(1.0, (similarity + 1.0) / 2.0))
146 return float(similarity)
148 except Exception as e:
149 logger.error(f"Failed to compute similarity: {e}")
150 return 0.0
152 def get_model_info(self) -> dict:
153 """Get information about the loaded model.
155 Returns:
156 Dictionary with model information
157 """
158 return {
159 "model_name": self.model_name,
160 "embedding_dimension": self.embedding_dimension,
161 "max_tokens_per_batch": self.max_tokens_per_batch,
162 "model_max_seq_length": getattr(self.model, "max_seq_length", "unknown"),
163 }
166# Global embedding service instance
167embedding_service = EmbeddingService()