embedding-server / app /services /embeddings_service.py
Faysal4200's picture
Upload 45 files
ec855e6 verified
Raw
History Blame Contribute Delete
6.97 kB
"""
EmbeddingsService:
- Loads tokenizer + model with retries (tenacity)
- Efficiently computes embeddings for large batches by chunking into BATCH_SIZE
- Mean-pools the last hidden states to produce embeddings
- Handles OOM errors by falling back to micro-batch single-item processing
"""
from ..config import settings
from ..logger import logger
import math
from typing import List, Optional
import torch
from transformers import AutoTokenizer, AutoModel
from tenacity import retry, stop_after_attempt, wait_exponential, retry_if_exception_type
class EmbeddingsService:
def __init__(self, model_name: Optional[str] = None):
self.model_name = model_name or settings.MODEL_NAME
# If DEVICE is set in env, honor it; otherwise auto-detect GPU if available
if settings.DEVICE:
self.device = torch.device(settings.DEVICE)
else:
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.model: Optional[AutoModel] = None
self.tokenizer: Optional[AutoTokenizer] = None
self.is_ready: bool = False
# Retry decorator: will attempt to load model multiple times with exponential backoff
@retry(
stop=stop_after_attempt(settings.RETRY_ATTEMPTS),
wait=wait_exponential(multiplier=settings.RETRY_BACKOFF_SECONDS, max=20), # max wait 20 second for next retry
retry=retry_if_exception_type(Exception),
reraise=True,
)
def load_model(self) -> None:
if self.is_ready:
return
logger.info(f"Loading model: {self.model_name}")
logger.info(f"Target device: {self.device}")
try:
#-------------------------------------------------
# Using private Hugging Face token to fetch model
self.tokenizer = AutoTokenizer.from_pretrained(
self.model_name,
use_fast=True,
token=settings.HUGGING_FACE_TOKEN
)
self.model = AutoModel.from_pretrained(
self.model_name,
use_safetensors=True,
token=settings.HUGGING_FACE_TOKEN
)
#-------------------------------------------------
self.model.to(self.device)
self.model.eval()
self.is_ready = True
logger.info(f"Model '{self.model_name}' loaded and ready.")
except Exception as exc:
logger.error(f"Failed to load model '{self.model_name}'")
logger.exception(exc)
raise
def _pool_embeddings(self, last_hidden_state: torch.Tensor) -> torch.Tensor:
"""Mean pooling across token dimension: (B, T, D) -> (B, D)."""
return last_hidden_state.mean(dim=1)
@torch.no_grad()
def generate_single_embedding(self, text: str) -> List[float]:
"""
Generate a single embedding synchronously.
"""
self.load_model()
if not text:
logger.warning("Received empty text for embedding")
raise ValueError("Empty text provided for embedding.")
log_text = (text[:100] + '...') if len(text) > 100 else text
logger.debug(f"Processing single embedding for text: '{log_text}'")
# debuging here text length by tokens
# token_count = len(self.tokenizer.tokenize(text))
# if token_count > settings.MAX_LENGTH:
# logger.warning(f"Text token count {token_count} is longer than max length ({settings.MAX_LENGTH}), truncating...")
try:
inputs = self.tokenizer(
text,
return_tensors="pt",
truncation=True,
padding=True,
max_length=settings.MAX_LENGTH,
).to(self.device)
outputs = self.model(**inputs)
emb = self._pool_embeddings(outputs.last_hidden_state)
return emb.cpu().numpy().flatten().tolist()
except Exception as e:
logger.error("Critical error during single embedding generation")
logger.exception(e)
raise
@torch.no_grad()
def generate_batch_embeddings(self, texts: List[str]) -> List[List[float]]:
"""
Efficient batch processing with OOM recovery (Sync).
"""
self.load_model()
if not texts:
return []
batch_size = max(1, settings.BATCH_SIZE)
embeddings: List[List[float]] = []
total = len(texts)
num_chunks = math.ceil(total / batch_size)
logger.info(f"Processing {total} texts in {num_chunks} chunks (batch_size={batch_size})")
for i in range(0, total, batch_size):
chunk = texts[i : i + batch_size]
try:
# debuging here every text
for index, text in enumerate(chunk):
log_text = (text[:100] + '...') if len(text) > 100 else text
logger.debug(f"Processing chunk-{i+1}, embedding for text-{index+1}: '{log_text}'")
# debuging here text length by tokens
# token_count = len(self.tokenizer.tokenize(text))
# if token_count > settings.MAX_LENGTH:
# logger.warning(f"Text token count {token_count} is longer than max length ({settings.MAX_LENGTH}), truncating...")
inputs = self.tokenizer(
chunk,
return_tensors="pt",
truncation=True,
padding=True,
max_length=settings.MAX_LENGTH,
).to(self.device)
outputs = self.model(**inputs)
chunk_embs = self._pool_embeddings(outputs.last_hidden_state)
embeddings.extend(chunk_embs.cpu().numpy().tolist())
except RuntimeError as e:
# Basic check for OOM
msg = str(e).lower()
if "out of memory" in msg or "cuda out of memory" in msg:
logger.warning(f"GPU OOM on chunk {i//batch_size + 1}. Attempting micro-batch fallback...")
try:
torch.cuda.empty_cache()
except: pass
for t in chunk:
embeddings.append(self.generate_single_embedding(t))
else:
logger.error(f"Runtime error in batch processing at chunk {i//batch_size + 1}")
logger.exception(e)
raise
except Exception as e:
logger.error(f"Unexpected error in batch processing at chunk {i//batch_size + 1}")
logger.exception(e)
raise
return embeddings