Batch embeddings across import_documents batches
_store_documents_with_chunks embedded each document's chunks in its own embed_chunks call; chunks missing embeddings are now flattened across the whole batch, embedded in one pass honoring embeddings.batch_size, and assigned back positionally.
This commit is contained in:
parent
b07426a883
commit
cc04f92f28
3 changed files with 106 additions and 2 deletions
|
|
@ -5,6 +5,10 @@
|
||||||
|
|
||||||
- `evaluations run --filter/-f CLAUSE`: SQL `WHERE` clause over document columns, applied to the retrieval benchmark's searches and to every capability search during QA. Recorded as `document_filter` in experiment metadata.
|
- `evaluations run --filter/-f CLAUSE`: SQL `WHERE` clause over document columns, applied to the retrieval benchmark's searches and to every capability search during QA. Recorded as `document_filter` in experiment metadata.
|
||||||
|
|
||||||
|
### Changed
|
||||||
|
|
||||||
|
- `import_documents` embeds chunks across the whole batch in one pass instead of per document.
|
||||||
|
|
||||||
### Removed
|
### Removed
|
||||||
|
|
||||||
- `wix` evaluation dataset and its reference config `evaluations/configs/wix.yaml`.
|
- `wix` evaluation dataset and its reference config `evaluations/configs/wix.yaml`.
|
||||||
|
|
|
||||||
|
|
@ -299,10 +299,21 @@ async def _store_documents_with_chunks(
|
||||||
Embeds any chunks that lack embeddings, then writes the documents, chunks,
|
Embeds any chunks that lack embeddings, then writes the documents, chunks,
|
||||||
and document_items tables once apiece. Restores all tables on any failure.
|
and document_items tables once apiece. Restores all tables on any failure.
|
||||||
"""
|
"""
|
||||||
embedded: list[list[Chunk]] = [
|
missing = [
|
||||||
await ensure_chunks_embedded(client._config, chunks, client.embedder)
|
chunk
|
||||||
for _, chunks, _ in prepared
|
for _, chunks, _ in prepared
|
||||||
|
for chunk in chunks
|
||||||
|
if chunk.embedding is None
|
||||||
]
|
]
|
||||||
|
if missing:
|
||||||
|
from haiku.rag.embeddings import embed_chunks
|
||||||
|
|
||||||
|
embedded_flat = await embed_chunks(missing, client.embedder, client._config)
|
||||||
|
# Assign positionally: duplicate chunk texts across documents make a
|
||||||
|
# content-keyed lookup ambiguous.
|
||||||
|
for chunk, with_embedding in zip(missing, embedded_flat):
|
||||||
|
chunk.embedding = with_embedding.embedding
|
||||||
|
embedded: list[list[Chunk]] = [chunks for _, chunks, _ in prepared]
|
||||||
|
|
||||||
def _extract_all_items():
|
def _extract_all_items():
|
||||||
return [extract_items("", d) for _, _, d in prepared]
|
return [extract_items("", d) for _, _, d in prepared]
|
||||||
|
|
|
||||||
|
|
@ -18,6 +18,7 @@ from haiku.rag.client.documents import (
|
||||||
check_source_accessible,
|
check_source_accessible,
|
||||||
)
|
)
|
||||||
from haiku.rag.config import Config
|
from haiku.rag.config import Config
|
||||||
|
from haiku.rag.embeddings import EmbedderWrapper
|
||||||
from haiku.rag.store.compression import decompress_json
|
from haiku.rag.store.compression import decompress_json
|
||||||
from haiku.rag.store.models.chunk import Chunk
|
from haiku.rag.store.models.chunk import Chunk
|
||||||
from haiku.rag.store.models.document import Document
|
from haiku.rag.store.models.document import Document
|
||||||
|
|
@ -901,6 +902,94 @@ async def test_client_import_documents_empty(temp_db_path):
|
||||||
assert after == before
|
assert after == before
|
||||||
|
|
||||||
|
|
||||||
|
class _CountingEmbedder(EmbedderWrapper):
|
||||||
|
def __init__(self, vector_dim: int):
|
||||||
|
super().__init__(None, vector_dim)
|
||||||
|
self.batches: list[int] = []
|
||||||
|
|
||||||
|
async def embed_documents(self, texts: list[str]) -> list[list[float]]:
|
||||||
|
self.batches.append(len(texts))
|
||||||
|
return [[0.1] * self.vector_dim for _ in texts]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_client_import_documents_batches_embeddings(temp_db_path):
|
||||||
|
"""Chunks missing embeddings are embedded in one pass across the whole
|
||||||
|
batch, not one embedder call per document. Duplicate chunk texts across
|
||||||
|
documents keep their per-document embeddings."""
|
||||||
|
dim = Config.embeddings.model.vector_dim
|
||||||
|
embedder = _CountingEmbedder(dim)
|
||||||
|
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
|
client.store.embedder = embedder
|
||||||
|
imports = [
|
||||||
|
DocumentImport(
|
||||||
|
docling_document=_docling_doc(name, text),
|
||||||
|
chunks=[Chunk(content=text, order=0)],
|
||||||
|
uri=f"mem://{name}",
|
||||||
|
title=name,
|
||||||
|
)
|
||||||
|
for name, text in (
|
||||||
|
("a", "Alpha document body"),
|
||||||
|
("b", "Beta document body"),
|
||||||
|
("c", "Alpha document body"),
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
docs = await client.import_documents(imports)
|
||||||
|
|
||||||
|
assert embedder.batches == [3]
|
||||||
|
rows = await (
|
||||||
|
client.store.chunks_table.query()
|
||||||
|
.select(["document_id", "vector"])
|
||||||
|
.to_list()
|
||||||
|
)
|
||||||
|
assert {row["document_id"] for row in rows} == {doc.id for doc in docs}
|
||||||
|
assert all(len(row["vector"]) == dim for row in rows)
|
||||||
|
|
||||||
|
|
||||||
|
async def test_client_import_documents_mixed_embeddings(temp_db_path):
|
||||||
|
"""Pre-embedded chunks keep their vectors; only the unembedded ones go
|
||||||
|
through the embedder, in one batch."""
|
||||||
|
dim = Config.embeddings.model.vector_dim
|
||||||
|
embedder = _CountingEmbedder(dim)
|
||||||
|
|
||||||
|
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||||
|
client.store.embedder = embedder
|
||||||
|
pre_embedded = DocumentImport(
|
||||||
|
docling_document=_docling_doc("b", "Beta document body"),
|
||||||
|
chunks=[
|
||||||
|
Chunk(content="Beta document body", embedding=[0.5] * dim, order=0)
|
||||||
|
],
|
||||||
|
uri="mem://b",
|
||||||
|
title="b",
|
||||||
|
)
|
||||||
|
unembedded = [
|
||||||
|
DocumentImport(
|
||||||
|
docling_document=_docling_doc(name, text),
|
||||||
|
chunks=[Chunk(content=text, order=0)],
|
||||||
|
uri=f"mem://{name}",
|
||||||
|
title=name,
|
||||||
|
)
|
||||||
|
for name, text in (("a", "Alpha document body"), ("c", "Gamma body"))
|
||||||
|
]
|
||||||
|
|
||||||
|
docs = await client.import_documents(
|
||||||
|
[unembedded[0], pre_embedded, unembedded[1]]
|
||||||
|
)
|
||||||
|
|
||||||
|
assert embedder.batches == [2]
|
||||||
|
by_uri = {doc.uri: doc.id for doc in docs}
|
||||||
|
rows = await (
|
||||||
|
client.store.chunks_table.query()
|
||||||
|
.select(["document_id", "vector"])
|
||||||
|
.to_list()
|
||||||
|
)
|
||||||
|
vectors = {row["document_id"]: list(row["vector"]) for row in rows}
|
||||||
|
assert vectors[by_uri["mem://b"]] == pytest.approx([0.5] * dim)
|
||||||
|
assert vectors[by_uri["mem://a"]] == pytest.approx([0.1] * dim)
|
||||||
|
assert vectors[by_uri["mem://c"]] == pytest.approx([0.1] * dim)
|
||||||
|
|
||||||
|
|
||||||
async def test_client_update_document_replaces_rows_with_bounded_versions(
|
async def test_client_update_document_replaces_rows_with_bounded_versions(
|
||||||
temp_db_path,
|
temp_db_path,
|
||||||
):
|
):
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue