Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
text_splitter.py507 linesDownload Raw Back to splitter
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