mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 09:25:14 +03:00
perf(core): Lazily import transformers (#38037)
This commit is contained in:
1 parent
9f2d56e376
commit
cccfbb1c5b
1 file changed
+24
-9
@@ -39,13 +39,6 @@ from langchain_core.runnables import Runnable, RunnableSerializable
|
||||
if TYPE_CHECKING:
|
||||
from langchain_core.outputs import LLMResult
|
||||
|
||||
try:
|
||||
from transformers import GPT2TokenizerFast # type: ignore[import-not-found]
|
||||
|
||||
_HAS_TRANSFORMERS = True
|
||||
except ImportError:
|
||||
_HAS_TRANSFORMERS = False
|
||||
|
||||
|
||||
class LangSmithParams(TypedDict, total=False):
|
||||
"""LangSmith parameters for tracing."""
|
||||
@@ -74,6 +67,27 @@ class LangSmithParams(TypedDict, total=False):
|
||||
"""Integration that created the trace."""
|
||||
|
||||
|
||||
@cache
|
||||
def _get_tokenizer_module() -> Callable[[str], Any] | None:
|
||||
"""Get the `from_pretrained` function from `GPT2TokenizerFast`.
|
||||
|
||||
Imported lazily so that merely importing this module doesn't pull in
|
||||
`transformers` (and, transitively, `pytorch`). Cached so the cost is paid once.
|
||||
|
||||
Returns:
|
||||
The `GPT2TokenizerFast.from_pretrained` function, or `None` if
|
||||
`transformers` is not installed.
|
||||
|
||||
"""
|
||||
try:
|
||||
from transformers import ( # type: ignore[import-not-found] # noqa: PLC0415
|
||||
GPT2TokenizerFast,
|
||||
)
|
||||
except ImportError:
|
||||
return None
|
||||
return cast("Callable[[str], Any]", GPT2TokenizerFast.from_pretrained)
|
||||
|
||||
|
||||
@cache # Cache the tokenizer
|
||||
def get_tokenizer() -> Any:
|
||||
"""Get a GPT-2 tokenizer instance.
|
||||
@@ -87,7 +101,8 @@ def get_tokenizer() -> Any:
|
||||
The GPT-2 tokenizer instance.
|
||||
|
||||
"""
|
||||
if not _HAS_TRANSFORMERS:
|
||||
loader_impl = _get_tokenizer_module()
|
||||
if loader_impl is None:
|
||||
msg = (
|
||||
"Could not import transformers python package. "
|
||||
"This is needed in order to calculate get_token_ids. "
|
||||
@@ -95,7 +110,7 @@ def get_tokenizer() -> Any:
|
||||
)
|
||||
raise ImportError(msg)
|
||||
# create a GPT-2 tokenizer instance
|
||||
return GPT2TokenizerFast.from_pretrained("gpt2")
|
||||
return loader_impl("gpt2")
|
||||
|
||||
|
||||
_GPT2_TOKENIZER_WARNED = False
|
||||
|
||||
Reference in new issue
Block a user