From 901fc04231f743788cb2417db403082cdf129bd7 Mon Sep 17 00:00:00 2001 From: jessegrabowski Date: Thu, 17 Sep 2026 11:04:17 +0800 Subject: [PATCH] Make safetensors an optional dependency --- README.md | 5 +++-- pyproject.toml | 2 +- pytensor_ml/checkpoint.py | 24 +++++++++++++++++------- pytensor_ml/pretrained.py | 7 +++++-- 4 files changed, 26 insertions(+), 12 deletions(-) diff --git a/README.md b/README.md index c706aab..3011436 100644 --- a/README.md +++ b/README.md @@ -19,8 +19,9 @@ any other PyTensor graph — a PyMC model included — as there is nothing else pip install pytensor-ml ``` -The only hard dependencies are `pytensor`, `numpy`, and `safetensors`. A backend beyond the default (`numba`, -`jax`, `torch`, `mlx`) is installed separately, and only loads when you actually compile against it. +The only hard dependencies are `pytensor` and `numpy`. Saving and loading weights needs `safetensors` +(`pip install safetensors`). A backend beyond the default (`numba`, `jax`, `torch`, `mlx`) is installed +separately, and only loads when you actually compile against it. ## Quickstart diff --git a/pyproject.toml b/pyproject.toml index 1fe3280..b731fc2 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -37,7 +37,6 @@ keywords = [ dependencies = [ "pytensor>=3.2.3,<4.0.0", "numpy", - "safetensors" ] [project.urls] @@ -48,6 +47,7 @@ Tracker = "https://github.com/pymc-devs/pytensor-ml/issues" [project.optional-dependencies] dev = [ + "safetensors", "pre-commit", "pytest", "pytest-cov", diff --git a/pytensor_ml/checkpoint.py b/pytensor_ml/checkpoint.py index 8bd630d..e70f564 100644 --- a/pytensor_ml/checkpoint.py +++ b/pytensor_ml/checkpoint.py @@ -10,8 +10,6 @@ from pytensor.compile.sharedvalue import SharedVariable from pytensor.tensor.random.type import RandomGeneratorType -from safetensors import safe_open -from safetensors.numpy import save_file # Generator state rides in the archive's metadata, which is a flat str-to-str map shared with whatever # wrote the file, so our entries are namespaced away from a foreign writer's (HuggingFace stores a @@ -19,6 +17,17 @@ _RNG_KEY_PREFIX = "rng/" +def import_safetensors(): + try: + import safetensors + import safetensors.numpy + except ModuleNotFoundError as error: + raise ModuleNotFoundError( + "pytensor-ml requires safetensors for checkpointing. Install it with `pip install safetensors`." + ) from error + return safetensors + + def holds_generator(variable: SharedVariable) -> bool: """Report whether a shared variable holds a random generator rather than a tensor.""" return isinstance(variable.type, RandomGeneratorType) @@ -154,9 +163,8 @@ def save_state(shared_variables: Sequence[SharedVariable], path: str | Path) -> optimizer state to capture a complete training checkpoint; both are ordinary shared variables and carry self-describing names (e.g. ``"fc1/weight"``, ``"fc1/weight/adam/first_moment"``). :func:`~pytensor_ml.pytensorf.collect_optimizer_state` finds the half no walk of the graph reaches. - A random - generator has no tensor to store, so its state is written to the archive's metadata under the same - key, which is what lets a stochastic network checkpoint at all. + A random generator has no tensor to store, so its state is written to the archive's metadata under + the same key, which is what lets a stochastic network checkpoint at all. Parameters ---------- @@ -194,7 +202,8 @@ def save_state(shared_variables: Sequence[SharedVariable], path: str | Path) -> generator_states[_RNG_KEY_PREFIX + key] = json.dumps(jsonable_rng_state(state)) else: tensors[key] = _as_saveable_array(variable.get_value(), variable.type.dtype) - save_file(tensors, fspath(path), metadata=generator_states or None) + safetensors = import_safetensors() + safetensors.numpy.save_file(tensors, fspath(path), metadata=generator_states or None) def _as_saveable_array(value: Any, dtype: str) -> np.ndarray: @@ -274,7 +283,8 @@ def load_state( source_by_key[key] = source target_by_key[key] = variable - with safe_open(fspath(path), framework="numpy") as archive: + safetensors = import_safetensors() + with safetensors.safe_open(fspath(path), framework="numpy") as archive: values = {key: archive.get_tensor(key) for key in archive.keys()} metadata = archive.metadata() or {} archived_states = { diff --git a/pytensor_ml/pretrained.py b/pytensor_ml/pretrained.py index ec1fb00..7e62c67 100644 --- a/pytensor_ml/pretrained.py +++ b/pytensor_ml/pretrained.py @@ -11,12 +11,12 @@ from pytensor.compile.sharedvalue import SharedVariable from pytensor.graph.basic import Variable from pytensor.tensor.random.type import RandomGeneratorType -from safetensors import safe_open from pytensor_ml.checkpoint import ( bit_generator_kind, generator_from_state, holds_generator, + import_safetensors, jsonable_rng_state, load_state, save_state, @@ -479,7 +479,10 @@ def from_pretrained( # raises, so drawing them first is work thrown away. with initial_values_from(EmptyInitializer()): data_inputs, outputs, keys = build_from_config(config) - with safe_open(_huggingface_weights(directory, variant), framework="numpy") as weights: + safetensors = import_safetensors() + with safetensors.safe_open( + _huggingface_weights(directory, variant), framework="numpy" + ) as weights: keys.load(weights.get_tensor, weights.keys()) return data_inputs, outputs