Add batching to embeddings for huge documents
This commit is contained in:
parent
68c3fa0f79
commit
32258d4e2a
3 changed files with 46 additions and 2 deletions
|
|
@ -12,6 +12,7 @@
|
||||||
- **Search performance**: Avoid loading full document blobs (docling_document, content) during search — use column projection to fetch only needed metadata (id, uri, title, metadata)
|
- **Search performance**: Avoid loading full document blobs (docling_document, content) during search — use column projection to fetch only needed metadata (id, uri, title, metadata)
|
||||||
- **Context expansion performance**: Load only docling columns during expand_context (skip content blob), and only when doc_item_refs exist
|
- **Context expansion performance**: Load only docling columns during expand_context (skip content blob), and only when doc_item_refs exist
|
||||||
- **Chunk expansion performance**: Fetch only chunks in the needed order range during context expansion instead of all chunks for a document
|
- **Chunk expansion performance**: Fetch only chunks in the needed order range during context expansion instead of all chunks for a document
|
||||||
|
- **Embedding batching**: Batch embedding calls in groups of 512 to avoid request size limits and timeouts with large documents
|
||||||
|
|
||||||
## [0.36.3] - 2026-04-01
|
## [0.36.3] - 2026-04-01
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -53,6 +53,9 @@ def contextualize(chunks: list["Chunk"]) -> list[str]:
|
||||||
return texts
|
return texts
|
||||||
|
|
||||||
|
|
||||||
|
EMBEDDING_BATCH_SIZE = 512
|
||||||
|
|
||||||
|
|
||||||
async def embed_chunks(
|
async def embed_chunks(
|
||||||
chunks: list["Chunk"], config: AppConfig = Config
|
chunks: list["Chunk"], config: AppConfig = Config
|
||||||
) -> list["Chunk"]:
|
) -> list["Chunk"]:
|
||||||
|
|
@ -61,6 +64,9 @@ async def embed_chunks(
|
||||||
Contextualizes chunks (prepends headings) before embedding for better
|
Contextualizes chunks (prepends headings) before embedding for better
|
||||||
semantic search. Returns new Chunk objects with embeddings set.
|
semantic search. Returns new Chunk objects with embeddings set.
|
||||||
|
|
||||||
|
Embeddings are generated in batches to avoid request size limits
|
||||||
|
and timeouts with large document sets.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
chunks: List of chunks to embed.
|
chunks: List of chunks to embed.
|
||||||
config: Configuration for embedder selection.
|
config: Configuration for embedder selection.
|
||||||
|
|
@ -75,7 +81,13 @@ async def embed_chunks(
|
||||||
|
|
||||||
embedder = get_embedder(config)
|
embedder = get_embedder(config)
|
||||||
texts = contextualize(chunks)
|
texts = contextualize(chunks)
|
||||||
embeddings = await embedder.embed_documents(texts)
|
|
||||||
|
# Batch embedding calls to avoid request size limits
|
||||||
|
all_embeddings: list[list[float]] = []
|
||||||
|
for i in range(0, len(texts), EMBEDDING_BATCH_SIZE):
|
||||||
|
batch = texts[i : i + EMBEDDING_BATCH_SIZE]
|
||||||
|
batch_embeddings = await embedder.embed_documents(batch)
|
||||||
|
all_embeddings.extend(batch_embeddings)
|
||||||
|
|
||||||
return [
|
return [
|
||||||
Chunk(
|
Chunk(
|
||||||
|
|
@ -89,7 +101,7 @@ async def embed_chunks(
|
||||||
document_meta=chunk.document_meta,
|
document_meta=chunk.document_meta,
|
||||||
embedding=embedding,
|
embedding=embedding,
|
||||||
)
|
)
|
||||||
for chunk, embedding in zip(chunks, embeddings)
|
for chunk, embedding in zip(chunks, all_embeddings)
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -160,6 +160,37 @@ async def test_embed_chunks_empty_list():
|
||||||
assert result == []
|
assert result == []
|
||||||
|
|
||||||
|
|
||||||
|
async def test_embed_chunks_batches_large_inputs(monkeypatch):
|
||||||
|
"""Test that embed_chunks batches calls when chunk count exceeds batch size."""
|
||||||
|
from haiku.rag.embeddings import EMBEDDING_BATCH_SIZE, EmbedderWrapper
|
||||||
|
|
||||||
|
call_sizes: list[int] = []
|
||||||
|
|
||||||
|
async def tracking_embed(self, texts):
|
||||||
|
call_sizes.append(len(texts))
|
||||||
|
return [[0.1] * 10 for _ in texts]
|
||||||
|
|
||||||
|
monkeypatch.setattr(EmbedderWrapper, "embed_documents", tracking_embed)
|
||||||
|
|
||||||
|
# Create more chunks than one batch
|
||||||
|
num_chunks = EMBEDDING_BATCH_SIZE + 100
|
||||||
|
chunks = [
|
||||||
|
Chunk(id=f"chunk-{i}", content=f"Content {i}", order=i)
|
||||||
|
for i in range(num_chunks)
|
||||||
|
]
|
||||||
|
|
||||||
|
result = await embed_chunks(chunks)
|
||||||
|
|
||||||
|
assert len(result) == num_chunks
|
||||||
|
assert len(call_sizes) == 2
|
||||||
|
assert call_sizes[0] == EMBEDDING_BATCH_SIZE
|
||||||
|
assert call_sizes[1] == 100
|
||||||
|
# Verify order is preserved
|
||||||
|
assert result[0].id == "chunk-0"
|
||||||
|
assert result[-1].id == f"chunk-{num_chunks - 1}"
|
||||||
|
assert all(r.embedding == [0.1] * 10 for r in result)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.vcr()
|
@pytest.mark.vcr()
|
||||||
async def test_embed_chunks_preserves_all_fields(allow_model_requests):
|
async def test_embed_chunks_preserves_all_fields(allow_model_requests):
|
||||||
"""Test that embed_chunks preserves all chunk fields."""
|
"""Test that embed_chunks preserves all chunk fields."""
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue