Underground-Digital/Workflow-Engine
0
1from __future__ import annotations2 3import copy4import logging5import re6from abc import ABC, abstractmethod7from collections.abc import Callable, Collection, Iterable, Sequence, Set8from dataclasses import dataclass9from typing import (10 Any,11 Literal,12 Optional,13 TypedDict,14 TypeVar,15 Union,16)17 18from core.rag.models.document import BaseDocumentTransformer, Document19 20logger = logging.getLogger(__name__)21 22TS = TypeVar("TS", bound="TextSplitter")23 24 25def _split_text_with_regex(text: str, separator: str, keep_separator: bool) -> list[str]:26 # Now that we have the separator, split the text27 if separator:28 if keep_separator:29 # The parentheses in the pattern keep the delimiters in the result.30 _splits = re.split(f"({re.escape(separator)})", text)31 splits = [_splits[i - 1] + _splits[i] for i in range(1, len(_splits), 2)]32 if len(_splits) % 2 != 0:33 splits += _splits[-1:]34 else:35 splits = re.split(separator, text)36 else:37 splits = list(text)38 return [s for s in splits if (s not in {"", "\n"})]39 40 41class TextSplitter(BaseDocumentTransformer, ABC):42 """Interface for splitting text into chunks."""43 44 def __init__(45 self,46 chunk_size: int = 4000,47 chunk_overlap: int = 200,48 length_function: Callable[[str], int] = len,49 keep_separator: bool = False,50 add_start_index: bool = False,51 ) -> None:52 """Create a new TextSplitter.53 54 Args:55 chunk_size: Maximum size of chunks to return56 chunk_overlap: Overlap in characters between chunks57 length_function: Function that measures the length of given chunks58 keep_separator: Whether to keep the separator in the chunks59 add_start_index: If `True`, includes chunk's start index in metadata60 """61 if chunk_overlap > chunk_size:62 raise ValueError(63 f"Got a larger chunk overlap ({chunk_overlap}) than chunk size ({chunk_size}), should be smaller."64 )65 self._chunk_size = chunk_size66 self._chunk_overlap = chunk_overlap67 self._length_function = length_function68 self._keep_separator = keep_separator69 self._add_start_index = add_start_index70 71 @abstractmethod72 def split_text(self, text: str) -> list[str]:73 """Split text into multiple components."""74 75 def create_documents(self, texts: list[str], metadatas: Optional[list[dict]] = None) -> list[Document]:76 """Create documents from a list of texts."""77 _metadatas = metadatas or [{}] * len(texts)78 documents = []79 for i, text in enumerate(texts):80 index = -181 for chunk in self.split_text(text):82 metadata = copy.deepcopy(_metadatas[i])83 if self._add_start_index:84 index = text.find(chunk, index + 1)85 metadata["start_index"] = index86 new_doc = Document(page_content=chunk, metadata=metadata)87 documents.append(new_doc)88 return documents89 90 def split_documents(self, documents: Iterable[Document]) -> list[Document]:91 """Split documents."""92 texts, metadatas = [], []93 for doc in documents:94 texts.append(doc.page_content)95 metadatas.append(doc.metadata)96 return self.create_documents(texts, metadatas=metadatas)97 98 def _join_docs(self, docs: list[str], separator: str) -> Optional[str]:99 text = separator.join(docs)100 text = text.strip()101 if text == "":102 return None103 else:104 return text105 106 def _merge_splits(self, splits: Iterable[str], separator: str, lengths: list[int]) -> list[str]:107 # We now want to combine these smaller pieces into medium size108 # chunks to send to the LLM.109 separator_len = self._length_function(separator)110 111 docs = []112 current_doc: list[str] = []113 total = 0114 index = 0115 for d in splits:116 _len = lengths[index]117 if total + _len + (separator_len if len(current_doc) > 0 else 0) > self._chunk_size:118 if total > self._chunk_size:119 logger.warning(120 f"Created a chunk of size {total}, which is longer than the specified {self._chunk_size}"121 )122 if len(current_doc) > 0:123 doc = self._join_docs(current_doc, separator)124 if doc is not None:125 docs.append(doc)126 # Keep on popping if:127 # - we have a larger chunk than in the chunk overlap128 # - or if we still have any chunks and the length is long129 while total > self._chunk_overlap or (130 total + _len + (separator_len if len(current_doc) > 0 else 0) > self._chunk_size and total > 0131 ):132 total -= self._length_function(current_doc[0]) + (separator_len if len(current_doc) > 1 else 0)133 current_doc = current_doc[1:]134 current_doc.append(d)135 total += _len + (separator_len if len(current_doc) > 1 else 0)136 index += 1137 doc = self._join_docs(current_doc, separator)138 if doc is not None:139 docs.append(doc)140 return docs141 142 @classmethod143 def from_huggingface_tokenizer(cls, tokenizer: Any, **kwargs: Any) -> TextSplitter:144 """Text splitter that uses HuggingFace tokenizer to count length."""145 try:146 from transformers import PreTrainedTokenizerBase147 148 if not isinstance(tokenizer, PreTrainedTokenizerBase):149 raise ValueError("Tokenizer received was not an instance of PreTrainedTokenizerBase")150 151 def _huggingface_tokenizer_length(text: str) -> int:152 return len(tokenizer.encode(text))153 154 except ImportError:155 raise ValueError(156 "Could not import transformers python package. Please install it with `pip install transformers`."157 )158 return cls(length_function=_huggingface_tokenizer_length, **kwargs)159 160 @classmethod161 def from_tiktoken_encoder(162 cls: type[TS],163 encoding_name: str = "gpt2",164 model_name: Optional[str] = None,165 allowed_special: Union[Literal["all"], Set[str]] = set(),166 disallowed_special: Union[Literal["all"], Collection[str]] = "all",167 **kwargs: Any,168 ) -> TS:169 """Text splitter that uses tiktoken encoder to count length."""170 try:171 import tiktoken172 except ImportError:173 raise ImportError(174 "Could not import tiktoken python package. "175 "This is needed in order to calculate max_tokens_for_prompt. "176 "Please install it with `pip install tiktoken`."177 )178 179 if model_name is not None:180 enc = tiktoken.encoding_for_model(model_name)181 else:182 enc = tiktoken.get_encoding(encoding_name)183 184 def _tiktoken_encoder(text: str) -> int:185 return len(186 enc.encode(187 text,188 allowed_special=allowed_special,189 disallowed_special=disallowed_special,190 )191 )192 193 if issubclass(cls, TokenTextSplitter):194 extra_kwargs = {195 "encoding_name": encoding_name,196 "model_name": model_name,197 "allowed_special": allowed_special,198 "disallowed_special": disallowed_special,199 }200 kwargs = {**kwargs, **extra_kwargs}201 202 return cls(length_function=_tiktoken_encoder, **kwargs)203 204 def transform_documents(self, documents: Sequence[Document], **kwargs: Any) -> Sequence[Document]:205 """Transform sequence of documents by splitting them."""206 return self.split_documents(list(documents))207 208 async def atransform_documents(self, documents: Sequence[Document], **kwargs: Any) -> Sequence[Document]:209 """Asynchronously transform a sequence of documents by splitting them."""210 raise NotImplementedError211 212 213class CharacterTextSplitter(TextSplitter):214 """Splitting text that looks at characters."""215 216 def __init__(self, separator: str = "\n\n", **kwargs: Any) -> None:217 """Create a new TextSplitter."""218 super().__init__(**kwargs)219 self._separator = separator220 221 def split_text(self, text: str) -> list[str]:222 """Split incoming text and return chunks."""223 # First we naively split the large input into a bunch of smaller ones.224 splits = _split_text_with_regex(text, self._separator, self._keep_separator)225 _separator = "" if self._keep_separator else self._separator226 _good_splits_lengths = [] # cache the lengths of the splits227 for split in splits:228 _good_splits_lengths.append(self._length_function(split))229 return self._merge_splits(splits, _separator, _good_splits_lengths)230 231 232class LineType(TypedDict):233 """Line type as typed dict."""234 235 metadata: dict[str, str]236 content: str237 238 239class HeaderType(TypedDict):240 """Header type as typed dict."""241 242 level: int243 name: str244 data: str245 246 247class MarkdownHeaderTextSplitter:248 """Splitting markdown files based on specified headers."""249 250 def __init__(self, headers_to_split_on: list[tuple[str, str]], return_each_line: bool = False):251 """Create a new MarkdownHeaderTextSplitter.252 253 Args:254 headers_to_split_on: Headers we want to track255 return_each_line: Return each line w/ associated headers256 """257 # Output line-by-line or aggregated into chunks w/ common headers258 self.return_each_line = return_each_line259 # Given the headers we want to split on,260 # (e.g., "#, ##, etc") order by length261 self.headers_to_split_on = sorted(headers_to_split_on, key=lambda split: len(split[0]), reverse=True)262 263 def aggregate_lines_to_chunks(self, lines: list[LineType]) -> list[Document]:264 """Combine lines with common metadata into chunks265 Args:266 lines: Line of text / associated header metadata267 """268 aggregated_chunks: list[LineType] = []269 270 for line in lines:271 if aggregated_chunks and aggregated_chunks[-1]["metadata"] == line["metadata"]:272 # If the last line in the aggregated list273 # has the same metadata as the current line,274 # append the current content to the last lines's content275 aggregated_chunks[-1]["content"] += " \n" + line["content"]276 else:277 # Otherwise, append the current line to the aggregated list278 aggregated_chunks.append(line)279 280 return [Document(page_content=chunk["content"], metadata=chunk["metadata"]) for chunk in aggregated_chunks]281 282 def split_text(self, text: str) -> list[Document]:283 """Split markdown file284 Args:285 text: Markdown file"""286 287 # Split the input text by newline character ("\n").288 lines = text.split("\n")289 # Final output290 lines_with_metadata: list[LineType] = []291 # Content and metadata of the chunk currently being processed292 current_content: list[str] = []293 current_metadata: dict[str, str] = {}294 # Keep track of the nested header structure295 # header_stack: List[Dict[str, Union[int, str]]] = []296 header_stack: list[HeaderType] = []297 initial_metadata: dict[str, str] = {}298 299 for line in lines:300 stripped_line = line.strip()301 # Check each line against each of the header types (e.g., #, ##)302 for sep, name in self.headers_to_split_on:303 # Check if line starts with a header that we intend to split on304 if stripped_line.startswith(sep) and (305 # Header with no text OR header is followed by space306 # Both are valid conditions that sep is being used a header307 len(stripped_line) == len(sep) or stripped_line[len(sep)] == " "308 ):309 # Ensure we are tracking the header as metadata310 if name is not None:311 # Get the current header level312 current_header_level = sep.count("#")313 314 # Pop out headers of lower or same level from the stack315 while header_stack and header_stack[-1]["level"] >= current_header_level:316 # We have encountered a new header317 # at the same or higher level318 popped_header = header_stack.pop()319 # Clear the metadata for the320 # popped header in initial_metadata321 if popped_header["name"] in initial_metadata:322 initial_metadata.pop(popped_header["name"])323 324 # Push the current header to the stack325 header: HeaderType = {326 "level": current_header_level,327 "name": name,328 "data": stripped_line[len(sep) :].strip(),329 }330 header_stack.append(header)331 # Update initial_metadata with the current header332 initial_metadata[name] = header["data"]333 334 # Add the previous line to the lines_with_metadata335 # only if current_content is not empty336 if current_content:337 lines_with_metadata.append(338 {339 "content": "\n".join(current_content),340 "metadata": current_metadata.copy(),341 }342 )343 current_content.clear()344 345 break346 else:347 if stripped_line:348 current_content.append(stripped_line)349 elif current_content:350 lines_with_metadata.append(351 {352 "content": "\n".join(current_content),353 "metadata": current_metadata.copy(),354 }355 )356 current_content.clear()357 358 current_metadata = initial_metadata.copy()359 360 if current_content:361 lines_with_metadata.append({"content": "\n".join(current_content), "metadata": current_metadata})362 363 # lines_with_metadata has each line with associated header metadata364 # aggregate these into chunks based on common metadata365 if not self.return_each_line:366 return self.aggregate_lines_to_chunks(lines_with_metadata)367 else:368 return [369 Document(page_content=chunk["content"], metadata=chunk["metadata"]) for chunk in lines_with_metadata370 ]371 372 373# should be in newer Python versions (3.10+)374# @dataclass(frozen=True, kw_only=True, slots=True)375@dataclass(frozen=True)376class Tokenizer:377 chunk_overlap: int378 tokens_per_chunk: int379 decode: Callable[[list[int]], str]380 encode: Callable[[str], list[int]]381 382 383def split_text_on_tokens(*, text: str, tokenizer: Tokenizer) -> list[str]:384 """Split incoming text and return chunks using tokenizer."""385 splits: list[str] = []386 input_ids = tokenizer.encode(text)387 start_idx = 0388 cur_idx = min(start_idx + tokenizer.tokens_per_chunk, len(input_ids))389 chunk_ids = input_ids[start_idx:cur_idx]390 while start_idx < len(input_ids):391 splits.append(tokenizer.decode(chunk_ids))392 start_idx += tokenizer.tokens_per_chunk - tokenizer.chunk_overlap393 cur_idx = min(start_idx + tokenizer.tokens_per_chunk, len(input_ids))394 chunk_ids = input_ids[start_idx:cur_idx]395 return splits396 397 398class TokenTextSplitter(TextSplitter):399 """Splitting text to tokens using model tokenizer."""400 401 def __init__(402 self,403 encoding_name: str = "gpt2",404 model_name: Optional[str] = None,405 allowed_special: Union[Literal["all"], Set[str]] = set(),406 disallowed_special: Union[Literal["all"], Collection[str]] = "all",407 **kwargs: Any,408 ) -> None:409 """Create a new TextSplitter."""410 super().__init__(**kwargs)411 try:412 import tiktoken413 except ImportError:414 raise ImportError(415 "Could not import tiktoken python package. "416 "This is needed in order to for TokenTextSplitter. "417 "Please install it with `pip install tiktoken`."418 )419 420 if model_name is not None:421 enc = tiktoken.encoding_for_model(model_name)422 else:423 enc = tiktoken.get_encoding(encoding_name)424 self._tokenizer = enc425 self._allowed_special = allowed_special426 self._disallowed_special = disallowed_special427 428 def split_text(self, text: str) -> list[str]:429 def _encode(_text: str) -> list[int]:430 return self._tokenizer.encode(431 _text,432 allowed_special=self._allowed_special,433 disallowed_special=self._disallowed_special,434 )435 436 tokenizer = Tokenizer(437 chunk_overlap=self._chunk_overlap,438 tokens_per_chunk=self._chunk_size,439 decode=self._tokenizer.decode,440 encode=_encode,441 )442 443 return split_text_on_tokens(text=text, tokenizer=tokenizer)444 445 446class RecursiveCharacterTextSplitter(TextSplitter):447 """Splitting text by recursively look at characters.448 449 Recursively tries to split by different characters to find one450 that works.451 """452 453 def __init__(454 self,455 separators: Optional[list[str]] = None,456 keep_separator: bool = True,457 **kwargs: Any,458 ) -> None:459 """Create a new TextSplitter."""460 super().__init__(keep_separator=keep_separator, **kwargs)461 self._separators = separators or ["\n\n", "\n", " ", ""]462 463 def _split_text(self, text: str, separators: list[str]) -> list[str]:464 final_chunks = []465 separator = separators[-1]466 new_separators = []467 468 for i, _s in enumerate(separators):469 if _s == "":470 separator = _s471 break472 if re.search(_s, text):473 separator = _s474 new_separators = separators[i + 1 :]475 break476 477 splits = _split_text_with_regex(text, separator, self._keep_separator)478 _good_splits = []479 _good_splits_lengths = [] # cache the lengths of the splits480 _separator = "" if self._keep_separator else separator481 482 for s in splits:483 s_len = self._length_function(s)484 if s_len < self._chunk_size:485 _good_splits.append(s)486 _good_splits_lengths.append(s_len)487 else:488 if _good_splits:489 merged_text = self._merge_splits(_good_splits, _separator, _good_splits_lengths)490 final_chunks.extend(merged_text)491 _good_splits = []492 _good_splits_lengths = []493 if not new_separators:494 final_chunks.append(s)495 else:496 other_info = self._split_text(s, new_separators)497 final_chunks.extend(other_info)498 499 if _good_splits:500 merged_text = self._merge_splits(_good_splits, _separator, _good_splits_lengths)501 final_chunks.extend(merged_text)502 503 return final_chunks504 505 def split_text(self, text: str) -> list[str]:506 return self._split_text(text, self._separators)507 