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:
parent
64b1096865
commit
8791f12c37
3 changed files with 63 additions and 13 deletions
|
|
@ -164,6 +164,8 @@ class HaikuRAG:
|
||||||
self._read_only = read_only
|
self._read_only = read_only
|
||||||
self._requested_sources = sources
|
self._requested_sources = sources
|
||||||
self._clients: dict[str, HaikuRAG] = {}
|
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._scope: DatabaseScope | None = None
|
||||||
self._session: SingleDatabaseSession | FederatedSession | None = None
|
self._session: SingleDatabaseSession | FederatedSession | None = None
|
||||||
self._owns_session = True
|
self._owns_session = True
|
||||||
|
|
@ -280,13 +282,21 @@ class HaikuRAG:
|
||||||
return get_embedder(config=self._config)
|
return get_embedder(config=self._config)
|
||||||
return self.store.embedder
|
return self.store.embedder
|
||||||
|
|
||||||
@cached_property
|
@property
|
||||||
def reranker(self) -> "RerankerBase | None":
|
def reranker(self) -> "RerankerBase | None":
|
||||||
"""The configured reranker, built once and reused across searches.
|
"""The configured reranker, built once and reused across searches.
|
||||||
|
|
||||||
None when reranking is disabled. Local rerankers load model weights on
|
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)
|
return get_reranker(config=self._config)
|
||||||
|
|
||||||
def _resolve_scope(self) -> DatabaseScope:
|
def _resolve_scope(self) -> DatabaseScope:
|
||||||
|
|
@ -367,7 +377,7 @@ class HaikuRAG:
|
||||||
"""The cached client borrowing this session, made once and kept."""
|
"""The cached client borrowing this session, made once and kept."""
|
||||||
facade = self._clients.get(name)
|
facade = self._clients.get(name)
|
||||||
if facade is None:
|
if facade is None:
|
||||||
facade = HaikuRAG._from_session(session)
|
facade = HaikuRAG._from_session(session, lender=self)
|
||||||
self._clients[name] = facade
|
self._clients[name] = facade
|
||||||
return facade
|
return facade
|
||||||
|
|
||||||
|
|
@ -396,13 +406,20 @@ class HaikuRAG:
|
||||||
return client
|
return client
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _from_session(cls, session: SingleDatabaseSession) -> "HaikuRAG":
|
def _from_session(
|
||||||
"""A client over a database another session opened and will close."""
|
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(
|
client = cls(
|
||||||
session.db_path, config=session.config, read_only=session.read_only
|
session.db_path, config=session.config, read_only=session.read_only
|
||||||
)
|
)
|
||||||
client._session = session
|
client._session = session
|
||||||
client._owns_session = False
|
client._owns_session = False
|
||||||
|
client._lender = lender
|
||||||
return client
|
return client
|
||||||
|
|
||||||
def _require_one_embedder(self, clients: "list[HaikuRAG]") -> None:
|
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
|
# a federating client that answered no query has nothing open and no
|
||||||
# store either.
|
# store either.
|
||||||
if isinstance(self._session, FederatedSession):
|
if isinstance(self._session, FederatedSession):
|
||||||
# The wrappers built over covered databases hold rerankers of their
|
# The wrappers only discard what they cached: the databases are the
|
||||||
# own; the databases themselves are the federated session's to close.
|
# federated session's to close, and the reranker they searched with
|
||||||
|
# is this client's, closed below.
|
||||||
for facade in self._clients.values():
|
for facade in self._clients.values():
|
||||||
await facade._release_own()
|
await facade._release_own()
|
||||||
self._clients.clear()
|
self._clients.clear()
|
||||||
|
|
@ -449,7 +467,7 @@ class HaikuRAG:
|
||||||
# The set shares this client's embedder and reranker, so this is the
|
# The set shares this client's embedder and reranker, so this is the
|
||||||
# only place they are closed — and only if anything built them.
|
# only place they are closed — and only if anything built them.
|
||||||
await self._aclose_cached("embedder")
|
await self._aclose_cached("embedder")
|
||||||
await self._aclose_cached("reranker")
|
await self._aclose_cached("_own_reranker")
|
||||||
return False
|
return False
|
||||||
if not self._owns_session:
|
if not self._owns_session:
|
||||||
await self._release_own()
|
await self._release_own()
|
||||||
|
|
@ -465,11 +483,11 @@ class HaikuRAG:
|
||||||
"""Release what this client built, leaving the database to its owner.
|
"""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
|
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
|
reference is only discarded. A borrowed reranker belongs to its lender,
|
||||||
its first text query.
|
so only one built here is closed.
|
||||||
"""
|
"""
|
||||||
self.__dict__.pop("embedder", None)
|
self.__dict__.pop("embedder", None)
|
||||||
await self._aclose_cached("reranker")
|
await self._aclose_cached("_own_reranker")
|
||||||
|
|
||||||
async def _aclose_cached(self, name: str) -> None:
|
async def _aclose_cached(self, name: str) -> None:
|
||||||
"""Close a cached_property this client materialized, and discard it.
|
"""Close a cached_property this client materialized, and discard it.
|
||||||
|
|
|
||||||
|
|
@ -240,11 +240,43 @@ class TestBorrowedDatabases:
|
||||||
|
|
||||||
async with HaikuRAG(config=config) as rag:
|
async with HaikuRAG(config=config) as rag:
|
||||||
(alpha,) = await rag.clients_for(["alpha"])
|
(alpha,) = await rag.clients_for(["alpha"])
|
||||||
alpha.__dict__["reranker"] = Reranker()
|
alpha.__dict__["_own_reranker"] = Reranker()
|
||||||
|
|
||||||
assert closed == ["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:
|
class TestLazyOpening:
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_entering_opens_nothing(self, tmp_path):
|
async def test_entering_opens_nothing(self, tmp_path):
|
||||||
|
|
|
||||||
|
|
@ -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.chunk_repository.search = fake_chunk_search # type: ignore[method-assign]
|
||||||
rag.__dict__["reranker"] = StubReranker()
|
rag.__dict__["_own_reranker"] = StubReranker()
|
||||||
rag._config.reranking.multimodal = multimodal
|
rag._config.reranking.multimodal = multimodal
|
||||||
|
|
||||||
await rag.search("totals", include_images=False)
|
await rag.search("totals", include_images=False)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue