Livrare LOT 1 - Didi
This commit is contained in:
commit
5380c3fc63
990 changed files with 133308 additions and 0 deletions
203
ai_platform/modules/didi_brain/shared/embedding_client.py
Normal file
203
ai_platform/modules/didi_brain/shared/embedding_client.py
Normal file
|
|
@ -0,0 +1,203 @@
|
|||
"""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()
|
||||
Loading…
Add table
Add a link
Reference in a new issue