feat: merge F-04 embed layer (BGE-M3 dense+sparse via FlagEmbedding)
This commit is contained in:
155
codex/embed.py
Normal file
155
codex/embed.py
Normal file
@@ -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
|
||||
Reference in New Issue
Block a user