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
« prev ^ index » next coverage.py v7.14.0, created at 2026-09-23 22:14 -0400
1"""Embedding generation service for table metadata."""
3from enum import Enum
4from typing import Any, List
6import numpy as np
7from loguru import logger
9from quber.settings import get_settings
12class EmbeddingProvider(Enum):
13 """Available embedding providers."""
15 LOCAL = "local"
16 OPENAI = "openai"
19class EmbeddingService:
20 """Generate embeddings for table metadata using local or OpenAI models."""
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.
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)
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
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}")
53 logger.info(f"Initialized {provider.value} embedding service")
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
64 logger.info(f"Loading local model: {model_name} on {device}")
65 self.model = SentenceTransformer(model_name)
67 # Move to specified device (CUDA if available)
68 if device == "cuda":
69 import torch
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")
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}")
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
92 api_key = get_settings().llm.openai_api_key
93 if not api_key:
94 raise ValueError("OPENAI_API_KEY environment variable not set")
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)")
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.
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)
113 Returns:
114 NumPy array of shape (len(texts), embedding_dim)
115 """
116 if not texts:
117 return np.array([])
119 if self.provider == EmbeddingProvider.LOCAL:
120 return self.embed_local(texts, batch_size, show_progress)
121 else:
122 return self.embed_openai(texts)
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
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")
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
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)
147 return np.array(all_embeddings)
149 def embed_table_metadata(self, title: str, description: str) -> "np.ndarray[Any, Any]":
150 """
151 Generate embedding for table metadata.
153 Combines title and description into a single text for embedding.
155 Args:
156 title: Table title (llm_title)
157 description: Table description (llm_description)
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}"
166 embeddings = self.embed([text])
167 return embeddings[0]
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.
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')
178 Returns:
179 Configured EmbeddingService instance
180 """
181 embedding_settings = get_settings().embeddings
182 if provider is None:
183 provider = embedding_settings.provider
185 if device is None:
186 device = embedding_settings.device
188 return EmbeddingService(provider=provider, device=device)