thunk
Generate automatic save and load methods for functions that manipulate arrays
Installation
uv add thunk
# or
pip install thunk
Usage
thunk.fn derives save/load methods from a callable's type annotations. It
inspects the signature and never runs the function body.
from dataclasses import dataclass
from typing import Annotated
import numpy as np
import thunk
@dataclass(frozen=True)
class Params:
y: np.ndarray
labels: list[str]
def simulator(
x: np.ndarray,
/,
params: Params,
*,
seed: int = 0,
chunk_size: Annotated[int, thunk.Skip()] = 4096,
) -> np.ndarray: ...
pfn = thunk.fn(simulator)
# Bind arguments and split them into persisted groups (does not run simulator).
inputs, opts = pfn.flatten(x, params, seed=42)
pfn.save_inputs("inputs.h5", inputs)
pfn.save_opts("opts.json", opts)
# ... or flatten and save both in one step:
pfn.save("inputs.h5", "opts.json", x, params, seed=42)
# Restore, optionally edit the configuration, and run.
inputs = pfn.load_inputs("inputs.h5")
opts = pfn.load_opts("opts.json")
result = pfn(inputs, opts)
# Or restore a legal Python call without executing simulator.
args, kwargs = pfn.load(inputs="inputs.h5", opts="opts.json")
result = simulator(*args, **kwargs)
pfn.save_output("output.h5", result, inputs=inputs, opts=opts)
restored = pfn.load_output("output.h5")
Roles decide where each parameter is stored:
| Role | Group | Storage |
|---|---|---|
Data() (default when the annotation contains arrays) |
inputs |
HDF5 |
Static() (default otherwise) |
opts |
JSON, via pydantic |
Skip() |
neither | supplied at call time |
Supported annotations: bool, int, float, str, None, np.ndarray,
jax.Array (with the JAX extra), Literal[...], T | None, list[T], tuple[...], dict[str, T], and
dataclasses of those. Values are validated strictly (no int → float, no
list ↔ tuple). Types outside that set can use
Data(DataSerializer(...), DataValidator(...)) or
Static(PlainSerializer(...), PlainValidator(...)).
Loading takes extras="forbid" | "ignore" for stored names that are no longer
parameters and missing="raise" | "default" for parameters absent from a file
(defaults are filled with a warning).
JAX arrays and random keys
Install the optional integration (JAX 0.11.2 or newer):
uv add 'thunk[jax]'
# or
pip install 'thunk[jax]'
Annotate arrays and typed PRNG keys with jax.Array. Both infer Data,
including inside the supported containers and dataclasses. NumPy and JAX
annotations require their respective array types; jax.typing.ArrayLike and
mixed-array unions are unsupported.
import jax
import jax.numpy as jnp
import thunk
def simulate(x: jax.Array, key: jax.Array, *, scale: float = 0.1) -> jax.Array:
return x + scale * jax.random.normal(key, x.shape, dtype=x.dtype)
pfn = thunk.fn(simulate) # jax.jit(simulate) also works
x = jnp.zeros((100, 3), dtype=jnp.float32)
key = jax.random.key(42, impl="threefry2x32")
pfn.save("inputs.h5", "opts.json", x, key, scale=0.2)
inputs = pfn.load_inputs("inputs.h5")
opts = pfn.load_opts("opts.json")
result = pfn(inputs, opts)
pfn.save_output("output.h5", result, inputs=inputs, opts=opts)
Values, shapes, numeric dtypes, and typed-key implementations are preserved.
Supported numeric dtypes are bool, int8/16/32/64, uint8/16/32/64,
float16/32/64, and complex64/128. Extended dtypes such as bfloat16 and float8
are rejected. Loading int64, uint64, float64, or complex128 requires
JAX_ENABLE_X64=1 or jax.config.update("jax_enable_x64", True); thunk raises
ValueTypeError rather than narrowing the stored dtype.
Typed keys support threefry2x32, threefry4x32, philox2x32, philox4x32,
rbg, and unsafe_rbg, including batched and empty key arrays. Restoration
uses the stored implementation regardless of JAX's current default RNG.
Legacy uint32 keys are stored as ordinary numeric arrays. Custom RNG
implementations are unsupported.
Weakly typed arrays (for example jnp.asarray(1)) are rejected: supply an
explicit dtype such as jnp.asarray(1, dtype=jnp.int32). Tracers, deleted
arrays, and arrays not fully addressable by the current process are also
rejected with a parameter or nested-field path. Persist outside JAX
transformations, before deletion or buffer donation, and gather distributed
data on the saving process first.
flatten() retains array identity and performs no transfers. Saving and
input content hashing synchronize and transfer array data to the host;
loading places arrays on JAX's default device. Device placement and sharding
are not restored. Custom pytrees and distributed checkpointing are outside
this integration's scope. Importing thunk or compiling NumPy annotations
does not import JAX, and existing NumPy files remain compatible.
Contributing
Setup
Install uv and Just, then run:
just setup
uv run prek install
The underlying command is uv sync --all-groups if Just is unavailable.
Development
just fmt # uv run ruff check --fix . && uv run ruff format .
just lint # uv run ruff check .
just typecheck # uv run ty check
just test # uv run pytest
just hooks # uv run prek run --all-files
just check # full validation, tests, and package build
Update from the template
The template runs uv lock, so updates must explicitly trust it:
uvx copier@9.18.2 update --trust
just check
Publishing
Configure pypi and testpypi GitHub environments with trusted publishers.
Push a tag matching the version in pyproject.toml, such as v0.1.0, to
publish to PyPI. Run the Release workflow manually to publish to TestPyPI.
License
This project is licensed under the MIT License. See LICENSE for details.
Metadata
Release files for thunk 0.1.1
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Source distribution (sdist)
| File | Size | Uploaded | |
|---|---|---|---|
| thunk-0.1.1.tar.gz | 22.3 kB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| thunk-0.1.1-py3-none-any.whl | Python 3 | none | any | Details |
Total release size: 49.7 kB
Release files / thunk-0.1.1.tar.gz
| Download URL | thunk-0.1.1.tar.gz |
|---|---|
| Size | 22.3 kB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
d05fd6f0f7763659ddfeb730bc5aba5948f9b6ccebb370e2a26d67226a306033
|
|
BLAKE2b-256 checksum How to use checksums |
4c9848c0f6c6b6e8dc0fd01afb0dba3ebda9233f0a3acb65bcb26acd708ac6b8
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
Yes |
| Uploaded via |
twine/7.0.0 CPython/3.13.14
|
Provenance
Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.
PyPI Publish Attestation
PyPI verified that this artifact, at this checksum, originated from the publisher listed below.
Signed by GitHub Actions, verified by PyPI on Oct 7, 2026.
Transparency logRelease files / thunk-0.1.1-py3-none-any.whl
| Download URL | thunk-0.1.1-py3-none-any.whl |
|---|---|
| Size | 27.4 kB |
| Tags | Python 3 |
|
SHA-256 checksum How to use checksums |
7498f70ba9f00bd661a3f824a2c54028b0ec8cac73ff5ea1e996d1c50df763c9
|
|
BLAKE2b-256 checksum How to use checksums |
3d714c6511c75b18013b1344d5371f6e1e683fe14491050303aa40456fa4d840
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
Yes |
| Uploaded via |
twine/7.0.0 CPython/3.13.14
|
Provenance
Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.
PyPI Publish Attestation
PyPI verified that this artifact, at this checksum, originated from the publisher listed below.
Signed by GitHub Actions, verified by PyPI on Oct 7, 2026.
Transparency log