feat: migrate text splitters to Component syntax (#2530)

* feat: migrate text splitters to Component syntax

* [autofix.ci] apply automated fixes

---------

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: Gabriel Luiz Freitas Almeida <gabriel@langflow.org>
This commit is contained in:
Nicolò Boschi 2024-07-04 19:29:33 +02:00 • committed by GitHub
commit 86aaab0cec
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 2923 additions and 2390 deletions

View file

@ -0,0 +1,58 @@
from abc import abstractmethod
from typing import Any
from langchain_text_splitters import TextSplitter
from langflow.custom import Component
from langflow.io import Output
from langflow.schema import Data
from langflow.utils.util import build_loader_repr_from_data
class LCTextSplitterComponent(Component):
trace_type = "text_splitter"
outputs = [
Output(display_name="Data", name="data", method="split_data"),
]
def _validate_outputs(self):
required_output_methods = ["text_splitter"]
output_names = [output.name for output in self.outputs]
for method_name in required_output_methods:
if method_name not in output_names:
raise ValueError(f"Output with name '{method_name}' must be defined.")
elif not hasattr(self, method_name):
raise ValueError(f"Method '{method_name}' must be defined.")
def split_data(self) -> list[Data]:
data_input = self.get_data_input()
documents = []
if not isinstance(data_input, list):
data_input: list[Any] = [data_input]
for _input in data_input:
if isinstance(_input, Data):
documents.append(_input.to_lc_document())
else:
documents.append(_input)
splitter = self.build_text_splitter()
docs = splitter.split_documents(documents)
data = self.to_data(docs)
self.repr_value = build_loader_repr_from_data(data)
return data
@abstractmethod
def get_data_input(self) -> Any:
"""
Get the data input.
"""
pass
@abstractmethod
def build_text_splitter(self) -> TextSplitter:
"""
Build the text splitter.
"""
pass

View file

@ -1,24 +1,58 @@
from typing import List from typing import List, Any
from langchain_text_splitters import CharacterTextSplitter from langchain_text_splitters import CharacterTextSplitter, TextSplitter
from langflow.custom import CustomComponent from langflow.base.textsplitters.model import LCTextSplitterComponent
from langflow.inputs import IntInput, DataInput, MessageTextInput
from langflow.schema import Data from langflow.schema import Data
from langflow.utils.util import unescape_string from langflow.utils.util import unescape_string
class CharacterTextSplitterComponent(CustomComponent): class CharacterTextSplitterComponent(LCTextSplitterComponent):
display_name = "CharacterTextSplitter" display_name = "CharacterTextSplitter"
description = "Splitting text that looks at characters." description = "Split text by number of characters."
documentation = "https://docs.langflow.org/components/text-splitters#charactertextsplitter"
name = "CharacterTextSplitter" name = "CharacterTextSplitter"
def build_config(self): inputs = [
return { IntInput(
"inputs": {"display_name": "Input", "input_types": ["Document", "Data"]}, name="chunk_size",
"chunk_overlap": {"display_name": "Chunk Overlap", "default": 200}, display_name="Chunk Size",
"chunk_size": {"display_name": "Chunk Size", "default": 1000}, info="The maximum length of each chunk.",
"separator": {"display_name": "Separator", "default": "\n"}, value=1000,
} ),
IntInput(
name="chunk_overlap",
display_name="Chunk Overlap",
info="The amount of overlap between chunks.",
value=200,
),
DataInput(
name="data_input",
display_name="Input",
info="The texts to split.",
input_types=["Document", "Data"],
),
MessageTextInput(
name="separator",
display_name="Separator",
info='The characters to split on.\nIf left empty defaults to "\\n\\n".',
),
]
def get_data_input(self) -> Any:
return self.data_input
def build_text_splitter(self) -> TextSplitter:
if self.separator:
separator = unescape_string(self.separator)
else:
separator = "\n\n"
return CharacterTextSplitter(
chunk_overlap=self.chunk_overlap,
chunk_size=self.chunk_size,
separator=separator,
)
def build( def build(
self, self,

View file

@ -1,85 +1,47 @@
from typing import List, Optional from typing import Any
from langchain_text_splitters import Language, RecursiveCharacterTextSplitter from langchain_text_splitters import Language, RecursiveCharacterTextSplitter, TextSplitter
from langflow.custom import CustomComponent from langflow.base.textsplitters.model import LCTextSplitterComponent
from langflow.schema import Data from langflow.inputs import IntInput, DataInput, DropdownInput
class LanguageRecursiveTextSplitterComponent(CustomComponent): class LanguageRecursiveTextSplitterComponent(LCTextSplitterComponent):
display_name: str = "Language Recursive Text Splitter" display_name: str = "Language Recursive Text Splitter"
description: str = "Split text into chunks of a specified length based on language." description: str = "Split text into chunks of a specified length based on language."
documentation: str = "https://docs.langflow.org/components/text-splitters#languagerecursivetextsplitter" documentation: str = "https://docs.langflow.org/components/text-splitters#languagerecursivetextsplitter"
name = "LanguageRecursiveTextSplitter" name = "LanguageRecursiveTextSplitter"
def build_config(self): inputs = [
options = [x.value for x in Language] IntInput(
return { name="chunk_size",
"inputs": {"display_name": "Input", "input_types": ["Document", "Data"]}, display_name="Chunk Size",
"separator_type": { info="The maximum length of each chunk.",
"display_name": "Separator Type", value=1000,
"info": "The type of separator to use.", ),
"field_type": "str", IntInput(
"options": options, name="chunk_overlap",
"value": "Python", display_name="Chunk Overlap",
}, info="The amount of overlap between chunks.",
"separators": { value=200,
"display_name": "Separators", ),
"info": "The characters to split on.", DataInput(
"is_list": True, name="data_input",
}, display_name="Input",
"chunk_size": { info="The texts to split.",
"display_name": "Chunk Size", input_types=["Document", "Data"],
"info": "The maximum length of each chunk.", ),
"field_type": "int", DropdownInput(
"value": 1000, name="code_language", display_name="Code Language", options=[x.value for x in Language], value="python"
}, ),
"chunk_overlap": { ]
"display_name": "Chunk Overlap",
"info": "The amount of overlap between chunks.",
"field_type": "int",
"value": 200,
},
"code": {"show": False},
}
def build( def get_data_input(self) -> Any:
self, return self.data_input
inputs: List[Data],
chunk_size: Optional[int] = 1000,
chunk_overlap: Optional[int] = 200,
separator_type: str = "Python",
) -> list[Data]:
"""
Split text into chunks of a specified length.
Args: def build_text_splitter(self) -> TextSplitter:
separators (list[str]): The characters to split on. return RecursiveCharacterTextSplitter.from_language(
chunk_size (int): The maximum length of each chunk. language=Language(self.code_language),
chunk_overlap (int): The amount of overlap between chunks. chunk_size=self.chunk_size,
length_function (function): The function to use to calculate the length of the text. chunk_overlap=self.chunk_overlap,
Returns:
list[str]: The chunks of text.
"""
# Make sure chunk_size and chunk_overlap are ints
if isinstance(chunk_size, str):
chunk_size = int(chunk_size)
if isinstance(chunk_overlap, str):
chunk_overlap = int(chunk_overlap)
splitter = RecursiveCharacterTextSplitter.from_language(
language=Language(separator_type),
chunk_size=chunk_size,
chunk_overlap=chunk_overlap,
) )
documents = []
for _input in inputs:
if isinstance(_input, Data):
documents.append(_input.to_lc_document())
else:
documents.append(_input)
docs = splitter.split_documents(documents)
data = self.to_data(docs)
return data

View file

@ -1,15 +1,13 @@
from langchain_text_splitters import RecursiveCharacterTextSplitter from typing import Any
from langchain_text_splitters import RecursiveCharacterTextSplitter, TextSplitter
from langflow.custom import Component from langflow.base.textsplitters.model import LCTextSplitterComponent
from langflow.inputs.inputs import DataInput, IntInput, MessageTextInput from langflow.inputs.inputs import DataInput, IntInput, MessageTextInput
from langflow.schema import Data from langflow.utils.util import unescape_string
from langflow.template.field.base import Output
from langflow.utils.util import build_loader_repr_from_data, unescape_string
class RecursiveCharacterTextSplitterComponent(Component): class RecursiveCharacterTextSplitterComponent(LCTextSplitterComponent):
display_name: str = "Recursive Character Text Splitter" display_name: str = "Recursive Character Text Splitter"
description: str = "Split text into chunks of a specified length." description: str = "Split text trying to keep all related text together."
documentation: str = "https://docs.langflow.org/components/text-splitters#recursivecharactertextsplitter" documentation: str = "https://docs.langflow.org/components/text-splitters#recursivecharactertextsplitter"
name = "RecursiveCharacterTextSplitter" name = "RecursiveCharacterTextSplitter"
@ -39,49 +37,20 @@ class RecursiveCharacterTextSplitterComponent(Component):
is_list=True, is_list=True,
), ),
] ]
outputs = [
Output(display_name="Data", name="data", method="split_data"),
]
def split_data(self) -> list[Data]: def get_data_input(self) -> Any:
""" return self.data_input
Split text into chunks of a specified length.
Args: def build_text_splitter(self) -> TextSplitter:
separators (list[str] | None): The characters to split on. if not self.separators:
chunk_size (int): The maximum length of each chunk. separators: list[str] | None = None
chunk_overlap (int): The amount of overlap between chunks. else:
Returns:
list[str]: The chunks of text.
"""
if self.separators == "":
self.separators: list[str] | None = None
elif self.separators:
# check if the separators list has escaped characters # check if the separators list has escaped characters
# if there are escaped characters, unescape them # if there are escaped characters, unescape them
self.separators = [unescape_string(x) for x in self.separators] separators = [unescape_string(x) for x in self.separators]
# Make sure chunk_size and chunk_overlap are ints return RecursiveCharacterTextSplitter(
if self.chunk_size: separators=separators,
self.chunk_size: int = int(self.chunk_size)
if self.chunk_overlap:
self.chunk_overlap: int = int(self.chunk_overlap)
splitter = RecursiveCharacterTextSplitter(
separators=self.separators,
chunk_size=self.chunk_size, chunk_size=self.chunk_size,
chunk_overlap=self.chunk_overlap, chunk_overlap=self.chunk_overlap,
) )
documents = []
if not isinstance(self.data_input, list):
self.data_input: list[Data] = [self.data_input]
for _input in self.data_input:
if isinstance(_input, Data):
documents.append(_input.to_lc_document())
else:
documents.append(_input)
docs = splitter.split_documents(documents)
data = self.to_data(docs)
self.repr_value = build_loader_repr_from_data(data)
return data

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long