didi-lot1-ai/ai_platform/modules/didi_brain/shared/embedding_client.py

203 lines
7.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""Embedding + reranking client.
- Embeddings: BAAI/bge-m3 via vLLM OpenAI-compat endpoint, 1024-dim, 8K context.
- Reranker: BAAI/bge-reranker-v2-m3 cross-encoder via vllm-rerank-api endpoint.
Both run on the same host (10.11.10.15) on different ports (8200 / 8100).
"""
from __future__ import annotations
import math
from collections.abc import AsyncIterator, Sequence
from contextlib import asynccontextmanager
from dataclasses import dataclass
from typing import Any
import httpx
from tenacity import (
AsyncRetrying,
retry_if_exception_type,
stop_after_attempt,
wait_exponential,
)
from shared.config import settings
from shared.logging import get_logger
log = get_logger(__name__)
class EmbeddingError(Exception):
pass
@dataclass(slots=True, frozen=True)
class RerankResult:
"""One reranked document, sorted by score descending."""
index: int # original position in the input documents list
score: float # cross-encoder score (higher = more relevant)
document: str # the text of that document
class EmbeddingClient:
"""Async client for BGE-M3 embeddings + BGE-reranker-v2-m3 reranking."""
def __init__(
self,
*,
embed_url: str | None = None,
rerank_url: str | None = None,
timeout: float = 60.0,
max_attempts: int = 3,
):
self._embed_base = (embed_url or settings.embedding_url).rstrip("/")
self._rerank_base = (rerank_url or settings.reranker_url).rstrip("/")
self._max_attempts = max_attempts
embed_headers = {"Content-Type": "application/json"}
if settings.embedding_api_key:
embed_headers["Authorization"] = f"Bearer {settings.embedding_api_key}"
rerank_headers = {"Content-Type": "application/json"}
if settings.reranker_api_key:
rerank_headers["Authorization"] = f"Bearer {settings.reranker_api_key}"
self._embed_http = httpx.AsyncClient(
base_url=self._embed_base,
headers=embed_headers,
timeout=httpx.Timeout(timeout, connect=10.0),
)
self._rerank_http = httpx.AsyncClient(
base_url=self._rerank_base,
headers=rerank_headers,
timeout=httpx.Timeout(timeout, connect=10.0),
)
async def aclose(self) -> None:
await self._embed_http.aclose()
await self._rerank_http.aclose()
async def __aenter__(self) -> EmbeddingClient:
return self
async def __aexit__(self, *args: Any) -> None:
await self.aclose()
# ------------------------------------------------------------------ embed
async def embed(self, texts: Sequence[str]) -> list[list[float]]:
"""Embed a batch of texts. Returns 1024-dim vectors in input order.
BGE-M3 max input is 8192 tokens; we don't truncate here — caller is
responsible for chunking long inputs (Atomic does this server-side).
Empty input returns empty list.
"""
if not texts:
return []
payload = {"model": settings.embedding_model, "input": list(texts)}
async for attempt in AsyncRetrying(
stop=stop_after_attempt(self._max_attempts),
wait=wait_exponential(multiplier=1, min=1, max=10),
retry=retry_if_exception_type((httpx.HTTPError, httpx.TimeoutException)),
reraise=True,
):
with attempt:
resp = await self._embed_http.post("/v1/embeddings", json=payload)
if resp.status_code >= 400:
raise EmbeddingError(
f"HTTP {resp.status_code} from embedding endpoint: "
f"{resp.text[:500]}"
)
data = resp.json()
items = data.get("data")
if not items or len(items) != len(texts):
raise EmbeddingError(
f"Embedding count mismatch: got {len(items or [])} for "
f"{len(texts)} inputs"
)
vecs = [item["embedding"] for item in items]
if vecs and len(vecs[0]) != settings.embedding_dim:
raise EmbeddingError(
f"Embedding dim mismatch: got {len(vecs[0])}, "
f"expected {settings.embedding_dim}"
)
return vecs
raise EmbeddingError("retry loop exited without result") # unreachable
async def embed_one(self, text: str) -> list[float]:
"""Convenience for single-text embedding."""
result = await self.embed([text])
return result[0]
# ----------------------------------------------------------------- rerank
async def rerank(
self,
query: str,
documents: Sequence[str],
*,
top_n: int | None = None,
) -> list[RerankResult]:
"""Cross-encoder rerank query × documents. Returns sorted by score desc.
Use this AFTER first-stage embedding retrieval to dramatically boost
precision on the top-N candidates (typical pattern: embed retrieves 50,
rerank picks 10).
"""
if not documents:
return []
payload: dict[str, Any] = {
"model": settings.reranker_model,
"query": query,
"documents": list(documents),
}
if top_n is not None:
payload["top_n"] = top_n
async for attempt in AsyncRetrying(
stop=stop_after_attempt(self._max_attempts),
wait=wait_exponential(multiplier=1, min=1, max=10),
retry=retry_if_exception_type((httpx.HTTPError, httpx.TimeoutException)),
reraise=True,
):
with attempt:
resp = await self._rerank_http.post("/v1/rerank", json=payload)
if resp.status_code >= 400:
raise EmbeddingError(
f"HTTP {resp.status_code} from reranker: {resp.text[:500]}"
)
data = resp.json()
results = data.get("results", [])
return [
RerankResult(
index=int(r["index"]),
score=float(r["score"]),
document=r.get("document") or documents[int(r["index"])],
)
for r in results
]
raise EmbeddingError("retry loop exited without result") # unreachable
# --------------------------------------------------------------------- math util
def cosine(a: Sequence[float], b: Sequence[float]) -> float:
"""Cosine similarity. BGE-M3 vectors are NOT pre-normalized; we compute it."""
if len(a) != len(b):
raise ValueError(f"vector dim mismatch: {len(a)} vs {len(b)}")
dot = sum(x * y for x, y in zip(a, b, strict=True))
na = math.sqrt(sum(x * x for x in a))
nb = math.sqrt(sum(x * x for x in b))
if na == 0 or nb == 0:
return 0.0
return dot / (na * nb)
@asynccontextmanager
async def embedding_client() -> AsyncIterator[EmbeddingClient]:
client = EmbeddingClient()
try:
yield client
finally:
await client.aclose()