Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,8 @@ additive features; the public API is not yet frozen).

## [Unreleased]

- Added an optional `[api-embed]` extra and `api:<model>` 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.
Expand Down
12 changes: 12 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down Expand Up @@ -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:<model>`
- `retrieval-lab[rerank]`: cross-encoder reranking
- `retrieval-lab[ann]`: HNSW approximate dense indexes
- `retrieval-lab[dev]`: pytest and Ruff for development
Expand All @@ -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.

Expand Down
2 changes: 2 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
13 changes: 10 additions & 3 deletions src/retrieval_lab/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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:<model>)")


def _make_chunker(spec: str, embedder=None):
Expand Down Expand Up @@ -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:<model>")
run.add_argument(
"--chunkers",
default="fixed",
Expand Down Expand Up @@ -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:<model>")
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)
Expand Down
6 changes: 6 additions & 0 deletions src/retrieval_lab/embedding/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,8 @@
"SentenceTransformerEmbedder",
"e5_embedder",
"bge_embedder",
"OpenAIEmbedder",
"openai_embedder",
]


Expand All @@ -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}")
78 changes: 78 additions & 0 deletions src/retrieval_lab/embedding/api.py
Original file line number Diff line number Diff line change
@@ -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)
10 changes: 10 additions & 0 deletions tests/test_cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
51 changes: 51 additions & 0 deletions tests/test_embedding.py
Original file line number Diff line number Diff line change
@@ -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):
Expand Down Expand Up @@ -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()
Loading