mcp 2.x already on main. Does not replace the variety chunker (one chunk per variety is the anti-hallucination contract). - Numbered [1] citations on search_docs / search_trials - Eval JSONL sidecar, k-curve, eval.pvalue, eval.trace - Nomic prefixes at embed time only; stored text unprefixed Closes #25
94 lines
3.3 KiB
Python
94 lines
3.3 KiB
Python
"""Embedding function for Chroma — Ollama-hosted nomic-embed-text by default.
|
|
|
|
Swappable: implement the same `embedding_function()` interface returning
|
|
a Chroma `EmbeddingFunction` and the rest of the pipeline doesn't care.
|
|
|
|
Defaults (override via env):
|
|
OLLAMA_URL one or more comma-separated URLs (load-balanced)
|
|
EMBED_MODEL model name; default 'nomic-embed-text'
|
|
EMBED_DIM expected embedding dim; default 768 (nomic-embed-text)
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import os
|
|
import logging
|
|
from typing import Any
|
|
|
|
try:
|
|
import httpx
|
|
from chromadb import Documents, EmbeddingFunction, Embeddings
|
|
except ImportError: # unit tests for apply_prefix / embed_texts
|
|
httpx = None # type: ignore[assignment]
|
|
EmbeddingFunction = object # type: ignore[misc,assignment]
|
|
Documents = list # type: ignore[misc,assignment]
|
|
Embeddings = list # type: ignore[misc,assignment]
|
|
|
|
log = logging.getLogger(__name__)
|
|
|
|
OLLAMA_URLS = [u.strip() for u in os.environ.get("OLLAMA_URL",
|
|
"http://localhost:11434").split(",") if u.strip()]
|
|
EMBED_MODEL = os.environ.get("EMBED_MODEL", "nomic-embed-text")
|
|
EMBED_DIM = int(os.environ.get("EMBED_DIM", "768"))
|
|
# nomic asymmetric prefixes — embed time only. Empty string disables.
|
|
EMBED_QUERY_PREFIX = os.environ.get("EMBED_QUERY_PREFIX", "search_query: ")
|
|
EMBED_DOC_PREFIX = os.environ.get("EMBED_DOC_PREFIX", "search_document: ")
|
|
|
|
|
|
def apply_prefix(text: str, prefix: str) -> str:
|
|
if not prefix:
|
|
return text
|
|
return prefix + text
|
|
|
|
|
|
def embed_texts(texts: list[str], *, prefix: str, ef: Any | None = None) -> list[list[float]]:
|
|
"""Embed texts after applying prefix. Callers store the original strings."""
|
|
fn = ef if ef is not None else embedding_function()
|
|
return list(fn([apply_prefix(t, prefix) for t in texts]))
|
|
|
|
|
|
class OllamaEmbeddings(EmbeddingFunction):
|
|
"""Calls /api/embed across N Ollama endpoints, naive round-robin.
|
|
|
|
For indexing throughput on multiple GPUs, run one Ollama container
|
|
per GPU (pinned via NVIDIA_VISIBLE_DEVICES) and pass all their URLs
|
|
in OLLAMA_URL — the embedder picks the next endpoint per batch.
|
|
"""
|
|
|
|
def __init__(self, urls: list[str] = OLLAMA_URLS, model: str = EMBED_MODEL):
|
|
self.urls = urls
|
|
self.model = model
|
|
self._next = 0
|
|
|
|
def __call__(self, input: Documents) -> Embeddings:
|
|
url = self.urls[self._next % len(self.urls)]
|
|
self._next += 1
|
|
with httpx.Client(timeout=300) as c:
|
|
r = c.post(f"{url}/api/embed",
|
|
json={"model": self.model, "input": list(input)})
|
|
r.raise_for_status()
|
|
data = r.json()
|
|
return data.get("embeddings") or []
|
|
|
|
def name(self) -> str: # newer chromadb requires this
|
|
return f"ollama:{self.model}"
|
|
|
|
@staticmethod
|
|
def build_from_config(config: dict) -> "OllamaEmbeddings": # newer chromadb
|
|
return OllamaEmbeddings(
|
|
urls=config.get("urls", OLLAMA_URLS),
|
|
model=config.get("model", EMBED_MODEL),
|
|
)
|
|
|
|
def get_config(self) -> dict: # newer chromadb
|
|
return {"urls": self.urls, "model": self.model}
|
|
|
|
def default_space(self) -> str:
|
|
return "cosine"
|
|
|
|
def supported_spaces(self) -> list[str]:
|
|
return ["cosine", "l2", "ip"]
|
|
|
|
|
|
def embedding_function() -> EmbeddingFunction:
|
|
return OllamaEmbeddings()
|