Coverage for src / quber / db / embeddings.py: 79%

86 statements  

« prev     ^ index     » next       coverage.py v7.14.0, created at 2026-09-23 22:14 -0400

1"""Embedding generation service for table metadata.""" 

2 

3from enum import Enum 

4from typing import Any, List 

5 

6import numpy as np 

7from loguru import logger 

8 

9from quber.settings import get_settings 

10 

11 

12class EmbeddingProvider(Enum): 

13 """Available embedding providers.""" 

14 

15 LOCAL = "local" 

16 OPENAI = "openai" 

17 

18 

19class EmbeddingService: 

20 """Generate embeddings for table metadata using local or OpenAI models.""" 

21 

22 def __init__( 

23 self, 

24 provider: str | EmbeddingProvider = "local", 

25 device: str = "cuda", 

26 model_name: str | None = None, 

27 ): 

28 """ 

29 Initialize the embedding service. 

30 

31 Args: 

32 provider: Embedding provider ('local' or 'openai') 

33 device: Device for local model ('cuda' or 'cpu') 

34 model_name: Optional custom model name (defaults per provider) 

35 """ 

36 if isinstance(provider, str): 

37 provider = EmbeddingProvider(provider) 

38 

39 self.provider = provider 

40 self.device = device 

41 self.model: Any = None 

42 self.client: Any = None 

43 self.model_name: str = "" 

44 self.embedding_dim: int = 0 

45 

46 if provider == EmbeddingProvider.LOCAL: 

47 self.init_local_model(model_name or "BAAI/bge-large-en-v1.5", device) 

48 elif provider == EmbeddingProvider.OPENAI: 

49 self.init_openai_client(model_name or "text-embedding-3-small") 

50 else: 

51 raise ValueError(f"Unsupported provider: {provider}") 

52 

53 logger.info(f"Initialized {provider.value} embedding service") 

54 

55 def init_local_model(self, model_name: str, device: str): 

56 """Initialize local sentence-transformers model.""" 

57 try: 

58 from sentence_transformers import SentenceTransformer 

59 except ImportError as err: 

60 raise ImportError( 

61 "sentence-transformers not installed. Install with: pip install sentence-transformers" 

62 ) from err 

63 

64 logger.info(f"Loading local model: {model_name} on {device}") 

65 self.model = SentenceTransformer(model_name) 

66 

67 # Move to specified device (CUDA if available) 

68 if device == "cuda": 

69 import torch 

70 

71 if torch.cuda.is_available(): 

72 self.model.to(device) 

73 logger.success(f"Model loaded on {device} (GPU: {torch.cuda.get_device_name(0)})") 

74 else: 

75 logger.warning("CUDA not available, using CPU") 

76 self.model.to("cpu") 

77 self.device = "cpu" 

78 else: 

79 self.model.to("cpu") 

80 

81 self.model_name = model_name 

82 self.embedding_dim = self.model.get_sentence_embedding_dimension() 

83 logger.info(f"Model embedding dimension: {self.embedding_dim}") 

84 

85 def init_openai_client(self, model_name: str): 

86 """Initialize OpenAI client.""" 

87 try: 

88 from openai import OpenAI 

89 except ImportError as err: 

90 raise ImportError("openai not installed. Install with: pip install openai") from err 

91 

92 api_key = get_settings().llm.openai_api_key 

93 if not api_key: 

94 raise ValueError("OPENAI_API_KEY environment variable not set") 

95 

96 self.client = OpenAI(api_key=api_key) 

97 self.model_name = model_name 

98 # OpenAI text-embedding-3-small is 1536 dims 

99 self.embedding_dim = 1536 if "3-small" in model_name else 3072 

100 logger.info(f"Using OpenAI model: {model_name} ({self.embedding_dim} dims)") 

101 

102 def embed( 

103 self, texts: List[str], batch_size: int = 32, show_progress: bool = False 

104 ) -> "np.ndarray[Any, Any]": 

105 """ 

106 Generate embeddings for a list of texts. 

107 

108 Args: 

109 texts: List of text strings to embed 

110 batch_size: Batch size for processing (local model only) 

111 show_progress: Show progress bar (local model only) 

112 

113 Returns: 

114 NumPy array of shape (len(texts), embedding_dim) 

115 """ 

116 if not texts: 

117 return np.array([]) 

118 

119 if self.provider == EmbeddingProvider.LOCAL: 

120 return self.embed_local(texts, batch_size, show_progress) 

121 else: 

122 return self.embed_openai(texts) 

123 

124 def embed_local(self, texts: List[str], batch_size: int, show_progress: bool) -> "np.ndarray[Any, Any]": 

125 """Generate embeddings using local model.""" 

126 logger.debug(f"Generating {len(texts)} embeddings (batch_size={batch_size})") 

127 embeddings = self.model.encode( 

128 texts, batch_size=batch_size, show_progress_bar=show_progress, convert_to_numpy=True 

129 ) 

130 return embeddings 

131 

132 def embed_openai(self, texts: List[str]) -> "np.ndarray[Any, Any]": 

133 """Generate embeddings using OpenAI API.""" 

134 logger.debug(f"Generating {len(texts)} embeddings via OpenAI") 

135 

136 # OpenAI has a limit of 2048 texts per request 

137 # For simplicity, batch in chunks of 100 

138 all_embeddings = [] 

139 chunk_size = 100 

140 

141 for i in range(0, len(texts), chunk_size): 

142 chunk = texts[i : i + chunk_size] 

143 response = self.client.embeddings.create(input=chunk, model=self.model_name) 

144 embeddings = [item.embedding for item in response.data] 

145 all_embeddings.extend(embeddings) 

146 

147 return np.array(all_embeddings) 

148 

149 def embed_table_metadata(self, title: str, description: str) -> "np.ndarray[Any, Any]": 

150 """ 

151 Generate embedding for table metadata. 

152 

153 Combines title and description into a single text for embedding. 

154 

155 Args: 

156 title: Table title (llm_title) 

157 description: Table description (llm_description) 

158 

159 Returns: 

160 Embedding vector as NumPy array 

161 """ 

162 # Combine title and description 

163 # Title is more important, so it comes first 

164 text = f"{title}\n\n{description}" 

165 

166 embeddings = self.embed([text]) 

167 return embeddings[0] 

168 

169 

170def get_embedding_service(provider: str | None = None, device: str | None = None) -> EmbeddingService: 

171 """ 

172 Factory function to create an embedding service based on environment config. 

173 

174 Args: 

175 provider: Override provider (defaults to EMBEDDING_PROVIDER env var or 'local') 

176 device: Override device (defaults to EMBEDDING_DEVICE env var or 'cuda') 

177 

178 Returns: 

179 Configured EmbeddingService instance 

180 """ 

181 embedding_settings = get_settings().embeddings 

182 if provider is None: 

183 provider = embedding_settings.provider 

184 

185 if device is None: 

186 device = embedding_settings.device 

187 

188 return EmbeddingService(provider=provider, device=device)