DocQuest / core /chunking.py
harao-ml's picture
Upload 17 files
30ed2ac verified
Raw
History Blame Contribute Delete
3.07 kB
import logging
from typing import List, Optional
import tiktoken
from langchain.text_splitter import RecursiveCharacterTextSplitter
from langchain.schema import Document
from config import CHUNK_SIZE_TOKENS, CHUNK_OVERLAP_TOKENS
logger = logging.getLogger(__name__)
class ChunkingPipeline:
def __init__(
self,
chunk_size: int = CHUNK_SIZE_TOKENS,
chunk_overlap: int = CHUNK_OVERLAP_TOKENS,
):
self._encoder = tiktoken.get_encoding("cl100k_base")
def _token_len(text: str) -> int:
return len(self._encoder.encode(text, disallowed_special=()))
self._token_len = _token_len
self._splitter = RecursiveCharacterTextSplitter(
chunk_size=chunk_size,
chunk_overlap=chunk_overlap,
length_function=_token_len,
separators=[
"\n\n\n",
"\n\n",
"\n",
". ",
"? ",
"! ",
"; ",
", ",
" ",
"",
],
keep_separator=False,
)
self.chunk_size = chunk_size
self.chunk_overlap = chunk_overlap
def split(self, text: str, metadata: Optional[dict] = None) -> List[Document]:
if not text or not text.strip():
raise ValueError("Cannot chunk empty text.")
base_meta = dict(metadata) if metadata else {}
raw_chunks: List[Document] = self._splitter.create_documents(
texts=[text],
metadatas=[base_meta],
)
if not raw_chunks:
raw_chunks = [Document(page_content=text, metadata=base_meta)]
total = len(raw_chunks)
char_offset = 0
for idx, chunk in enumerate(raw_chunks):
token_count = self._token_len(chunk.page_content)
chunk.metadata.update(
{
"chunk_id": idx,
"total_chunks": total,
"chunk_tokens": token_count,
"start_char": char_offset,
}
)
char_offset += len(chunk.page_content)
return raw_chunks
def count_tokens(self, text: str) -> int:
return self._token_len(text)
def estimate_chunks(self, text: str) -> int:
total_tokens = self.count_tokens(text)
if total_tokens <= self.chunk_size:
return 1
stride = self.chunk_size - self.chunk_overlap
return max(1, int((total_tokens - self.chunk_overlap) / stride) + 1)
def fits_in_single_call(self, text: str, context_window: int = 6000) -> bool:
return self.count_tokens(text) <= context_window
def get_strategy(self, text: str) -> str:
n = self.count_tokens(text)
if n <= 6000:
return "single"
elif n <= 80000:
return "map_reduce"
else:
return "sample"
def __repr__(self) -> str:
return f"ChunkingPipeline(chunk_size={self.chunk_size}, overlap={self.chunk_overlap})"