Embed a search query once for the whole selection
Each database owns an embedder, so embedding per database cost a round trip each. One database still embeds inside the repository, which returns early for a filter that matches nothing.
This commit is contained in:
parent
c3182d0fb6
commit
16e0add64b
6 changed files with 87 additions and 47 deletions
|
|
@ -5,7 +5,7 @@
|
|||
### Added
|
||||
|
||||
- `api_key` on model and embedding-model config, overriding the provider's environment variable. Honored on the `openai` and `ollama` providers, `vllm` embedders and rerankers, the picture-description VLM endpoint, and `doctor`'s endpoint probes; other providers raise.
|
||||
- `lancedb.databases`: a name-to-location mapping for searching several databases at once, mutually exclusive with `lancedb.uri`. `client.search(..., sources=[...])` selects which to search, `sources=None` searches all of them, and `SearchResult.source` carries the configured name a result came from. `Document.source` names it on a document from a listing or a lookup. Candidates are fused by the configured reranker over the union, or by reciprocal rank fusion when none is configured. Databases searched together must have been written with the same embedder; two that disagree raise `ConfigMismatchError`. `SearchResult.format_for_agent` names the database, so the model can attribute evidence to one while it answers. `haiku-rag search`, `ask`, `analyze` and `chat` cover the configured set and label each result and citation with its database; every other command works on one, named with `--database NAME` or `--db PATH`.
|
||||
- `lancedb.databases`: a name-to-location mapping for searching several databases at once, mutually exclusive with `lancedb.uri`. `client.search(..., sources=[...])` selects which to search, `sources=None` searches all of them, and `SearchResult.source` carries the configured name a result came from. `Document.source` names it on a document from a listing or a lookup. Candidates are fused by the configured reranker over the union, or by reciprocal rank fusion when none is configured. Databases searched together must have been written with the same embedder; two that disagree raise `ConfigMismatchError`. The query is embedded once for the whole selection. `SearchResult.format_for_agent` names the database, so the model can attribute evidence to one while it answers. `haiku-rag search`, `ask`, `analyze` and `chat` cover the configured set and label each result and citation with its database; every other command works on one, named with `--database NAME` or `--db PATH`.
|
||||
- `client.ask(..., sources=[...])` asks across the selected databases, and `Citation.source` names the one a cited chunk came from. The cite fallback for an id absent from the run's results looks only in the selected databases, so a question scoped to some cannot cite another.
|
||||
- `client.analyze(..., sources=[...])` analyzes across the selected databases: the sandbox mounts their documents under one flat `/documents/{id}/` namespace, resolving each id to the database holding it, and in-code `search()` covers the same selection.
|
||||
|
||||
|
|
|
|||
|
|
@ -39,8 +39,19 @@ async def search(
|
|||
if limit is None:
|
||||
limit = client._config.search.limit
|
||||
|
||||
resolved = search_type or ("hybrid" if isinstance(query, str) else "vector")
|
||||
# One database embeds inside the repository, which returns early for a filter
|
||||
# that matches nothing, so a text query that finds no documents never embeds.
|
||||
query_vector = (
|
||||
None if isinstance(query, str) else await _embed_query(client, query, resolved)
|
||||
)
|
||||
candidates = await _fetch(
|
||||
client, query, _fetch_limit(client, query, limit), search_type, filter
|
||||
client,
|
||||
query,
|
||||
_fetch_limit(client, query, limit),
|
||||
resolved,
|
||||
filter,
|
||||
query_vector,
|
||||
)
|
||||
chunk_results = await _rank(client, query, candidates, limit)
|
||||
|
||||
|
|
@ -77,10 +88,17 @@ async def search_sources(
|
|||
selected = await client.clients_for(names)
|
||||
client._require_one_embedder(selected)
|
||||
|
||||
# One over-fetch decision, and one reranker, for the whole set.
|
||||
# One over-fetch decision, one query vector, and one reranker, for the whole
|
||||
# set. The databases in a selection share an embedder, so the vector is the
|
||||
# same wherever it is computed.
|
||||
fetch_limit = _fetch_limit(client, query, limit)
|
||||
resolved = search_type or ("hybrid" if isinstance(query, str) else "vector")
|
||||
query_vector = await _embed_query(selected[0], query, resolved)
|
||||
per_source = await asyncio.gather(
|
||||
*(_fetch(c, query, fetch_limit, search_type, filter) for c in selected)
|
||||
*(
|
||||
_fetch(c, query, fetch_limit, resolved, filter, query_vector)
|
||||
for c in selected
|
||||
)
|
||||
)
|
||||
|
||||
ranked = await _fuse(client, selected, query, per_source, limit)
|
||||
|
|
@ -180,23 +198,19 @@ def _fetch_limit(
|
|||
return limit * _RERANK_OVERFETCH if client.reranker else limit
|
||||
|
||||
|
||||
async def _fetch(
|
||||
client: "HaikuRAG",
|
||||
query: "str | bytes | PILImage.Image",
|
||||
limit: int,
|
||||
search_type: SearchType | None,
|
||||
filter: str | None,
|
||||
) -> list[tuple[Chunk, float]]:
|
||||
"""Candidates from one database, ranked by that database.
|
||||
async def _embed_query(
|
||||
client: "HaikuRAG", query: "str | bytes | PILImage.Image", search_type: SearchType
|
||||
) -> list[float] | None:
|
||||
"""The query as a vector, or None when the search needs no vector.
|
||||
|
||||
`limit` is how many to fetch, already including any over-fetch the caller
|
||||
wants. Deciding that here would have each database consult its own reranker,
|
||||
and a local reranker loads model weights per instance.
|
||||
Computed by the caller so that searching several databases embeds once: the
|
||||
databases in a selection share an embedder, and embedding per database costs
|
||||
a round trip each on a remote endpoint.
|
||||
"""
|
||||
if search_type == "fts":
|
||||
return None
|
||||
if isinstance(query, str):
|
||||
if search_type is None:
|
||||
search_type = "hybrid"
|
||||
return await client.chunk_repository.search(query, limit, search_type, filter)
|
||||
return await client.embedder.embed_query(query)
|
||||
|
||||
embedder = client.embedder
|
||||
if not embedder.supports_images:
|
||||
|
|
@ -205,10 +219,27 @@ async def _fetch(
|
|||
"embeddings.model.multimodal: true on a vllm, voyageai, or cohere "
|
||||
"model."
|
||||
)
|
||||
query_vector = await embedder.embed_image(query)
|
||||
return await embedder.embed_image(query)
|
||||
|
||||
|
||||
async def _fetch(
|
||||
client: "HaikuRAG",
|
||||
query: "str | bytes | PILImage.Image",
|
||||
limit: int,
|
||||
search_type: SearchType,
|
||||
filter: str | None,
|
||||
query_vector: list[float] | None,
|
||||
) -> list[tuple[Chunk, float]]:
|
||||
"""Candidates from one database, ranked by that database.
|
||||
|
||||
`limit` is how many to fetch, already including any over-fetch the caller
|
||||
wants. Deciding that here would have each database consult its own reranker,
|
||||
and a local reranker loads model weights per instance.
|
||||
"""
|
||||
return await client.chunk_repository.search(
|
||||
query="",
|
||||
query=query if isinstance(query, str) else "",
|
||||
limit=limit,
|
||||
search_type=search_type,
|
||||
filter=filter,
|
||||
query_vector=query_vector,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -213,7 +213,9 @@ class ChunkRepository:
|
|||
limit: Maximum number of results to return.
|
||||
search_type: "vector", "fts", or "hybrid" (default).
|
||||
filter: Optional SQL WHERE clause to filter documents before searching chunks.
|
||||
query_vector: Pre-computed query embedding; forces vector-only search.
|
||||
query_vector: Pre-computed query embedding, used instead of embedding
|
||||
``query``. Searching several databases embeds once and passes it
|
||||
to each.
|
||||
|
||||
Returns:
|
||||
List of (chunk, score) tuples ordered by relevance.
|
||||
|
|
@ -239,37 +241,27 @@ class ChunkRepository:
|
|||
id_list = ", ".join(f"'{d}'" for d in docs_df["id"])
|
||||
chunk_filter = f"document_id IN ({id_list})"
|
||||
|
||||
if query_vector is not None:
|
||||
# Image-as-query: vector-only against the pre-computed embedding.
|
||||
results = (
|
||||
self.store.chunks_table.query()
|
||||
.nearest_to(query_vector)
|
||||
.column("vector")
|
||||
.refine_factor(self.store._config.search.vector_refine_factor)
|
||||
)
|
||||
elif search_type == "vector":
|
||||
query_embedding = await self.embedder.embed_query(query)
|
||||
results = (
|
||||
self.store.chunks_table.query()
|
||||
.nearest_to(query_embedding)
|
||||
.column("vector")
|
||||
.refine_factor(self.store._config.search.vector_refine_factor)
|
||||
)
|
||||
elif search_type == "fts":
|
||||
if search_type == "fts":
|
||||
results = self.store.chunks_table.query().nearest_to_text(
|
||||
query, columns="content_fts"
|
||||
)
|
||||
else: # hybrid (default)
|
||||
query_embedding = await self.embedder.embed_query(query)
|
||||
reranker = RRFReranker()
|
||||
else:
|
||||
query_embedding = (
|
||||
query_vector
|
||||
if query_vector is not None
|
||||
else await self.embedder.embed_query(query)
|
||||
)
|
||||
results = (
|
||||
self.store.chunks_table.query()
|
||||
.nearest_to(query_embedding)
|
||||
.column("vector")
|
||||
.nearest_to_text(query, columns="content_fts")
|
||||
.refine_factor(self.store._config.search.vector_refine_factor)
|
||||
.rerank(reranker)
|
||||
)
|
||||
# An image query has no text to match, so it stays vector-only.
|
||||
if search_type != "vector" and query.strip():
|
||||
results = results.nearest_to_text(query, columns="content_fts").rerank(
|
||||
RRFReranker()
|
||||
)
|
||||
|
||||
if chunk_filter is not None:
|
||||
results = results.where(chunk_filter)
|
||||
|
|
|
|||
|
|
@ -338,6 +338,23 @@ class TestLookupByIdentifier:
|
|||
assert await rag.get_document_by_uri("test://nowhere") is None
|
||||
|
||||
|
||||
class TestOneQueryVector:
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_search_embeds_the_query_once_for_the_whole_set(
|
||||
self, tmp_path, query_embedding
|
||||
):
|
||||
"""Each database owns an embedder, so embedding per database costs a
|
||||
round trip each on a remote endpoint."""
|
||||
config = _config(tmp_path, ["alpha", "beta", "gamma"])
|
||||
for name in ("alpha", "beta", "gamma"):
|
||||
await _seed(config, name, [f"{name} one"])
|
||||
|
||||
async with HaikuRAG(config=config, read_only=True) as rag:
|
||||
await rag.search("one")
|
||||
|
||||
assert query_embedding == ["one"]
|
||||
|
||||
|
||||
class TestOneEmbedderAcrossTheSet:
|
||||
"""A set is searched with one query vector, so a database written with
|
||||
another model would answer from a different space."""
|
||||
|
|
|
|||
|
|
@ -386,7 +386,7 @@ async def test_reranker_built_once_across_searches(temp_db_path, monkeypatch):
|
|||
|
||||
monkeypatch.setattr("haiku.rag.client.get_reranker", fake_get_reranker)
|
||||
|
||||
async def fake_chunk_search(query, limit, search_type, filter):
|
||||
async def fake_chunk_search(query, limit, search_type, filter, query_vector):
|
||||
return [(Chunk(content="x", metadata={}), 0.5)]
|
||||
|
||||
async with HaikuRAG(temp_db_path, create=True) as rag:
|
||||
|
|
@ -434,7 +434,7 @@ async def test_search_attaches_picture_bytes_for_multimodal_reranker(
|
|||
metadata={"doc_item_refs": ["#/pictures/1"], "labels": ["picture"]},
|
||||
)
|
||||
|
||||
async def fake_chunk_search(query, limit, search_type, filter):
|
||||
async def fake_chunk_search(query, limit, search_type, filter, query_vector):
|
||||
return [(text_chunk, 0.9), (picture_chunk, 0.8), (detached_chunk, 0.7)]
|
||||
|
||||
async with HaikuRAG(temp_db_path, create=True) as rag:
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import pytest
|
||||
|
||||
from haiku.rag.client import HaikuRAG
|
||||
from haiku.rag.client.search import _fetch, _rank
|
||||
from haiku.rag.client.search import _rank
|
||||
from haiku.rag.store.models import Chunk
|
||||
|
||||
|
||||
|
|
@ -54,7 +54,7 @@ async def test_image_fetch_never_builds_the_reranker(
|
|||
type(rag.embedder), "supports_images", property(lambda self: True)
|
||||
)
|
||||
|
||||
await _fetch(rag, b"image-bytes", 5, None, None)
|
||||
await rag.search(b"image-bytes", limit=5)
|
||||
|
||||
# No over-fetch: nothing will re-rank these.
|
||||
assert seen["limit"] == 5
|
||||
|
|
|
|||
Loading…
Reference in a new issue