haliax 1.4.dev343__tar.gz → 1.4.dev345__tar.gz
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.
- {haliax-1.4.dev343 → haliax-1.4.dev345}/PKG-INFO +1 -1
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/scan.md +6 -0
- haliax-1.4.dev345/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/__init__.py +29 -4
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/scan.py +10 -1
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/jax_utils.py +7 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/random.py +22 -26
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_scan.py +48 -6
- haliax-1.4.dev343/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev343 → haliax-1.4.dev345}/.coveragerc +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/.flake8 +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/.gitignore +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/LICENSE +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/README.md +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/api.md +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/css/material.css +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/faq.md +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/fp8.md +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/index.md +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/indexing.md +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/matmul.md +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/nn.md +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/partitioning.md +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/rearrange.md +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/requirements.txt +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/state-dict.md +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/tutorial.md +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/vmap.md +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/mkdocs.yml +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/pyproject.toml +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/core.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/types.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/util.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/core_test.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_attention.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_axis.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_conv.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_debug.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_dot.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_hof.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_int8.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_nn.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_ops.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_pool.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_random.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_utils.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev345
|
|
4
4
|
Summary: Named Tensors for Legible Deep Learning in JAX
|
|
5
5
|
Project-URL: Homepage, https://github.com/stanford-crfm/haliax
|
|
6
6
|
Project-URL: Bug Tracker, https://github.com/stanford-crfm/haliax/issues/
|
|
@@ -289,6 +289,12 @@ which is double that required by the default policy, but it reduces the amount o
|
|
|
289
289
|
Both `save_carries` and `save_inputs` can either be a boolean or the string "offload". If "offload", then the
|
|
290
290
|
checkpointed values will be offloaded to the host during the forward pass, and reloaded during the backward pass.
|
|
291
291
|
|
|
292
|
+
In addition, you can offload block internals by passing a list of strings to `offload_block_internals`:
|
|
293
|
+
|
|
294
|
+
```
|
|
295
|
+
policy = ScanCheckpointPolicy(save_carries=True, save_block_internals=["y"], offload_block_internals=["z"])
|
|
296
|
+
```
|
|
297
|
+
|
|
292
298
|
|
|
293
299
|
### Summary of String and Boolean Aliases
|
|
294
300
|
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "1.4.dev345"
|
|
@@ -121,13 +121,38 @@ def full_like(a: NamedArray, fill_value: T, dtype: Optional[DTypeLike] = None) -
|
|
|
121
121
|
return NamedArray(jnp.full_like(a.array, fill_value, dtype=dtype), a.axes)
|
|
122
122
|
|
|
123
123
|
|
|
124
|
-
def arange(axis:
|
|
125
|
-
"""
|
|
124
|
+
def arange(axis: AxisSpec, *, start=0, step=1, dtype: Optional[DTypeLike] = None) -> NamedArray:
|
|
125
|
+
"""
|
|
126
|
+
Version of jnp.arange that returns a NamedArray.
|
|
127
|
+
|
|
128
|
+
This version differs from jnp.arange (beyond the obvious NamedArray) in two ways:
|
|
129
|
+
|
|
130
|
+
1) It can work with a start that is a tracer (i.e. a JAX expression), whereas jax arange is not able to handle
|
|
131
|
+
tracers.
|
|
132
|
+
2) Axis can be more than one axis, in which case it's equivalent to arange of the product of sizes, followed by
|
|
133
|
+
reshape.
|
|
134
|
+
|
|
135
|
+
Examples
|
|
136
|
+
|
|
137
|
+
```python
|
|
138
|
+
X, Y = hax.make_axes(X=3, Y=4)
|
|
139
|
+
# Create a NamedArray along a single axis
|
|
140
|
+
arr = hax.arange(X) # equivalent to jnp.arange(0, 3, 1)
|
|
141
|
+
# 2D
|
|
142
|
+
arr = hax.arange((X, Y)) # equivalent to jnp.arange(0, 12, 1).reshape(3, 4)
|
|
143
|
+
```
|
|
144
|
+
|
|
145
|
+
"""
|
|
146
|
+
from haliax.jax_utils import to_jax_shape
|
|
147
|
+
from haliax.util import ensure_tuple
|
|
148
|
+
|
|
126
149
|
# if start is a tracer, we need to be a bit cleverer since arange doesn't support tracers
|
|
127
150
|
# return NamedArray(jnp.arange(start, stop, step, dtype=dtype), (axis,))
|
|
151
|
+
size = axis_size(axis)
|
|
128
152
|
|
|
129
|
-
arr = jax.lax.iota(dtype=dtype or jnp.result_type(start), size=
|
|
130
|
-
|
|
153
|
+
arr = jax.lax.iota(dtype=dtype or jnp.result_type(start), size=size) * step + start
|
|
154
|
+
arr = arr.reshape(to_jax_shape(axis))
|
|
155
|
+
return NamedArray(arr, ensure_tuple(axis))
|
|
131
156
|
|
|
132
157
|
|
|
133
158
|
# TODO: add overrides for arraylike start/stop to linspace, logspace, geomspace
|
|
@@ -106,6 +106,13 @@ class ScanCheckpointPolicy:
|
|
|
106
106
|
|
|
107
107
|
See Also: https://docs.jax.dev/en/latest/gradient-checkpointing.html#custom-policies-for-offload
|
|
108
108
|
"""
|
|
109
|
+
|
|
110
|
+
offload_block_internals: list[str] = dataclasses.field(default_factory=list)
|
|
111
|
+
"""
|
|
112
|
+
List of named block internals to offload to the host. This is useful for reducing memory usage on the device
|
|
113
|
+
while still avoiding rematerialization.
|
|
114
|
+
"""
|
|
115
|
+
|
|
109
116
|
prevent_cse: bool = False
|
|
110
117
|
"""
|
|
111
118
|
Whether to prevent common subexpression elimination in the checkpointed function.
|
|
@@ -168,7 +175,6 @@ class ScanCheckpointPolicy:
|
|
|
168
175
|
if self.disable:
|
|
169
176
|
return callable
|
|
170
177
|
elif self.simple:
|
|
171
|
-
print("simple")
|
|
172
178
|
return eqx.filter_checkpoint(callable, prevent_cse=self.prevent_cse)
|
|
173
179
|
else:
|
|
174
180
|
policy = self._to_jax_policy(carry_name, input_name)
|
|
@@ -202,6 +208,9 @@ class ScanCheckpointPolicy:
|
|
|
202
208
|
if isinstance(self.save_block_internals, Sequence):
|
|
203
209
|
our_names_to_save.extend(self.save_block_internals)
|
|
204
210
|
|
|
211
|
+
if self.offload_block_internals:
|
|
212
|
+
our_names_to_offload.extend(self.offload_block_internals)
|
|
213
|
+
|
|
205
214
|
if not our_names_to_save and not our_names_to_offload and not self.save_block_internals:
|
|
206
215
|
return None
|
|
207
216
|
|
|
@@ -263,3 +263,10 @@ def multilevel_scan(f, carry, xs, outer_size, length, reverse=False, unroll=1):
|
|
|
263
263
|
return x
|
|
264
264
|
|
|
265
265
|
return carry, jax.tree.map(_deshape, scanned)
|
|
266
|
+
|
|
267
|
+
|
|
268
|
+
def to_jax_shape(shape):
|
|
269
|
+
from haliax.core import Axis, ensure_tuple
|
|
270
|
+
|
|
271
|
+
shape = ensure_tuple(shape)
|
|
272
|
+
return tuple(axis.size if isinstance(axis, Axis) else axis for axis in shape)
|
|
@@ -12,7 +12,7 @@ from haliax.core import NamedArray, NamedOrNumeric, broadcast_to
|
|
|
12
12
|
from haliax.util import ensure_tuple
|
|
13
13
|
|
|
14
14
|
from .axis import Axis, AxisSelector, AxisSpec, selects_axis
|
|
15
|
-
from .jax_utils import named_call
|
|
15
|
+
from .jax_utils import named_call, to_jax_shape
|
|
16
16
|
from .partitioning import physical_axis_name, physical_axis_size, pspec_for_axis
|
|
17
17
|
|
|
18
18
|
|
|
@@ -23,7 +23,7 @@ def uniform(
|
|
|
23
23
|
shape = ensure_tuple(shape)
|
|
24
24
|
minval = broadcast_to(minval, shape).array
|
|
25
25
|
maxval = broadcast_to(maxval, shape).array
|
|
26
|
-
jax_shape =
|
|
26
|
+
jax_shape = to_jax_shape(shape)
|
|
27
27
|
jax_array = jrandom.uniform(key=key, shape=jax_shape, dtype=dtype, minval=minval, maxval=maxval)
|
|
28
28
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
29
29
|
|
|
@@ -31,7 +31,7 @@ def uniform(
|
|
|
31
31
|
@named_call
|
|
32
32
|
def normal(key, shape: AxisSpec, dtype=float):
|
|
33
33
|
shape = ensure_tuple(shape)
|
|
34
|
-
jax_shape =
|
|
34
|
+
jax_shape = to_jax_shape(shape)
|
|
35
35
|
jax_array = jrandom.normal(key=key, shape=jax_shape, dtype=dtype)
|
|
36
36
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
37
37
|
|
|
@@ -40,7 +40,7 @@ def normal(key, shape: AxisSpec, dtype=float):
|
|
|
40
40
|
def bernoulli(key, shape: AxisSpec, p: NamedOrNumeric):
|
|
41
41
|
shape = ensure_tuple(shape)
|
|
42
42
|
p = broadcast_to(p, shape).array
|
|
43
|
-
jax_shape =
|
|
43
|
+
jax_shape = to_jax_shape(shape)
|
|
44
44
|
jax_array = jrandom.bernoulli(key=key, p=p, shape=jax_shape)
|
|
45
45
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
46
46
|
|
|
@@ -50,7 +50,7 @@ def randint(key, shape: AxisSpec, minval: NamedOrNumeric, maxval: NamedOrNumeric
|
|
|
50
50
|
shape = ensure_tuple(shape)
|
|
51
51
|
minval = broadcast_to(minval, shape).array
|
|
52
52
|
maxval = broadcast_to(maxval, shape).array
|
|
53
|
-
jax_shape =
|
|
53
|
+
jax_shape = to_jax_shape(shape)
|
|
54
54
|
jax_array = jrandom.randint(key=key, shape=jax_shape, minval=minval, maxval=maxval, dtype=dtype)
|
|
55
55
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
56
56
|
|
|
@@ -59,7 +59,7 @@ def randint(key, shape: AxisSpec, minval: NamedOrNumeric, maxval: NamedOrNumeric
|
|
|
59
59
|
def poisson(key, shape: AxisSpec, lam: NamedOrNumeric, dtype=int):
|
|
60
60
|
shape = ensure_tuple(shape)
|
|
61
61
|
lam = broadcast_to(lam, shape).array
|
|
62
|
-
jax_shape =
|
|
62
|
+
jax_shape = to_jax_shape(shape)
|
|
63
63
|
jax_array = jrandom.poisson(key=key, lam=lam, shape=jax_shape, dtype=dtype)
|
|
64
64
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
65
65
|
|
|
@@ -67,7 +67,7 @@ def poisson(key, shape: AxisSpec, lam: NamedOrNumeric, dtype=int):
|
|
|
67
67
|
@named_call
|
|
68
68
|
def exponential(key, shape: AxisSpec, dtype=float):
|
|
69
69
|
shape = ensure_tuple(shape)
|
|
70
|
-
jax_shape =
|
|
70
|
+
jax_shape = to_jax_shape(shape)
|
|
71
71
|
jax_array = jrandom.exponential(key=key, shape=jax_shape, dtype=dtype)
|
|
72
72
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
73
73
|
|
|
@@ -76,7 +76,7 @@ def exponential(key, shape: AxisSpec, dtype=float):
|
|
|
76
76
|
def gamma(key, shape: AxisSpec, a: NamedOrNumeric, dtype=float):
|
|
77
77
|
shape = ensure_tuple(shape)
|
|
78
78
|
a = broadcast_to(a, shape).array
|
|
79
|
-
jax_shape =
|
|
79
|
+
jax_shape = to_jax_shape(shape)
|
|
80
80
|
jax_array = jrandom.gamma(key=key, a=a, shape=jax_shape, dtype=dtype)
|
|
81
81
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
82
82
|
|
|
@@ -86,7 +86,7 @@ def beta(key, shape: AxisSpec, a: NamedOrNumeric, b: NamedOrNumeric, dtype=float
|
|
|
86
86
|
shape = ensure_tuple(shape)
|
|
87
87
|
a = broadcast_to(a, shape).array
|
|
88
88
|
b = broadcast_to(b, shape).array
|
|
89
|
-
jax_shape =
|
|
89
|
+
jax_shape = to_jax_shape(shape)
|
|
90
90
|
jax_array = jrandom.beta(key=key, a=a, b=b, shape=jax_shape, dtype=dtype)
|
|
91
91
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
92
92
|
|
|
@@ -94,7 +94,7 @@ def beta(key, shape: AxisSpec, a: NamedOrNumeric, b: NamedOrNumeric, dtype=float
|
|
|
94
94
|
@named_call
|
|
95
95
|
def laplace(key, shape: AxisSpec, dtype=float):
|
|
96
96
|
shape = ensure_tuple(shape)
|
|
97
|
-
jax_shape =
|
|
97
|
+
jax_shape = to_jax_shape(shape)
|
|
98
98
|
jax_array = jrandom.laplace(key=key, shape=jax_shape, dtype=dtype)
|
|
99
99
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
100
100
|
|
|
@@ -102,7 +102,7 @@ def laplace(key, shape: AxisSpec, dtype=float):
|
|
|
102
102
|
@named_call
|
|
103
103
|
def cauchy(key, shape: AxisSpec, dtype=float):
|
|
104
104
|
shape = ensure_tuple(shape)
|
|
105
|
-
jax_shape =
|
|
105
|
+
jax_shape = to_jax_shape(shape)
|
|
106
106
|
jax_array = jrandom.cauchy(key=key, shape=jax_shape, dtype=dtype)
|
|
107
107
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
108
108
|
|
|
@@ -110,7 +110,7 @@ def cauchy(key, shape: AxisSpec, dtype=float):
|
|
|
110
110
|
@named_call
|
|
111
111
|
def logistic(key, shape: AxisSpec, dtype=float):
|
|
112
112
|
shape = ensure_tuple(shape)
|
|
113
|
-
jax_shape =
|
|
113
|
+
jax_shape = to_jax_shape(shape)
|
|
114
114
|
jax_array = jrandom.logistic(key=key, shape=jax_shape, dtype=dtype)
|
|
115
115
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
116
116
|
|
|
@@ -120,7 +120,7 @@ def truncated_normal(key, shape: AxisSpec, lower: NamedOrNumeric, upper: NamedOr
|
|
|
120
120
|
shape = ensure_tuple(shape)
|
|
121
121
|
lower = broadcast_to(lower, shape).array
|
|
122
122
|
upper = broadcast_to(upper, shape).array
|
|
123
|
-
jax_shape =
|
|
123
|
+
jax_shape = to_jax_shape(shape)
|
|
124
124
|
jax_array = jrandom.truncated_normal(key=key, lower=lower, upper=upper, shape=jax_shape, dtype=dtype)
|
|
125
125
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
126
126
|
|
|
@@ -201,7 +201,7 @@ def generate_sharded(fn, axis: Optional[AxisSelector] = None):
|
|
|
201
201
|
@named_call
|
|
202
202
|
def ball(key, shape: AxisSpec, D: Axis, p: float = 2.0, dtype=float):
|
|
203
203
|
shape = ensure_tuple(shape)
|
|
204
|
-
jax_shape =
|
|
204
|
+
jax_shape = to_jax_shape(shape)
|
|
205
205
|
jax_array = jrandom.ball(key=key, shape=jax_shape, d=D.size, p=p, dtype=dtype)
|
|
206
206
|
return haliax.auto_sharded(NamedArray(jax_array, shape + (D,)))
|
|
207
207
|
|
|
@@ -225,7 +225,7 @@ def choice(
|
|
|
225
225
|
if p is not None:
|
|
226
226
|
assert p.resolve_axis(ensure_tuple(axis)) == p.axes, f"p must be 1D with axis {axis} or be None"
|
|
227
227
|
|
|
228
|
-
jax_shape =
|
|
228
|
+
jax_shape = to_jax_shape(shape)
|
|
229
229
|
jax_p = p.array if p is not None else None
|
|
230
230
|
|
|
231
231
|
jax_array = jrandom.choice(key, a.array, jax_shape, replace=replace, p=jax_p, axis=index)
|
|
@@ -264,7 +264,7 @@ def categorical(key, logits: NamedArray, axis: AxisSelector, shape: Optional[Axi
|
|
|
264
264
|
index = logits._lookup_indices(axis)
|
|
265
265
|
assert index is not None, f"axis {axis} not in logits"
|
|
266
266
|
|
|
267
|
-
jax_shape =
|
|
267
|
+
jax_shape = to_jax_shape(shape)
|
|
268
268
|
|
|
269
269
|
jax_array = jrandom.categorical(key, logits.array, axis=index, shape=jax_shape)
|
|
270
270
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
@@ -273,7 +273,7 @@ def categorical(key, logits: NamedArray, axis: AxisSelector, shape: Optional[Axi
|
|
|
273
273
|
@named_call
|
|
274
274
|
def gumbel(key, shape: AxisSpec, dtype=float):
|
|
275
275
|
shape = ensure_tuple(shape)
|
|
276
|
-
jax_shape =
|
|
276
|
+
jax_shape = to_jax_shape(shape)
|
|
277
277
|
jax_array = jrandom.gumbel(key, jax_shape, dtype=dtype)
|
|
278
278
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
279
279
|
|
|
@@ -288,7 +288,7 @@ def permutation(key, x: NamedArray, axis: AxisSelector, independent: bool = Fals
|
|
|
288
288
|
@named_call
|
|
289
289
|
def rademacher(key, shape: AxisSpec, dtype=float):
|
|
290
290
|
shape = ensure_tuple(shape)
|
|
291
|
-
jax_shape =
|
|
291
|
+
jax_shape = to_jax_shape(shape)
|
|
292
292
|
jax_array = jrandom.rademacher(key, jax_shape, dtype=dtype)
|
|
293
293
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
294
294
|
|
|
@@ -297,7 +297,7 @@ def rademacher(key, shape: AxisSpec, dtype=float):
|
|
|
297
297
|
def t(key, shape: AxisSpec, df: NamedOrNumeric, dtype=float):
|
|
298
298
|
shape = ensure_tuple(shape)
|
|
299
299
|
df = broadcast_to(df, shape)
|
|
300
|
-
jax_shape =
|
|
300
|
+
jax_shape = to_jax_shape(shape)
|
|
301
301
|
jax_array = jrandom.t(key, df.array, jax_shape, dtype=dtype)
|
|
302
302
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
303
303
|
|
|
@@ -307,7 +307,7 @@ def weibull_min(key, shape: AxisSpec, scale: NamedOrNumeric, concentration: Name
|
|
|
307
307
|
shape = ensure_tuple(shape)
|
|
308
308
|
scale = broadcast_to(scale, shape)
|
|
309
309
|
concentration = broadcast_to(concentration, shape)
|
|
310
|
-
jax_shape =
|
|
310
|
+
jax_shape = to_jax_shape(shape)
|
|
311
311
|
jax_array = jrandom.weibull_min(key, scale.array, concentration.array, jax_shape, dtype=dtype)
|
|
312
312
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
313
313
|
|
|
@@ -316,7 +316,7 @@ def weibull_min(key, shape: AxisSpec, scale: NamedOrNumeric, concentration: Name
|
|
|
316
316
|
def pareto(key, shape: AxisSpec, b: NamedOrNumeric, dtype=float):
|
|
317
317
|
shape = ensure_tuple(shape)
|
|
318
318
|
b = broadcast_to(b, shape)
|
|
319
|
-
jax_shape =
|
|
319
|
+
jax_shape = to_jax_shape(shape)
|
|
320
320
|
jax_array = jrandom.pareto(key, b.array, jax_shape, dtype=dtype)
|
|
321
321
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
322
322
|
|
|
@@ -325,15 +325,11 @@ def pareto(key, shape: AxisSpec, b: NamedOrNumeric, dtype=float):
|
|
|
325
325
|
def loggamma(key, shape: AxisSpec, a: NamedOrNumeric, dtype=float):
|
|
326
326
|
shape = ensure_tuple(shape)
|
|
327
327
|
a = broadcast_to(a, shape)
|
|
328
|
-
jax_shape =
|
|
328
|
+
jax_shape = to_jax_shape(shape)
|
|
329
329
|
jax_array = jrandom.loggamma(key, a.array, jax_shape, dtype=dtype)
|
|
330
330
|
return haliax.auto_sharded(NamedArray(jax_array, shape))
|
|
331
331
|
|
|
332
332
|
|
|
333
|
-
def _to_jax_shape(shape):
|
|
334
|
-
return tuple(axis.size if isinstance(axis, Axis) else axis for axis in shape)
|
|
335
|
-
|
|
336
|
-
|
|
337
333
|
__all__ = [
|
|
338
334
|
"generate_sharded",
|
|
339
335
|
"uniform",
|
|
@@ -171,40 +171,56 @@ E = hax.Axis("E", 10)
|
|
|
171
171
|
|
|
172
172
|
|
|
173
173
|
@pytest.mark.parametrize(
|
|
174
|
-
"name,policy,expected_scan_shapes",
|
|
174
|
+
"name,policy,expected_scan_shapes,check_offloading",
|
|
175
175
|
[
|
|
176
|
-
(
|
|
177
|
-
|
|
176
|
+
(
|
|
177
|
+
"disabled",
|
|
178
|
+
ScanCheckpointPolicy(disable=True),
|
|
179
|
+
[(E.size,), (Block.size, E.size), (Block.size, E.size)],
|
|
180
|
+
None,
|
|
181
|
+
),
|
|
182
|
+
("carry_true", True, [(E.size,), (Block.size, E.size)], None),
|
|
178
183
|
(
|
|
179
184
|
"carry",
|
|
180
185
|
ScanCheckpointPolicy(save_carries=True, save_block_internals=False),
|
|
181
186
|
[(E.size,), (Block.size, E.size)],
|
|
187
|
+
None,
|
|
182
188
|
),
|
|
183
189
|
(
|
|
184
190
|
"everything",
|
|
185
191
|
ScanCheckpointPolicy(save_carries=True, save_inputs=True, save_block_internals=True),
|
|
186
192
|
[(E.size,), (Block.size, E.size), (Block.size, E.size)],
|
|
193
|
+
None,
|
|
187
194
|
),
|
|
188
195
|
(
|
|
189
196
|
"internals",
|
|
190
197
|
ScanCheckpointPolicy(save_carries=False, save_block_internals=True),
|
|
191
198
|
[(E.size,), (Block.size, E.size), (Block.size, E.size)],
|
|
199
|
+
None,
|
|
192
200
|
),
|
|
193
201
|
(
|
|
194
202
|
"cos",
|
|
195
203
|
ScanCheckpointPolicy(save_carries=False, save_block_internals=["cos"]),
|
|
196
204
|
[(E.size,), (Block.size, E.size)],
|
|
205
|
+
None,
|
|
197
206
|
),
|
|
198
207
|
(
|
|
199
208
|
"sin",
|
|
200
209
|
ScanCheckpointPolicy(save_carries=True, save_block_internals=["sin"]),
|
|
201
210
|
[(E.size,), (Block.size, E.size), (Block.size, E.size)],
|
|
211
|
+
None,
|
|
212
|
+
),
|
|
213
|
+
("simple", ScanCheckpointPolicy(simple=True), [(E.size,), (Block.size, E.size)], None),
|
|
214
|
+
("nested", ScanCheckpointPolicy(simple=True, nested=2), [(E.size,), (2, E.size)], None),
|
|
215
|
+
(
|
|
216
|
+
"sin_offload",
|
|
217
|
+
ScanCheckpointPolicy(save_carries=True, offload_block_internals=["sin"]),
|
|
218
|
+
[(E.size,), (Block.size, E.size), (Block.size, E.size)],
|
|
219
|
+
["sin"],
|
|
202
220
|
),
|
|
203
|
-
("simple", ScanCheckpointPolicy(simple=True), [(E.size,), (Block.size, E.size)]),
|
|
204
|
-
("nested", ScanCheckpointPolicy(simple=True, nested=2), [(E.size,), (2, E.size)]),
|
|
205
221
|
],
|
|
206
222
|
)
|
|
207
|
-
def test_checkpoint_carries(name, policy, expected_scan_shapes):
|
|
223
|
+
def test_checkpoint_carries(name, policy, expected_scan_shapes, check_offloading):
|
|
208
224
|
class Module(eqx.Module):
|
|
209
225
|
named: hax.NamedArray
|
|
210
226
|
|
|
@@ -243,3 +259,29 @@ def test_checkpoint_carries(name, policy, expected_scan_shapes):
|
|
|
243
259
|
print(residual)
|
|
244
260
|
|
|
245
261
|
assert out_shapes == expected_scan_shapes, f"{name}: Expected {expected_scan_shapes}, got {out_shapes}"
|
|
262
|
+
|
|
263
|
+
# Add check for offloading if specified
|
|
264
|
+
if check_offloading is not None:
|
|
265
|
+
for name in check_offloading:
|
|
266
|
+
print(f"Checking offloading for {name}")
|
|
267
|
+
target = None
|
|
268
|
+
found_saved = False
|
|
269
|
+
for expr in jaxpr.jaxpr.eqns:
|
|
270
|
+
if expr.primitive.name == "scan":
|
|
271
|
+
inner_jaxpr = expr.params["jaxpr"]
|
|
272
|
+
for eqn in inner_jaxpr.eqns:
|
|
273
|
+
if eqn.primitive.name == "name":
|
|
274
|
+
this_name = eqn.params["name"]
|
|
275
|
+
if this_name == name:
|
|
276
|
+
# TODO in theory we can save more than one thing with the same name
|
|
277
|
+
# not gonna worry about that for now
|
|
278
|
+
target = eqn.outvars[0]
|
|
279
|
+
elif eqn.primitive.name == "device_put":
|
|
280
|
+
if eqn.invars[0] == target:
|
|
281
|
+
found_saved = True
|
|
282
|
+
break
|
|
283
|
+
# found scan
|
|
284
|
+
break
|
|
285
|
+
|
|
286
|
+
assert target is not None, f"Could not find named value for {name}"
|
|
287
|
+
assert found_saved, f"Could not find offloaded value for {name}"
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev343"
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|