Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 3 additions & 2 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,6 @@ keywords = [
dependencies = [
"pytensor>=3.2.3,<4.0.0",
"numpy",
"safetensors"
]

[project.urls]
Expand All @@ -48,6 +47,7 @@ Tracker = "https://github.com/pymc-devs/pytensor-ml/issues"

[project.optional-dependencies]
dev = [
"safetensors",
"pre-commit",
"pytest",
"pytest-cov",
Expand Down
24 changes: 17 additions & 7 deletions pytensor_ml/checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,15 +10,24 @@

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
# "format" key there).
_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)
Expand Down Expand Up @@ -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
----------
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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 = {
Expand Down
7 changes: 5 additions & 2 deletions pytensor_ml/pretrained.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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

Expand Down
Loading