diff --git a/CHANGELOG.md b/CHANGELOG.md index bda57732..5a549e05 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -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. +### Changed + +- `import_documents` embeds chunks across the whole batch in one pass instead of per document. + ### Removed - `wix` evaluation dataset and its reference config `evaluations/configs/wix.yaml`. diff --git a/haiku_rag_slim/haiku/rag/client/documents.py b/haiku_rag_slim/haiku/rag/client/documents.py index 8facaba2..1f5c4aae 100644 --- a/haiku_rag_slim/haiku/rag/client/documents.py +++ b/haiku_rag_slim/haiku/rag/client/documents.py @@ -299,10 +299,21 @@ async def _store_documents_with_chunks( Embeds any chunks that lack embeddings, then writes the documents, chunks, and document_items tables once apiece. Restores all tables on any failure. """ - embedded: list[list[Chunk]] = [ - await ensure_chunks_embedded(client._config, chunks, client.embedder) + missing = [ + chunk 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(): return [extract_items("", d) for _, _, d in prepared] diff --git a/tests/test_client.py b/tests/test_client.py index db2da16f..939d1f07 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -18,6 +18,7 @@ from haiku.rag.client.documents import ( check_source_accessible, ) from haiku.rag.config import Config +from haiku.rag.embeddings import EmbedderWrapper from haiku.rag.store.compression import decompress_json from haiku.rag.store.models.chunk import Chunk 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 +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( temp_db_path, ):