diff --git a/codex/embed.py b/codex/embed.py new file mode 100644 index 0000000..fcfe67e --- /dev/null +++ b/codex/embed.py @@ -0,0 +1,155 @@ +"""Hybrid dense + sparse embeddings via BGE-M3 (ADR-0002). + +Uses :class:`FlagEmbedding.BGEM3FlagModel` for encoding because the +``return_dense`` / ``return_sparse`` kwargs are part of the FlagEmbedding +API — not ``sentence-transformers``. The dense output is identical to a +vanilla ``SentenceTransformer`` load of ``BAAI/bge-m3`` (same weights, +same model); sparse output is a list of ``{token_id: weight}`` dicts per +text. + +Notes for callers +----------------- +* Empty input is handled explicitly — no model call is issued. +* Dense vectors are L2-normalised row-wise so cosine similarity reduces + to a dot product. +* :func:`get_embedder` returns a process-wide singleton so the model is + loaded at most once per Python process. +""" + +from __future__ import annotations + +import numpy as np +import torch +from FlagEmbedding import BGEM3FlagModel + + +class Embedder: + """Thin wrapper over :class:`BGEM3FlagModel` with a stable API. + + The wrapper hides the FlagEmbedding-specific encode-dict and exposes + three operations relevant to the ingestion pipeline: + + * :meth:`encode_dense` — dense float32 matrix, L2-normalised. + * :meth:`encode_sparse` — list of ``{token_id: weight}`` dicts. + * :meth:`encode` — both in a single forward pass. + """ + + def __init__( + self, + model_name: str = "BAAI/bge-m3", + dim: int = 1024, + device: str | None = None, + batch_size: int = 32, + ) -> None: + resolved_device = device or ("cuda" if torch.cuda.is_available() else "cpu") + self._model = BGEM3FlagModel( + model_name, + use_fp16=False, + devices=[resolved_device], + ) + self.dim = dim + self.batch_size = batch_size + + # ------------------------------------------------------------------ + # Internal helpers + # ------------------------------------------------------------------ + @staticmethod + def _l2_normalise(matrix: np.ndarray) -> np.ndarray: + """Return ``matrix`` with every row scaled to unit L2 norm. + + Zero-rows are left untouched (we divide by 1 rather than 0 to + avoid ``nan``s — the BGE-M3 encoder never emits zero vectors in + practice, but the guard keeps the function total). + """ + norms = np.linalg.norm(matrix, axis=1, keepdims=True) + norms = np.where(norms == 0.0, 1.0, norms) + return (matrix / norms).astype(np.float32) + + @staticmethod + def _coerce_sparse(weights: list[dict[int | str, float]]) -> list[dict[int, float]]: + """Cast token ids to ``int`` and weights to ``float``. + + FlagEmbedding returns ids as strings in some versions and as + ints in others; we normalise to ``int`` so downstream code can + rely on a stable key type. + """ + return [ + {int(token_id): float(weight) for token_id, weight in row.items()} for row in weights + ] + + # ------------------------------------------------------------------ + # Public API + # ------------------------------------------------------------------ + def encode_dense(self, texts: list[str]) -> np.ndarray: + """Encode ``texts`` to a dense, L2-normalised float32 matrix. + + Returns a ``(0, dim)`` zero-matrix for an empty input. + """ + if not texts: + return np.zeros((0, self.dim), dtype=np.float32) + + result = self._model.encode( + texts, + batch_size=self.batch_size, + return_dense=True, + return_sparse=False, + ) + dense = np.asarray(result["dense_vecs"], dtype=np.float32) + return self._l2_normalise(dense) + + def encode_sparse(self, texts: list[str]) -> list[dict[int, float]]: + """Encode ``texts`` to a list of sparse ``{token_id: weight}`` dicts. + + Returns ``[]`` for an empty input. + """ + if not texts: + return [] + + result = self._model.encode( + texts, + batch_size=self.batch_size, + return_dense=False, + return_sparse=True, + ) + return self._coerce_sparse(result["lexical_weights"]) + + def encode(self, texts: list[str]) -> tuple[np.ndarray, list[dict[int, float]]]: + """Encode ``texts`` to dense **and** sparse in one forward pass. + + Returns ``(zeros((0, dim)), [])`` for an empty input. + """ + if not texts: + return np.zeros((0, self.dim), dtype=np.float32), [] + + result = self._model.encode( + texts, + batch_size=self.batch_size, + return_dense=True, + return_sparse=True, + ) + dense = self._l2_normalise(np.asarray(result["dense_vecs"], dtype=np.float32)) + sparse = self._coerce_sparse(result["lexical_weights"]) + return dense, sparse + + +# --------------------------------------------------------------------------- +# Process-wide singleton +# --------------------------------------------------------------------------- +_embedder: Embedder | None = None + + +def get_embedder() -> Embedder: + """Return the process-wide :class:`Embedder` singleton. + + The first call constructs the embedder using values from + :class:`codex.config.Settings`; subsequent calls return the cached + instance. Tests can reset the cache by setting + ``codex.embed._embedder`` back to ``None``. + """ + global _embedder + if _embedder is None: + from codex.config import Settings + + s = Settings() + _embedder = Embedder(model_name=s.embedding_model, dim=s.embedding_dim) + return _embedder diff --git a/tests/embed/__init__.py b/tests/embed/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/embed/test_embedder.py b/tests/embed/test_embedder.py new file mode 100644 index 0000000..16a62c7 --- /dev/null +++ b/tests/embed/test_embedder.py @@ -0,0 +1,186 @@ +"""Tests for codex.embed. + +The real :class:`FlagEmbedding.BGEM3FlagModel` would download 2 GB of +weights and load them onto the device. We mock it out so the suite +runs offline in milliseconds. +""" + +from __future__ import annotations + +from typing import Any + +import numpy as np +import pytest + +import codex.embed as embed_module +from codex.embed import Embedder, get_embedder + +# --------------------------------------------------------------------------- +# Fakes / fixtures +# --------------------------------------------------------------------------- + + +class _FakeBGEModel: + """Stand-in for :class:`BGEM3FlagModel` that records calls. + + ``encode()`` returns deterministic fake vectors so the L2-norm + assertions are exact. The ``call_count`` attribute lets tests + verify the single-forward-pass invariant of :meth:`Embedder.encode`. + """ + + DIM = 1024 + + def __init__(self, *args: Any, **kwargs: Any) -> None: + self.init_args = args + self.init_kwargs = kwargs + self.call_count = 0 + self.last_kwargs: dict[str, Any] | None = None + + def encode( + self, + sentences: list[str], + **kwargs: Any, + ) -> dict[str, Any]: + self.call_count += 1 + self.last_kwargs = kwargs + n = len(sentences) + rng = np.random.default_rng(seed=42) + return { + "dense_vecs": rng.random((n, self.DIM)).astype(np.float32), + "lexical_weights": [{0: 0.5, 7: 0.25} for _ in range(n)], + } + + +@pytest.fixture +def fake_model(monkeypatch: pytest.MonkeyPatch) -> type[_FakeBGEModel]: + """Replace ``BGEM3FlagModel`` in ``codex.embed`` with the fake class.""" + monkeypatch.setattr(embed_module, "BGEM3FlagModel", _FakeBGEModel) + return _FakeBGEModel + + +@pytest.fixture(autouse=True) +def reset_singleton() -> None: + """Reset the module-level singleton between tests.""" + embed_module._embedder = None + + +# --------------------------------------------------------------------------- +# encode_dense +# --------------------------------------------------------------------------- + + +def test_encode_dense_shape_dtype_and_norm(fake_model: type[_FakeBGEModel]) -> None: + """Dense output is (N, dim) float32 with unit L2 rows.""" + e = Embedder() + out = e.encode_dense(["a", "b"]) + + assert out.shape == (2, 1024) + assert out.dtype == np.float32 + norms = np.linalg.norm(out, axis=1) + np.testing.assert_allclose(norms, [1.0, 1.0], atol=1e-5) + + +def test_encode_dense_empty_input_returns_empty_matrix( + fake_model: type[_FakeBGEModel], +) -> None: + """Empty input -> (0, dim) without invoking the model.""" + e = Embedder() + # Reach into the wrapped model to verify it is not called. + fake = e._model + assert isinstance(fake, _FakeBGEModel) + + out = e.encode_dense([]) + + assert out.shape == (0, 1024) + assert out.dtype == np.float32 + assert fake.call_count == 0 + + +# --------------------------------------------------------------------------- +# encode_sparse +# --------------------------------------------------------------------------- + + +def test_encode_sparse_returns_list_of_dicts(fake_model: type[_FakeBGEModel]) -> None: + """Sparse output has one dict per input string.""" + e = Embedder() + out = e.encode_sparse(["a", "b"]) + + assert len(out) == 2 + for row in out: + assert isinstance(row, dict) + for key, value in row.items(): + assert isinstance(key, int) + assert isinstance(value, float) + + +def test_encode_sparse_empty_input_returns_empty_list( + fake_model: type[_FakeBGEModel], +) -> None: + """Empty input -> [] without invoking the model.""" + e = Embedder() + fake = e._model + assert isinstance(fake, _FakeBGEModel) + + out = e.encode_sparse([]) + + assert out == [] + assert fake.call_count == 0 + + +# --------------------------------------------------------------------------- +# encode (combined) +# --------------------------------------------------------------------------- + + +def test_encode_returns_tuple_and_uses_single_forward_pass( + fake_model: type[_FakeBGEModel], +) -> None: + """encode() must issue exactly one model.encode() call for efficiency.""" + e = Embedder() + fake = e._model + assert isinstance(fake, _FakeBGEModel) + + dense, sparse = e.encode(["a"]) + + assert isinstance(dense, np.ndarray) + assert dense.shape == (1, 1024) + assert dense.dtype == np.float32 + np.testing.assert_allclose(np.linalg.norm(dense, axis=1), [1.0], atol=1e-5) + + assert isinstance(sparse, list) + assert len(sparse) == 1 + assert isinstance(sparse[0], dict) + + assert fake.call_count == 1 + assert fake.last_kwargs is not None + assert fake.last_kwargs.get("return_dense") is True + assert fake.last_kwargs.get("return_sparse") is True + + +def test_encode_empty_input(fake_model: type[_FakeBGEModel]) -> None: + """Empty input -> ((0, dim), []).""" + e = Embedder() + fake = e._model + assert isinstance(fake, _FakeBGEModel) + + dense, sparse = e.encode([]) + + assert dense.shape == (0, 1024) + assert dense.dtype == np.float32 + assert sparse == [] + assert fake.call_count == 0 + + +# --------------------------------------------------------------------------- +# Singleton +# --------------------------------------------------------------------------- + + +def test_get_embedder_returns_singleton(fake_model: type[_FakeBGEModel]) -> None: + """Two calls return the same object.""" + first = get_embedder() + second = get_embedder() + + assert first is second + assert isinstance(first, Embedder)