Build one reranker for the databases searched together

Every client covering a database built its own, so a set of five loaded the same
local model weights five times over, beside the federator's. A client covering a
database for another borrows that one's, and only its owner closes it.
This commit is contained in:
Yiorgis Gozadinos 2026-08-27 16:01:44 +03:00
parent 64b1096865
commit 8791f12c37
No known key found for this signature in database
3 changed files with 63 additions and 13 deletions

View file

@ -164,6 +164,8 @@ class HaikuRAG:
self._read_only = read_only
self._requested_sources = sources
self._clients: dict[str, HaikuRAG] = {}
# The client this one covers a database for, whose reranker it borrows.
self._lender: HaikuRAG | None = None
self._scope: DatabaseScope | None = None
self._session: SingleDatabaseSession | FederatedSession | None = None
self._owns_session = True
@ -280,13 +282,21 @@ class HaikuRAG:
return get_embedder(config=self._config)
return self.store.embedder
@cached_property
@property
def reranker(self) -> "RerankerBase | None":
"""The configured reranker, built once and reused across searches.
None when reranking is disabled. Local rerankers load model weights on
construction, so building per search would reload them on every query.
construction, so one per database in a set would load the same weights
that many times over: a client covering a database for another borrows
that one's, built on the first query to reach any of them.
"""
if self._lender is not None:
return self._lender.reranker
return self._own_reranker
@cached_property
def _own_reranker(self) -> "RerankerBase | None":
return get_reranker(config=self._config)
def _resolve_scope(self) -> DatabaseScope:
@ -367,7 +377,7 @@ class HaikuRAG:
"""The cached client borrowing this session, made once and kept."""
facade = self._clients.get(name)
if facade is None:
facade = HaikuRAG._from_session(session)
facade = HaikuRAG._from_session(session, lender=self)
self._clients[name] = facade
return facade
@ -396,13 +406,20 @@ class HaikuRAG:
return client
@classmethod
def _from_session(cls, session: SingleDatabaseSession) -> "HaikuRAG":
"""A client over a database another session opened and will close."""
def _from_session(
cls, session: SingleDatabaseSession, lender: "HaikuRAG | None" = None
) -> "HaikuRAG":
"""A client over a database another session opened and will close.
`lender` is the client that opened it, whose reranker this one borrows
rather than building a second copy of the same model.
"""
client = cls(
session.db_path, config=session.config, read_only=session.read_only
)
client._session = session
client._owns_session = False
client._lender = lender
return client
def _require_one_embedder(self, clients: "list[HaikuRAG]") -> None:
@ -440,8 +457,9 @@ class HaikuRAG:
# a federating client that answered no query has nothing open and no
# store either.
if isinstance(self._session, FederatedSession):
# The wrappers built over covered databases hold rerankers of their
# own; the databases themselves are the federated session's to close.
# The wrappers only discard what they cached: the databases are the
# federated session's to close, and the reranker they searched with
# is this client's, closed below.
for facade in self._clients.values():
await facade._release_own()
self._clients.clear()
@ -449,7 +467,7 @@ class HaikuRAG:
# The set shares this client's embedder and reranker, so this is the
# only place they are closed — and only if anything built them.
await self._aclose_cached("embedder")
await self._aclose_cached("reranker")
await self._aclose_cached("_own_reranker")
return False
if not self._owns_session:
await self._release_own()
@ -465,11 +483,11 @@ class HaikuRAG:
"""Release what this client built, leaving the database to its owner.
The embedder belongs to the store and is closed with it, so the cached
reference is only discarded. The reranker is this client's own, built on
its first text query.
reference is only discarded. A borrowed reranker belongs to its lender,
so only one built here is closed.
"""
self.__dict__.pop("embedder", None)
await self._aclose_cached("reranker")
await self._aclose_cached("_own_reranker")
async def _aclose_cached(self, name: str) -> None:
"""Close a cached_property this client materialized, and discard it.

View file

@ -240,11 +240,43 @@ class TestBorrowedDatabases:
async with HaikuRAG(config=config) as rag:
(alpha,) = await rag.clients_for(["alpha"])
alpha.__dict__["reranker"] = Reranker()
alpha.__dict__["_own_reranker"] = Reranker()
assert closed == ["reranker"]
class TestSharingTheReranker:
@pytest.mark.asyncio
async def test_the_set_builds_and_closes_one_reranker(self, tmp_path, monkeypatch):
"""A local reranker loads model weights, so one per database in a set
would load the same weights that many times."""
import haiku.rag.client as client_module
built: list[object] = []
closed: list[object] = []
class Reranker:
def __init__(self):
built.append(self)
async def aclose(self):
closed.append(self)
monkeypatch.setattr(client_module, "get_reranker", lambda config: Reranker())
config = _config(tmp_path, ["alpha", "beta"])
await _seed(config, "alpha", ["alpha document about cats"])
await _seed(config, "beta", ["beta document about cats"])
async with HaikuRAG(config=config) as rag:
alpha, beta = await rag.clients_for(["alpha", "beta"])
assert alpha.reranker is beta.reranker is rag.reranker
assert len(built) == 1
assert closed == built
class TestLazyOpening:
@pytest.mark.asyncio
async def test_entering_opens_nothing(self, tmp_path):

View file

@ -452,7 +452,7 @@ async def test_search_attaches_picture_bytes_for_multimodal_reranker(
],
)
rag.chunk_repository.search = fake_chunk_search # type: ignore[method-assign]
rag.__dict__["reranker"] = StubReranker()
rag.__dict__["_own_reranker"] = StubReranker()
rag._config.reranking.multimodal = multimodal
await rag.search("totals", include_images=False)