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

1"""Embedding generation service using Google's EmbeddingGemma model.""" 

2 

3 

4import numpy as np 

5from loguru import logger 

6from sentence_transformers import SentenceTransformer 

7 

8from ..config import settings 

9from ..utils.tokenizer import tokenizer 

10 

11 

12class EmbeddingService: 

13 """Service for generating text embeddings using EmbeddingGemma.""" 

14 

15 def __init__(self, model_name: str | None = None): 

16 """Initialize the embedding service. 

17 

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 

24 

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 

32 

33 def generate_embedding(self, text: str) -> np.ndarray: 

34 """Generate embedding for a single text. 

35 

36 Args: 

37 text: Input text 

38 

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) 

45 

46 # Tokenize and batch the text if necessary 

47 batches = tokenizer.batch_text(text, self.max_tokens_per_batch) 

48 

49 if not batches: 

50 logger.warning("No valid batches created from text") 

51 return np.zeros(self.embedding_dimension, dtype=np.float32) 

52 

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) 

62 

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) 

69 

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) 

77 

78 logger.debug(f"Generated embedding with {len(batches)} batches") 

79 return combined_embedding.astype(np.float32) 

80 

81 except Exception as e: 

82 logger.error(f"Failed to generate embedding: {e}") 

83 return np.zeros(self.embedding_dimension, dtype=np.float32) 

84 

85 def generate_embeddings(self, texts: list[str]) -> list[np.ndarray]: 

86 """Generate embeddings for multiple texts. 

87 

88 Args: 

89 texts: List of input texts 

90 

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) 

98 

99 logger.info(f"Generated embeddings for {len(texts)} texts") 

100 return embeddings 

101 

102 def generate_query_embedding(self, query: str) -> np.ndarray: 

103 """Generate embedding for a query text. 

104 

105 This method is optimized for query embeddings and uses the same 

106 batching strategy as document embeddings for consistency. 

107 

108 Args: 

109 query: Query text 

110 

111 Returns: 

112 Normalized query embedding vector 

113 """ 

114 # Use the same logic as generate_embedding for consistency 

115 return self.generate_embedding(query) 

116 

117 def compute_similarity( 

118 self, embedding1: np.ndarray, embedding2: np.ndarray 

119 ) -> float: 

120 """Compute cosine similarity between two embeddings. 

121 

122 Args: 

123 embedding1: First embedding vector 

124 embedding2: Second embedding vector 

125 

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) 

133 

134 if norm1 == 0 or norm2 == 0: 

135 return 0.0 

136 

137 normalized1 = embedding1 / norm1 

138 normalized2 = embedding2 / norm2 

139 

140 # Compute cosine similarity 

141 similarity = np.dot(normalized1, normalized2) 

142 

143 # Clamp to [0, 1] range (cosine can be negative) 

144 similarity = max(0.0, min(1.0, (similarity + 1.0) / 2.0)) 

145 

146 return float(similarity) 

147 

148 except Exception as e: 

149 logger.error(f"Failed to compute similarity: {e}") 

150 return 0.0 

151 

152 def get_model_info(self) -> dict: 

153 """Get information about the loaded model. 

154 

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 } 

164 

165 

166# Global embedding service instance 

167embedding_service = EmbeddingService()