caj 0.0.1__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
caj/__init__.py ADDED
@@ -0,0 +1,116 @@
1
+ import functools
2
+ import hashlib
3
+ import os
4
+ import sys
5
+ from collections.abc import Callable
6
+ from pathlib import Path
7
+ from typing import overload
8
+ from warnings import warn
9
+
10
+ if sys.version_info < (3, 15):
11
+ from typing_extensions import sentinel
12
+
13
+ import jax
14
+ from platformdirs import user_cache_path
15
+
16
+ from ._cache import Cache
17
+ from ._serialization import deserialize_pytree, serialize_jaxpr, serialize_pytree
18
+ from ._writers import HashWriter, WriteLimitError
19
+
20
+ DEFAULT_CACHE_DIR = user_cache_path(__package__)
21
+ DEFAULT_CACHE_MAX_BYTES = 1_000_000_000
22
+
23
+ _MISSING = sentinel("_MISSING")
24
+
25
+
26
+ @overload
27
+ def cache[**P, R](func: Callable[P, R], /) -> Callable[P, R]: ...
28
+
29
+
30
+ @overload
31
+ def cache[**P, R](
32
+ func: Callable[P, R],
33
+ /,
34
+ *,
35
+ dir: str | os.PathLike[str],
36
+ max_bytes: int | None = ...,
37
+ ) -> Callable[P, R]: ...
38
+
39
+
40
+ @overload
41
+ def cache[**P, R](
42
+ *, dir: str | os.PathLike[str], max_bytes: int | None = ...
43
+ ) -> Callable[[Callable[P, R]], Callable[P, R]]: ...
44
+
45
+
46
+ def cache[**P, R](
47
+ func: Callable[P, R] | _MISSING = _MISSING,
48
+ /,
49
+ *,
50
+ dir: str | os.PathLike[str] | _MISSING = _MISSING,
51
+ max_bytes: int | _MISSING = _MISSING,
52
+ ) -> Callable[P, R] | Callable[[Callable[P, R]], Callable[P, R]]:
53
+ if dir is not _MISSING:
54
+ dir = Path(dir).absolute()
55
+
56
+ if max_bytes is _MISSING:
57
+ max_bytes = DEFAULT_CACHE_MAX_BYTES
58
+ elif max_bytes is not None and max_bytes <= 0:
59
+ raise ValueError(f"max_bytes must be None or positive, got {max_bytes!r}")
60
+
61
+ elif max_bytes is not _MISSING:
62
+ raise TypeError("max_bytes is only allowed if dir is also specified")
63
+
64
+ else:
65
+ dir = DEFAULT_CACHE_DIR
66
+ max_bytes = DEFAULT_CACHE_MAX_BYTES
67
+
68
+ cache = Cache(dir, max_bytes=max_bytes)
69
+
70
+ def decorator(func: Callable[P, R], /) -> Callable[P, R]:
71
+ jaxpr_func = jax.make_jaxpr(func, return_shape=True)
72
+
73
+ @functools.wraps(func)
74
+ def wrapper(*args: P.args, **kwargs: P.kwargs) -> R:
75
+ jaxpr, ret_shape = jaxpr_func(*args, **kwargs)
76
+
77
+ h = hashlib.blake2b(digest_size=16)
78
+
79
+ f = HashWriter(h)
80
+ serialize_jaxpr(f, jaxpr)
81
+ serialize_pytree(f, args)
82
+ serialize_pytree(f, kwargs)
83
+
84
+ key = h.hexdigest()
85
+
86
+ try:
87
+ with cache.read(key) as f:
88
+ ret = deserialize_pytree(f, like=ret_shape)
89
+ except KeyError:
90
+ pass
91
+ except (OSError, RuntimeError) as e:
92
+ warn(
93
+ f"{__package__}: failed to load cached entry: {e}",
94
+ RuntimeWarning,
95
+ stacklevel=2,
96
+ )
97
+ else:
98
+ return ret
99
+
100
+ ret = func(*args, **kwargs)
101
+
102
+ try:
103
+ with cache.write(key) as f:
104
+ serialize_pytree(f, ret, exc_workaround=WriteLimitError)
105
+ except (OSError, WriteLimitError) as e:
106
+ warn(
107
+ f"{__package__}: failed to save cache to {cache.dir}: {e}",
108
+ RuntimeWarning,
109
+ stacklevel=2,
110
+ )
111
+
112
+ return ret
113
+
114
+ return wrapper
115
+
116
+ return decorator(func) if func is not _MISSING else decorator
caj/_cache.py ADDED
@@ -0,0 +1,113 @@
1
+ import contextlib
2
+ import os
3
+ import sys
4
+ import time
5
+ from collections.abc import Generator
6
+ from pathlib import Path
7
+ from warnings import warn
8
+
9
+ if sys.version_info >= (3, 14):
10
+ from io import Writer
11
+ else:
12
+ from typing_extensions import Writer
13
+
14
+ from safewrite import atomic_write
15
+
16
+ from ._typing import SupportsReadSeek
17
+ from ._writers import BoundedWriter
18
+
19
+
20
+ class Cache:
21
+ _SUFFIX = f".{__package__}"
22
+
23
+ def __init__(self, dir: Path, /, *, max_bytes: int | None = None) -> None:
24
+ self.dir = dir
25
+ self.max_bytes = max_bytes
26
+
27
+ @contextlib.contextmanager
28
+ def read(self, key: str, /) -> Generator[SupportsReadSeek[bytes], None, None]:
29
+ path = self._path(key)
30
+ try:
31
+ f = path.open("rb")
32
+ except FileNotFoundError as e:
33
+ raise KeyError(key) from e
34
+ try:
35
+ yield f
36
+ finally:
37
+ f.close()
38
+ self._hit(key)
39
+
40
+ @contextlib.contextmanager
41
+ def write(self, key: str, /) -> Generator[Writer[bytes], None, None]:
42
+ path = self._path(key)
43
+ path.parent.mkdir(parents=True, exist_ok=True)
44
+ with atomic_write(path, mode="wb") as f:
45
+ if self.max_bytes is not None:
46
+ f = BoundedWriter(f, max_bytes=self.max_bytes)
47
+ yield f
48
+ self._cull()
49
+
50
+ def _path(self, key: str, /) -> Path:
51
+ return self.dir / f"{key}{self._SUFFIX}"
52
+
53
+ def _hit(self, key: str, /) -> None:
54
+ path = self._path(key)
55
+ with contextlib.suppress(OSError):
56
+ os.utime(path, ns=(time.time_ns(), path.stat().st_mtime_ns))
57
+
58
+ def _cull(self) -> None:
59
+ if self.max_bytes is None:
60
+ return
61
+
62
+ cache_entries: list[tuple[Path, int, int]] = []
63
+
64
+ try:
65
+ with os.scandir(self.dir) as dir_entries:
66
+ for dir_entry in dir_entries:
67
+ path = Path(dir_entry.path)
68
+
69
+ if path.suffix != self._SUFFIX:
70
+ continue
71
+
72
+ try:
73
+ if not dir_entry.is_file(follow_symlinks=False):
74
+ continue
75
+ st = dir_entry.stat(follow_symlinks=False)
76
+ except OSError:
77
+ continue
78
+
79
+ cache_entries.append((path, st.st_atime_ns, st.st_size))
80
+ except OSError:
81
+ warn(
82
+ f"{__package__}: failed to cull cache in {self.dir} to {self.max_bytes} bytes; "
83
+ "unable to list files",
84
+ RuntimeWarning,
85
+ stacklevel=3,
86
+ )
87
+ return
88
+
89
+ total_size = sum(entry[2] for entry in cache_entries)
90
+
91
+ if total_size <= self.max_bytes:
92
+ return
93
+
94
+ cache_entries.sort(key=lambda entry: entry[1])
95
+
96
+ for cache_entry in cache_entries:
97
+ assert cache_entry[0].suffix == self._SUFFIX
98
+ try:
99
+ cache_entry[0].unlink(missing_ok=True)
100
+ except OSError:
101
+ continue
102
+
103
+ total_size -= cache_entry[2]
104
+
105
+ if total_size <= self.max_bytes:
106
+ return
107
+
108
+ warn(
109
+ f"{__package__}: failed to cull cache in {self.dir} to {self.max_bytes} bytes; "
110
+ f"current size is at least {total_size} bytes",
111
+ RuntimeWarning,
112
+ stacklevel=3,
113
+ )
caj/_serialization.py ADDED
@@ -0,0 +1,71 @@
1
+ import sys
2
+ from typing import Any
3
+
4
+ if sys.version_info >= (3, 14):
5
+ from io import Writer
6
+ else:
7
+ from typing_extensions import Writer
8
+
9
+ import equinox as eqx
10
+ import jax
11
+ import jax.numpy as jnp
12
+ import jaxlib
13
+
14
+ from ._typing import SupportsReadSeek
15
+
16
+
17
+ def serialize_jaxpr(f: Writer[bytes], jaxpr: Any, /) -> None:
18
+ f.write(jax.__version__.encode())
19
+ f.write(jaxlib.__version__.encode())
20
+ f.write(str(jaxpr).encode())
21
+ serialize_pytree(f, jaxpr.consts)
22
+
23
+
24
+ # Workaround for https://github.com/patrick-kidger/equinox/issues/1255
25
+ def serialize_filter_spec(f: Writer[bytes], x: object, /) -> None:
26
+ if isinstance(x, jax.Array) and jnp.issubdtype(x.dtype, jax.dtypes.prng_key):
27
+ x = jax.random.key_data(x)
28
+ return eqx.default_serialise_filter_spec(f, x)
29
+
30
+
31
+ def serialize_pytree(
32
+ f: Writer[bytes],
33
+ pytree: object,
34
+ /,
35
+ *,
36
+ exc_workaround: type[Exception] | None = None,
37
+ ) -> None:
38
+ if exc_workaround is not None:
39
+ # Workaround for https://github.com/patrick-kidger/equinox/issues/1255
40
+ assert issubclass(exc_workaround, Exception)
41
+ assert not issubclass(exc_workaround, RuntimeError)
42
+
43
+ try:
44
+ eqx.tree_serialise_leaves(f, pytree, filter_spec=serialize_filter_spec)
45
+ except RuntimeError as e:
46
+ cause = e.__cause__
47
+ while isinstance(cause, RuntimeError):
48
+ cause = cause.__cause__
49
+ if isinstance(cause, exc_workaround):
50
+ raise cause from e # ty: ignore[invalid-raise]
51
+ raise
52
+ else:
53
+ eqx.tree_serialise_leaves(f, pytree, filter_spec=serialize_filter_spec)
54
+
55
+
56
+ # Workaround for https://github.com/patrick-kidger/equinox/issues/1255
57
+ def deserialize_filter_spec(f: SupportsReadSeek[bytes], x: object, /) -> object:
58
+ ret = eqx.default_deserialise_filter_spec(f, x)
59
+ if isinstance(x, (jax.Array, jax.ShapeDtypeStruct)) and jnp.issubdtype(
60
+ x.dtype, jax.dtypes.prng_key
61
+ ):
62
+ return jax.random.wrap_key_data(
63
+ ret, impl=jax.random.key_impl(jnp.zeros((), x.dtype))
64
+ )
65
+ return ret
66
+
67
+
68
+ def deserialize_pytree[T](f: SupportsReadSeek[bytes], /, *, like: T) -> T:
69
+ return eqx.tree_deserialise_leaves(
70
+ f, like=like, filter_spec=deserialize_filter_spec
71
+ )
caj/_typing.py ADDED
@@ -0,0 +1,16 @@
1
+ import sys
2
+ from collections.abc import Buffer
3
+ from typing import Protocol
4
+
5
+ if sys.version_info >= (3, 14):
6
+ from io import Reader
7
+ else:
8
+ from typing_extensions import Reader
9
+
10
+
11
+ class SupportsReadSeek[T](Reader[T], Protocol):
12
+ def seek(self, offset: int, whence: int, /) -> object: ...
13
+
14
+
15
+ class SupportsUpdate(Protocol):
16
+ def update(self, data: Buffer, /) -> None: ...
caj/_writers.py ADDED
@@ -0,0 +1,47 @@
1
+ import sys
2
+ from collections.abc import Buffer
3
+ from typing import override
4
+
5
+ if sys.version_info >= (3, 14):
6
+ from io import Writer
7
+ else:
8
+ from typing_extensions import Writer
9
+
10
+
11
+ from ._typing import SupportsUpdate
12
+
13
+
14
+ class HashWriter(Writer[Buffer]):
15
+ def __init__(self, hasher: SupportsUpdate, /) -> None:
16
+ self._hasher = hasher
17
+
18
+ @override
19
+ def write(self, data: Buffer, /) -> int:
20
+ data = memoryview(data)
21
+ self._hasher.update(data)
22
+ return data.nbytes
23
+
24
+
25
+ class WriteLimitError(Exception):
26
+ pass
27
+
28
+
29
+ class BoundedWriter(Writer[Buffer]):
30
+ def __init__(self, writer: Writer[Buffer], /, *, max_bytes: int) -> None:
31
+ self._writer = writer
32
+ self._max_bytes = max_bytes
33
+ self._bytes_written = 0
34
+
35
+ @override
36
+ def write(self, data: Buffer, /) -> int:
37
+ data = memoryview(data)
38
+ if self._bytes_written + data.nbytes > self._max_bytes:
39
+ raise WriteLimitError(
40
+ f"write would exceed limit of {self._max_bytes} bytes"
41
+ )
42
+
43
+ ret = self._writer.write(data)
44
+ assert 0 <= ret <= data.nbytes
45
+
46
+ self._bytes_written += ret
47
+ return ret
caj/py.typed ADDED
File without changes
@@ -0,0 +1,111 @@
1
+ Metadata-Version: 2.4
2
+ Name: caj
3
+ Version: 0.0.1
4
+ Summary: Automatic persistent caching for JAX-based computations
5
+ Author: Gabriel S. Gerlero
6
+ Author-email: Gabriel S. Gerlero <ggerlero@cimec.unl.edu.ar>
7
+ License-Expression: Apache-2.0
8
+ Classifier: Development Status :: 4 - Beta
9
+ Classifier: Intended Audience :: Developers
10
+ Classifier: Intended Audience :: Science/Research
11
+ Classifier: Operating System :: OS Independent
12
+ Classifier: Programming Language :: Python
13
+ Classifier: Programming Language :: Python :: 3
14
+ Classifier: Programming Language :: Python :: 3.12
15
+ Classifier: Programming Language :: Python :: 3.13
16
+ Classifier: Programming Language :: Python :: 3.14
17
+ Classifier: Programming Language :: Python :: 3.15
18
+ Classifier: Topic :: Scientific/Engineering
19
+ Classifier: Topic :: Software Development :: Libraries
20
+ Classifier: Typing :: Typed
21
+ Requires-Dist: equinox>=0.13,<0.14
22
+ Requires-Dist: jax>=0.5,<0.12
23
+ Requires-Dist: platformdirs>=2.1,<5
24
+ Requires-Dist: safewrite>=0.1.3,<0.2
25
+ Requires-Dist: typing-extensions>=4.16,<5
26
+ Requires-Python: >=3.12
27
+ Description-Content-Type: text/markdown
28
+
29
+ <div align="center">
30
+ <a href="https://github.com/gerlero/caj"><img src="https://raw.githubusercontent.com/gerlero/caj/main/logo.png" alt="caj" width="200"/></a>
31
+
32
+ **Automatic persistent caching for JAX function invocations**
33
+
34
+ [![CI](https://github.com/gerlero/caj/actions/workflows/ci.yml/badge.svg)](https://github.com/gerlero/caj/actions/workflows/ci.yml)
35
+ [![Codecov](https://codecov.io/gh/gerlero/caj/branch/main/graph/badge.svg)](https://codecov.io/gh/gerlero/caj)
36
+ [![Ruff](https://img.shields.io/endpoint?url=https://raw.githubusercontent.com/astral-sh/ruff/main/assets/badge/v2.json)](https://github.com/astral-sh/ruff)
37
+ [![ty](https://img.shields.io/endpoint?url=https://raw.githubusercontent.com/astral-sh/ty/main/assets/badge/v0.json)](https://github.com/astral-sh/ty)
38
+ [![uv](https://img.shields.io/endpoint?url=https://raw.githubusercontent.com/astral-sh/uv/main/assets/badge/v0.json)](https://github.com/astral-sh/uv)
39
+ [![Publish](https://github.com/gerlero/caj/actions/workflows/pypi-publish.yml/badge.svg)](https://github.com/gerlero/caj/actions/workflows/pypi-publish.yml)
40
+ [![PyPI](https://img.shields.io/pypi/v/caj)](https://pypi.org/project/caj/)
41
+ [![PyPI - Python Version](https://img.shields.io/pypi/pyversions/caj)](https://pypi.org/project/caj/)
42
+ </div>
43
+
44
+ **caj** is a simple persistent cache for JAX-based code.
45
+
46
+ Decorate a function with `@cache`, and **caj** will store the return on disk. Later calls with the same inputs will load the result directly from the cache instead of performing the computation again.
47
+
48
+ ```python
49
+ import jax.numpy as jnp
50
+ from caj import cache
51
+
52
+
53
+ @cache
54
+ @jax.jit
55
+ def compute(x):
56
+ return jnp.linalg.eigvalsh(x)
57
+
58
+
59
+ x = jnp.eye(1000)
60
+
61
+ y = compute(x)
62
+ ```
63
+
64
+ Running the above code twice will only compute the eigenvalues once: the second time, **caj** will just load the result from disk.
65
+
66
+ **caj** takes advantage of the computation tracing provided by JAX so it can detect changes in the decorated function or any other functions it calls and not return cached results that were computed with different code.
67
+
68
+ ## Installation
69
+
70
+ Install **caj** from PyPI:
71
+
72
+ ```console
73
+ pip install caj
74
+ ```
75
+
76
+ ## Usage
77
+
78
+ For most uses, the default cache is enough:
79
+
80
+ ```python
81
+ from caj import cache
82
+
83
+
84
+ @cache
85
+ def f(x): ...
86
+ ```
87
+
88
+ By default, **caj** stores entries in a per-user cache directory and limits the cache to **1 GB**.
89
+
90
+ A custom cache directory can be specified with `dir`:
91
+
92
+ ```python
93
+ @cache(dir=".cache")
94
+ def f(x): ...
95
+ ```
96
+
97
+ The maximum size of a custom cache can also be configured via `max_bytes`:
98
+
99
+ ```python
100
+ @cache(dir=".cache", max_bytes=2_000_000_000)
101
+ def f(x): ...
102
+ ```
103
+
104
+ Or, set `max_bytes=None` for no size limit:
105
+
106
+ ```python
107
+ @cache(dir=".cache", max_bytes=None)
108
+ def f(x): ...
109
+ ```
110
+
111
+ When a size limit is enabled, **caj** will remove the least recently used entries as necessary to keep the cache within the target size.
@@ -0,0 +1,9 @@
1
+ caj/__init__.py,sha256=eGwFjfCcbXANcflNUTlXYU1qjpt3ZuotGrp9HDVdCJQ,3241
2
+ caj/_cache.py,sha256=2tW4htp4_qcDaEceZSLWltqt699_OSQrD4pnF_KZRgc,3336
3
+ caj/_serialization.py,sha256=Bg2DuFQts3VZfCLqiDjrgZb9xbRoXi_4rGXI2WNu-_Y,2258
4
+ caj/_typing.py,sha256=Izwkkk6vpLoP1Ay7f6-Tv-Ng5ZceGPPMF9rvm_WepwM,378
5
+ caj/_writers.py,sha256=YuNzNxGwBhMjhjwqte8murAof_pPwk3AOawOl7X09Lk,1166
6
+ caj/py.typed,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
7
+ caj-0.0.1.dist-info/WHEEL,sha256=cmC5s21ojypbVslldL7IJq3hjZH-tINy4rziKePFsG0,81
8
+ caj-0.0.1.dist-info/METADATA,sha256=6xcuacj9HXXYlCOQsA_GzHwS-drPahNoaegwmdaUg1g,3946
9
+ caj-0.0.1.dist-info/RECORD,,
@@ -0,0 +1,4 @@
1
+ Wheel-Version: 1.0
2
+ Generator: uv 0.12.23
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any