Spaces:
Running
Running
| """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}") | |
| 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) | |
| 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() | |
| 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] | |