haiku.rag/tests/test_chunker.py
Yiorgis Gozadinos 6c948a4afd
Chucker & tests
2025-06-15 21:08:06 +02:00

44 lines
1.6 KiB
Python

import pytest
from datasets import Dataset
from haiku.rag.chunker import Chunker
@pytest.mark.asyncio
async def test_chunker(qa_corpus: Dataset):
chunker = Chunker()
doc = qa_corpus[0]["document_extracted"]
chunks = await Chunker().chunk(doc)
# Ensure that the text is split into multiple chunks
assert len(chunks) > 1
# Ensure that each chunk corresponds to roughly Config.CHUNK_SIZE tokens
for chunk in chunks[:-1]:
encoded_tokens = Chunker.encoder.encode(chunk, disallowed_special=())
assert len(encoded_tokens) <= Chunker().chunk_size
assert len(encoded_tokens) > Chunker().chunk_size * 0.9
# Ensure that the last chunk is less than Config.CHUNK_SIZE tokens
assert (
len(Chunker.encoder.encode(chunks[-1], disallowed_special=()))
< Chunker().chunk_size
)
# Test overlap between consecutive chunks
for i in range(len(chunks) - 1):
current_chunk = chunks[i]
next_chunk = chunks[i + 1]
current_tokens = Chunker.encoder.encode(current_chunk, disallowed_special=())
next_tokens = Chunker.encoder.encode(next_chunk, disallowed_special=())
overlap_size = min(chunker.chunk_overlap, len(current_tokens))
current_overlap_tokens = current_tokens[-overlap_size:]
next_overlap_tokens = next_tokens[:overlap_size]
# The overlapping tokens should be identical
assert current_overlap_tokens == next_overlap_tokens
assert len(current_overlap_tokens) == min(
chunker.chunk_overlap, len(current_tokens)
)