medbillcodes-api / app /embeddings.py
medbillcodes-deploy
Deploy cloud pilot API
1ddeb51
Raw
History Blame Contribute Delete
3.93 kB
"""Embedding generation — local SentenceTransformers or remote TEI/HF HTTP API.
Default model: Alibaba-NLP/gte-multilingual-base (768 dims).
When EMBEDDING_URL is set, vectors are fetched over HTTPS (no local torch),
which is the preferred path for a no-laptop cloud pilot.
"""
from __future__ import annotations
import logging
from functools import lru_cache
import httpx
from .config import settings
from .weave_trace import redact_inputs, summarize_embedding, traced
logger = logging.getLogger(__name__)
def _embedding_auth_headers() -> dict[str, str]:
token = (
settings.embedding_api_key
or settings.hf_token
or settings.llm_api_key
or settings.wandb_api_key
or ""
).strip()
if not token:
return {}
return {"Authorization": f"Bearer {token}"}
def _normalize(vec: list[float]) -> list[float]:
import math
norm = math.sqrt(sum(x * x for x in vec)) or 1.0
return [x / norm for x in vec]
def _embed_remote(texts: list[str]) -> list[list[float]]:
"""Call TEI / OpenAI-compatible embeddings, or HF feature-extraction style."""
base = settings.embedding_url.rstrip("/")
headers = {"Content-Type": "application/json", **_embedding_auth_headers()}
# Prefer OpenAI-compatible /v1/embeddings (TEI and many HF endpoints).
url = base if base.endswith("/embeddings") else f"{base}/v1/embeddings"
payload = {"model": settings.embedding_model, "input": texts}
with httpx.Client(timeout=120.0) as client:
resp = client.post(url, json=payload, headers=headers)
if resp.status_code == 404 and not base.endswith("/embeddings"):
# Fallback: TEI native /embed
resp = client.post(
f"{base}/embed",
json={"inputs": texts if len(texts) > 1 else texts[0]},
headers=headers,
)
resp.raise_for_status()
data = resp.json()
if isinstance(data, list) and data and isinstance(data[0], (int, float)):
return [_normalize([float(x) for x in data])]
if isinstance(data, list) and data and isinstance(data[0], list):
return [_normalize([float(x) for x in row]) for row in data]
if isinstance(data, dict) and "data" in data:
rows = sorted(data["data"], key=lambda r: r.get("index", 0))
return [_normalize([float(x) for x in r["embedding"]]) for r in rows]
raise RuntimeError(f"Unrecognized embedding response shape from {url}")
@lru_cache(maxsize=1)
def _model():
from sentence_transformers import SentenceTransformer
logger.info("Loading embedding model %s", settings.embedding_model)
# gte-multilingual-base requires trust_remote_code for its custom pooling.
return SentenceTransformer(settings.embedding_model, trust_remote_code=True)
@traced(
"embed_text",
kind="tool",
postprocess_inputs=redact_inputs,
postprocess_output=summarize_embedding,
)
def embed_text(text: str) -> list[float]:
"""Return a single embedding for the given text."""
if settings.embedding_url:
return _embed_remote([text])[0]
vec = _model().encode(
[text], normalize_embeddings=True, convert_to_numpy=True
)[0]
return vec.tolist()
@traced(
"embed_batch",
kind="tool",
postprocess_inputs=redact_inputs,
postprocess_output=summarize_embedding,
)
def embed_batch(texts: list[str]) -> list[list[float]]:
"""Batch embed multiple texts (used during ingestion)."""
if not texts:
return []
if settings.embedding_url:
out: list[list[float]] = []
# Keep batches modest for remote rate limits.
chunk = 32
for i in range(0, len(texts), chunk):
out.extend(_embed_remote(texts[i : i + chunk]))
return out
vecs = _model().encode(
texts, normalize_embeddings=True, convert_to_numpy=True, batch_size=32
)
return [v.tolist() for v in vecs]