Skip to main content

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).

Deterministic output paths and locked inputs

Options files contain only serialized option values (for example, {"seed": 42}), with no metadata envelope. This intentionally replaces the old options format; legacy envelopes are not supported. load_opts() checks names and values against current annotations, but cannot detect historical schema changes. Use a lockfile when exact persisted schemas matter:

lock_path = pfn.save_locked("inputs.h5", "opts.json", x, params, seed=42)
inputs, opts = pfn.load_lock(lock_path)

output_path = pfn.output_path(inputs, opts, base_dir="outputs/simulator-v1")
# Create the output directory before saving; thunk does not create directories.
result = pfn(inputs, opts)
pfn.save_output(output_path, result, inputs=inputs, opts=opts)

output_path() validates both groups and returns <full SHA-256 digest>.h5, optionally under base_dir. It never accesses files or executes the function, and does not require a supported return annotation. The key includes persisted parameter schemas and values, excluding function identity, return annotations, Skip arguments, and file paths. Use directories such as simulator-v1 to separate computations and revisions. This method does not execute the function or reuse cached results; use cached() or thunk.cache() for that. Hashes are recomputed so mutations are reflected.

Parameter groups are hashed in signature order; nested dictionary order remains significant. Array layout and byte order are normalized, while dtype, shape, and values remain significant. Custom serializers must produce deterministic representations. When both groups are passed to save_output(), its HDF5 attributes include digest, digest_version, and the two group digests.

save_locked() returns <digest>.lock.json beside the options file. Its strict, versioned metadata records each group's relative path, digest, and schema fingerprints. Paths may contain ..; relocating the files together preserves the lock. Input HDF5 files also retain their fingerprints. load_lock() verifies schemas and content before returning (inputs, opts); it offers no extras or default-filling policies. Content or key mismatches raise DigestMismatchError, schema mismatches raise SchemaMismatchError, malformed metadata raises StorageFormatError, and missing files raise FileNotFoundError.

All three files are prepared before replacing any destination, and the lockfile is published last. Each replacement is atomic; the group is not a filesystem transaction. Parent directories must already exist, destinations must be distinct, and older lockfiles remain in place. Editing either referenced file requires a new lockfile. Both save() and save_locked() forward all function keyword arguments, including one named lockfile.

Disk caching

Cache a function's output using its persisted arguments:

cached_simulator = thunk.cache(simulator, outdir=".cache/simulator-v1")
result = cached_simulator(x, params, seed=42, chunk_size=4096)
# The same persisted arguments load the saved result without running simulator.
result = cached_simulator(x, params, seed=42, chunk_size=4096)

Or use already flattened or restored groups:

inputs, opts = pfn.flatten(x, params, seed=42)
result = pfn.cached(
    inputs,
    opts,
    outdir=".cache/simulator-v1",
    skipped={"chunk_size": 4096},
)

Both interfaces share entries at <outdir>/<digest>.h5 and require an explicit outdir (a string or path-like object). Relative and absolute paths are supported. Directories are created automatically on a miss.

Each output directory must identify one computation and revision, including bound instance state, partial arguments omitted from the exposed signature, captured values, and external dependencies that affect results. Use a new directory when those change. There is no automatic function, source, or package-version hashing. Persisted arguments must determine results together with that computation and revision; Skip values must only control execution details that do not affect results. Persist random seeds or keys explicitly. Do not mutate inputs during execution, and do not rely on side effects being replayed on cache hits.

A supported return annotation is required. Hits validate the output schema and recorded argument digests before restoring the result. These digests identify inputs; they are not checksums of the output payload. Invalid files, schema or digest mismatches, and permission errors propagate rather than triggering recomputation. Only absent entries are misses.

Pass refresh=True to recompute and atomically replace an entry. Failed computation or saving preserves an existing entry. For thunk.cache, cache controls are fixed when constructing the wrapper: a wrapper created with refresh=True recomputes on every invocation. All arguments passed to the wrapper itself belong to the underlying function, even arguments named outdir or refresh. Refresh only replaces invoked entries; a new output directory starts a separate cache without deleting existing entries.

Caching saves outputs only; input files and lockfiles are optional. Concurrent misses may execute the function more than once, with the last successful atomic replacement winning. There is no eviction, expiration, or concurrency locking.

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.2.0

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for thunk 0.2.0
File Size Uploaded
thunk-0.2.0.tar.gz 27.2 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for thunk 0.2.0
File Interpreter ABI Platform
thunk-0.2.0-py3-none-any.whl Python 3 none any Details

Total release size: 59.8 kB

Release files / thunk-0.2.0.tar.gz

Download URL thunk-0.2.0.tar.gz
Size 27.2 kB
Tags Source
SHA-256 checksum
How to use checksums
73f40729acd7983a0360808ba98882c82207c9f216cbb6d304faf61340b6f2ec
BLAKE2b-256 checksum
How to use checksums
f61eb0bfc6e32f375c71734fcb905ca6af5eda7343c3d1967d9adca2a938235f
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 8, 2026.

Transparency log

Release files / thunk-0.2.0-py3-none-any.whl

Download URL thunk-0.2.0-py3-none-any.whl
Size 32.6 kB
Tags Python 3
SHA-256 checksum
How to use checksums
c7a00311ba22065a052f84054d19df184c5dff5487384fc064c79d1683cb3938
BLAKE2b-256 checksum
How to use checksums
45156d3d46c6680f47ecf92892de8c56149ee7292e3d29c2483d332228d94c25
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 8, 2026.

Transparency log

Release history Release notifications | RSS feed

This release

0.2.0 This release

2 release files

0.1.1

2 release files

0.1.0

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page