mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 09:25:14 +03:00
## Summary - Moves `nltk`, `spacy`, `sentence-transformers`, and `konlpy` imports back inside class constructors/functions so they are only loaded when the respective splitter is actually instantiated - Adds a subprocess-based regression test to verify no heavy packages are imported at `langchain_text_splitters` load time ## Why PR #32325 moved these optional dependency imports to module-level `try/except` blocks (to satisfy ruff's `PLC0415` rule). Since `__init__.py` imports all four splitter modules, this caused `import langchain_text_splitters` to eagerly load all optional heavy packages, resulting in: - A PyTorch NVML warning (`UserWarning: Can't initialize NVML`) on non-GPU machines - A ~650MB memory spike on import (74MB → 736MB), vs ~50MB in 0.3.x The fix restores the lazy import pattern with `# noqa: PLC0415` to suppress the linter rule, which is the correct trade-off when a dependency has high instantiation cost. ## Review notes - The `PLC0415` suppressions are intentional — these are optional heavy dependencies that should never be loaded unless the user explicitly instantiates the splitter class - The regression test uses a subprocess for proper isolation (the test file itself imports `langchain_text_splitters` at the top, so `sys.modules` checks within the same process would not reflect a clean import state) Fixes #35437. > **AI disclaimer:** This PR was developed with assistance from Claude Code (Anthropic AI). --------- Co-authored-by: AshwathB-debug <ashwathbalaji04@gmail.com> Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com> Co-authored-by: Mason Daugherty <github@mdrxy.com>
135 lines
5.1 KiB
Python
135 lines
5.1 KiB
Python
"""Sentence transformers text splitter."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from importlib import import_module
|
|
from typing import Any, cast
|
|
|
|
from typing_extensions import override
|
|
|
|
from langchain_text_splitters.base import TextSplitter, Tokenizer, split_text_on_tokens
|
|
|
|
|
|
class SentenceTransformersTokenTextSplitter(TextSplitter):
|
|
"""Splitting text to tokens using sentence model tokenizer."""
|
|
|
|
def __init__(
|
|
self,
|
|
chunk_overlap: int = 50,
|
|
model_name: str = "sentence-transformers/all-mpnet-base-v2",
|
|
tokens_per_chunk: int | None = None,
|
|
model_kwargs: dict[str, Any] | None = None,
|
|
**kwargs: Any,
|
|
) -> None:
|
|
"""Create a new `TextSplitter`.
|
|
|
|
Args:
|
|
chunk_overlap: The number of tokens to overlap between chunks.
|
|
model_name: The name of the sentence transformer model to use.
|
|
tokens_per_chunk: The number of tokens per chunk.
|
|
|
|
If `None`, uses the maximum tokens allowed by the model.
|
|
model_kwargs: Additional parameters for model initialization.
|
|
Parameters of sentence_transformers.SentenceTransformer can be used.
|
|
|
|
Raises:
|
|
ImportError: If the `sentence_transformers` package is not installed.
|
|
ValueError: If `tokens_per_chunk` exceeds the model's maximum token limit.
|
|
"""
|
|
super().__init__(**kwargs, chunk_overlap=chunk_overlap)
|
|
|
|
try:
|
|
sentence_transformers = cast("Any", import_module("sentence_transformers"))
|
|
sentence_transformer_cls = sentence_transformers.SentenceTransformer
|
|
except ImportError as err:
|
|
msg = (
|
|
"Could not import sentence_transformers python package. "
|
|
"This is needed in order to use SentenceTransformersTokenTextSplitter. "
|
|
"Please install it with `pip install sentence-transformers`."
|
|
)
|
|
raise ImportError(msg) from err
|
|
|
|
self.model_name = model_name
|
|
self._model = sentence_transformer_cls(self.model_name, **(model_kwargs or {}))
|
|
self.tokenizer = self._model.tokenizer
|
|
self._initialize_chunk_configuration(tokens_per_chunk=tokens_per_chunk)
|
|
|
|
def _initialize_chunk_configuration(self, *, tokens_per_chunk: int | None) -> None:
|
|
self.maximum_tokens_per_chunk = self._model.max_seq_length
|
|
|
|
if tokens_per_chunk is None:
|
|
if self.maximum_tokens_per_chunk is None:
|
|
msg = (
|
|
"The model does not have a maximum token limit, "
|
|
"and tokens_per_chunk was not provided. "
|
|
"Please provide a value for tokens_per_chunk."
|
|
)
|
|
raise ValueError(msg)
|
|
self.tokens_per_chunk = self.maximum_tokens_per_chunk
|
|
else:
|
|
self.tokens_per_chunk = tokens_per_chunk
|
|
|
|
if (
|
|
self.maximum_tokens_per_chunk is not None
|
|
and self.tokens_per_chunk > self.maximum_tokens_per_chunk
|
|
):
|
|
msg = (
|
|
f"The token limit of the models '{self.model_name}'"
|
|
f" is: {self.maximum_tokens_per_chunk}."
|
|
f" Argument tokens_per_chunk={self.tokens_per_chunk}"
|
|
f" > maximum token limit."
|
|
)
|
|
raise ValueError(msg)
|
|
|
|
@override
|
|
def split_text(self, text: str) -> list[str]:
|
|
"""Splits the input text into smaller components by splitting text on tokens.
|
|
|
|
This method encodes the input text using a private `_encode` method, then
|
|
strips the start and stop token IDs from the encoded result. It returns the
|
|
processed segments as a list of strings.
|
|
|
|
Args:
|
|
text: The input text to be split.
|
|
|
|
Returns:
|
|
A list of string components derived from the input text after encoding and
|
|
processing.
|
|
"""
|
|
|
|
def encode_strip_start_and_stop_token_ids(text: str) -> list[int]:
|
|
return self._encode(text)[1:-1]
|
|
|
|
tokenizer = Tokenizer(
|
|
chunk_overlap=self._chunk_overlap,
|
|
tokens_per_chunk=self.tokens_per_chunk,
|
|
decode=self.tokenizer.decode,
|
|
encode=encode_strip_start_and_stop_token_ids,
|
|
)
|
|
|
|
return split_text_on_tokens(text=text, tokenizer=tokenizer)
|
|
|
|
def count_tokens(self, *, text: str) -> int:
|
|
"""Counts the number of tokens in the given text.
|
|
|
|
This method encodes the input text using a private `_encode` method and
|
|
calculates the total number of tokens in the encoded result.
|
|
|
|
Args:
|
|
text: The input text for which the token count is calculated.
|
|
|
|
Returns:
|
|
The number of tokens in the encoded text.
|
|
"""
|
|
return len(self._encode(text))
|
|
|
|
_max_length_equal_32_bit_integer: int = 2**32
|
|
|
|
def _encode(self, text: str) -> list[int]:
|
|
token_ids_with_start_and_end_token_ids = self.tokenizer.encode(
|
|
text,
|
|
max_length=self._max_length_equal_32_bit_integer,
|
|
truncation="do_not_truncate",
|
|
)
|
|
return cast("list[int]", token_ids_with_start_and_end_token_ids)
|