From 71c71f6e0ba08a36c9fc3c7138b4bd3f98be93f6 Mon Sep 17 00:00:00 2001 From: Dawn-OuYang Date: Thu, 27 Aug 2026 01:24:14 +0800 Subject: [PATCH] Add OpenAI API embedder Signed-off-by: Dawn-OuYang --- CHANGELOG.md | 2 + README.md | 12 ++++ pyproject.toml | 2 + src/retrieval_lab/cli.py | 13 ++++- src/retrieval_lab/embedding/__init__.py | 6 ++ src/retrieval_lab/embedding/api.py | 78 +++++++++++++++++++++++++ tests/test_cli.py | 10 ++++ tests/test_embedding.py | 51 ++++++++++++++++ 8 files changed, 171 insertions(+), 3 deletions(-) create mode 100644 src/retrieval_lab/embedding/api.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 039a01d..e9697ae 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -6,6 +6,8 @@ additive features; the public API is not yet frozen). ## [Unreleased] +- Added an optional `[api-embed]` extra and `api:` CLI selector for hosted OpenAI + embedding models, while keeping the default install and tests keyless. - Docs: the live demo report and README now state up front that the published benchmark is an example run on a small synthetic API-documentation corpus, so a first-time visitor knows what they are looking at. diff --git a/README.md b/README.md index 3500b67..40e5735 100644 --- a/README.md +++ b/README.md @@ -54,6 +54,15 @@ retrieval-lab run \ --html report.html ``` +To compare a hosted embedding model, install the API extra and select an `api:` model. The +OpenAI client reads credentials from `OPENAI_API_KEY`. + +```bash +pip install "retrieval-lab[api-embed]" +retrieval-lab run --corpus docs.jsonl --queries queries.jsonl \ + --embed-models api:text-embedding-3-small +``` + Open `report.html` directly in a browser. It has no server or external frontend dependencies, supports light and dark themes, and remains usable without network access. @@ -189,6 +198,7 @@ retrieval-lab geometry --corpus docs.jsonl --embed-model e5 - `retrieval-lab`: lightweight core with NumPy and the deterministic embedder - `retrieval-lab[real-embed]`: E5 and BGE through sentence-transformers +- `retrieval-lab[api-embed]`: hosted OpenAI embedding models via `api:` - `retrieval-lab[rerank]`: cross-encoder reranking - `retrieval-lab[ann]`: HNSW approximate dense indexes - `retrieval-lab[dev]`: pytest and Ruff for development @@ -200,6 +210,8 @@ Python 3.10–3.12 is supported. - Results are only as representative as the labeled query set. - Missing valid gold alternatives make measured recall a lower bound. - Latency and index cost depend on the machine running the benchmark. +- Hosted API embedders depend on provider availability, pricing, rate limits, and the API key + available in the local environment. - Stage attribution requires a decomposable retrieval pipeline; black-box retrievers can only be scored at their observable output. diff --git a/pyproject.toml b/pyproject.toml index 835f0a5..23f1007 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -31,6 +31,8 @@ dependencies = [ [project.optional-dependencies] # Real embedding models (e5, bge, ...). Large downloads; opt-in. real-embed = ["sentence-transformers>=2.2"] +# Hosted embedding APIs such as OpenAI. Requires provider credentials at runtime. +api-embed = ["openai>=1.0"] # Cross-encoder reranking. rerank = ["sentence-transformers>=2.2"] # Approximate nearest-neighbour dense index. diff --git a/src/retrieval_lab/cli.py b/src/retrieval_lab/cli.py index 970f3b6..6e3e430 100644 --- a/src/retrieval_lab/cli.py +++ b/src/retrieval_lab/cli.py @@ -50,6 +50,13 @@ def _csv(value: str) -> list[str]: def _make_embedder(name: str, cache: EmbeddingCache): if name in ("det", "det-hash"): return DeterministicEmbedder(dim=2048, name=name, cache=cache) + if name.startswith("api:"): + from retrieval_lab.embedding import openai_embedder + + model = name.split(":", 1)[1] + emb = openai_embedder(model, cache=cache) + emb.name = name + return emb if name in ("e5", "bge"): from retrieval_lab.embedding import bge_embedder, e5_embedder @@ -58,7 +65,7 @@ def _make_embedder(name: str, cache: EmbeddingCache): # Re-key under the short CLI name so config ids stay readable. emb.name = name return emb - raise ValueError(f"unknown embed model {name!r} (use det, e5, or bge)") + raise ValueError(f"unknown embed model {name!r} (use det, e5, bge, or api:)") def _make_chunker(spec: str, embedder=None): @@ -310,7 +317,7 @@ def build_parser() -> argparse.ArgumentParser: run = sub.add_parser("run", help="sweep configs over a corpus + query set") run.add_argument("--corpus", required=True, help="documents JSONL ({id, text, meta?})") run.add_argument("--queries", required=True, help="queries JSONL (with source-span gold)") - run.add_argument("--embed-models", default="det", help="csv: det,e5,bge") + run.add_argument("--embed-models", default="det", help="csv: det,e5,bge,api:") run.add_argument( "--chunkers", default="fixed", @@ -392,7 +399,7 @@ def build_parser() -> argparse.ArgumentParser: geometry = sub.add_parser("geometry", help="embedding-space diagnostics (risk indicators)") geometry.add_argument("--corpus", required=True, help="documents JSONL") - geometry.add_argument("--embed-model", default="det", help="det, e5, or bge") + geometry.add_argument("--embed-model", default="det", help="det, e5, bge, or api:") geometry.add_argument("--chunker", default="fixed", help="chunker for the corpus vectors") geometry.add_argument("--queries", default=None, help="optional queries JSONL for mismatch") geometry.set_defaults(func=_cmd_geometry) diff --git a/src/retrieval_lab/embedding/__init__.py b/src/retrieval_lab/embedding/__init__.py index 338d024..9ecb780 100644 --- a/src/retrieval_lab/embedding/__init__.py +++ b/src/retrieval_lab/embedding/__init__.py @@ -13,6 +13,8 @@ "SentenceTransformerEmbedder", "e5_embedder", "bge_embedder", + "OpenAIEmbedder", + "openai_embedder", ] @@ -21,4 +23,8 @@ def __getattr__(name: str): from retrieval_lab.embedding import sentence_transformer as _st return getattr(_st, name) + if name in {"OpenAIEmbedder", "openai_embedder"}: + from retrieval_lab.embedding import api as _api + + return getattr(_api, name) raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/src/retrieval_lab/embedding/api.py b/src/retrieval_lab/embedding/api.py new file mode 100644 index 0000000..528c2fc --- /dev/null +++ b/src/retrieval_lab/embedding/api.py @@ -0,0 +1,78 @@ +"""Hosted embedding providers behind the ``[api-embed]`` extra. + +The OpenAI client is imported lazily so the default install and test suite stay keyless. +""" + +from __future__ import annotations + +from collections.abc import Sequence + +import numpy as np + +from retrieval_lab.embedding.base import Embedder, EmbeddingCache, l2_normalize + + +class OpenAIEmbedder(Embedder): + """OpenAI embeddings adapter. + + The API key is resolved by the OpenAI SDK from the usual environment variables, so + tests can inject a fake client and normal runs can use ``OPENAI_API_KEY``. + """ + + def __init__( + self, + model_name: str = "text-embedding-3-small", + cache: EmbeddingCache | None = None, + *, + client=None, + dim: int = 1536, + dimensions: int | None = None, + ) -> None: + if client is None: + try: + from openai import OpenAI + except ImportError as exc: # pragma: no cover - only without the extra + raise ImportError( + "OpenAIEmbedder needs the '[api-embed]' extra: " + "pip install 'retrieval-lab[api-embed]'" + ) from exc + client = OpenAI() + + name = f"api:{model_name}" + if dimensions is not None: + name = f"{name}:dim={dimensions}" + dim = dimensions + super().__init__(name=name, dim=dim, cache=cache) + self.model_name = model_name + self.dimensions = dimensions + self._client = client + + def _embed_raw(self, texts: list[str]) -> np.ndarray: + kwargs = {"model": self.model_name, "input": texts} + if self.dimensions is not None: + kwargs["dimensions"] = self.dimensions + response = self._client.embeddings.create(**kwargs) + vectors = [_embedding_vector(item) for item in response.data] + return l2_normalize(np.asarray(vectors, dtype=np.float32)) + + def embed_query(self, texts: Sequence[str]) -> np.ndarray: + return self.embed(texts) + + def embed_passage(self, texts: Sequence[str]) -> np.ndarray: + return self.embed(texts) + + +def _embedding_vector(item) -> Sequence[float]: + if isinstance(item, dict): + return item["embedding"] + return item.embedding + + +def openai_embedder( + model_name: str = "text-embedding-3-small", + cache: EmbeddingCache | None = None, + *, + dim: int = 1536, + dimensions: int | None = None, +) -> OpenAIEmbedder: + return OpenAIEmbedder(model_name, cache=cache, dim=dim, dimensions=dimensions) diff --git a/tests/test_cli.py b/tests/test_cli.py index fb7694f..5257f1e 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -116,6 +116,16 @@ def test_unknown_embed_model_is_an_error(demo, capsys): assert "unknown embed model" in capsys.readouterr().err +def test_api_embed_model_missing_extra_is_a_clean_input_error(demo, capsys): + docs, queries, _tmp = demo + code = main([ + "run", "--corpus", str(docs), "--queries", str(queries), + "--embed-models", "api:text-embedding-3-small", "--min-sample", "1", + ]) + assert code == EXIT_INPUT_ERROR + assert "retrieval-lab[api-embed]" in capsys.readouterr().err + + def test_run_writes_html(demo, tmp_path): docs, queries, _tmp = demo html = tmp_path / "report.html" diff --git a/tests/test_embedding.py b/tests/test_embedding.py index 45fe58d..a4581ef 100644 --- a/tests/test_embedding.py +++ b/tests/test_embedding.py @@ -1,8 +1,10 @@ """Phase 1 — the keyless deterministic embedder + content-addressed cache (spec §I.7).""" import numpy as np +import pytest from retrieval_lab.embedding import DeterministicEmbedder, EmbeddingCache +from retrieval_lab.embedding.api import OpenAIEmbedder def cos(a, b): @@ -72,3 +74,52 @@ def test_empty_input_returns_empty_matrix(): e = DeterministicEmbedder(dim=32) out = e.embed([]) assert out.shape == (0, 32) + + +def test_openai_embedder_uses_injected_client_and_cache(): + class FakeEmbeddings: + def __init__(self): + self.calls = [] + + def create(self, **kwargs): + self.calls.append(kwargs) + vectors = [[1.0, 0.0, 0.0], [0.0, 2.0, 0.0]] + return type( + "EmbeddingResponse", + (), + {"data": [type("Embedding", (), {"embedding": v}) for v in vectors]}, + )() + + class FakeClient: + def __init__(self): + self.embeddings = FakeEmbeddings() + + client = FakeClient() + cache = EmbeddingCache() + embedder = OpenAIEmbedder("text-embedding-3-small", cache=cache, client=client, dim=3) + + first = embedder.embed(["alpha", "beta"]) + second = embedder.embed(["alpha", "beta"]) + + assert len(client.embeddings.calls) == 1 + assert client.embeddings.calls[0]["model"] == "text-embedding-3-small" + assert client.embeddings.calls[0]["input"] == ["alpha", "beta"] + assert np.array_equal(first, second) + assert np.allclose(np.linalg.norm(first, axis=1), 1.0) + assert len(cache) == 2 + + +def test_openai_embedder_missing_extra_is_actionable(monkeypatch): + import builtins + + real_import = builtins.__import__ + + def fake_import(name, *args, **kwargs): + if name == "openai": + raise ImportError("missing") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", fake_import) + + with pytest.raises(ImportError, match=r"\[api-embed\]"): + OpenAIEmbedder()