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.
Files changed (106) hide show
  1. {haliax-1.4.dev343 → haliax-1.4.dev345}/PKG-INFO +1 -1
  2. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/scan.md +6 -0
  3. haliax-1.4.dev345/src/haliax/__about__.py +1 -0
  4. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/__init__.py +29 -4
  5. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/scan.py +10 -1
  6. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/jax_utils.py +7 -0
  7. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/random.py +22 -26
  8. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_scan.py +48 -6
  9. haliax-1.4.dev343/src/haliax/__about__.py +0 -1
  10. {haliax-1.4.dev343 → haliax-1.4.dev345}/.coveragerc +0 -0
  11. {haliax-1.4.dev343 → haliax-1.4.dev345}/.flake8 +0 -0
  12. {haliax-1.4.dev343 → haliax-1.4.dev345}/.github/workflows/publish_dev.yaml +0 -0
  13. {haliax-1.4.dev343 → haliax-1.4.dev345}/.github/workflows/run_pre_commit.yaml +0 -0
  14. {haliax-1.4.dev343 → haliax-1.4.dev345}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  15. {haliax-1.4.dev343 → haliax-1.4.dev345}/.github/workflows/run_tests.yaml +0 -0
  16. {haliax-1.4.dev343 → haliax-1.4.dev345}/.gitignore +0 -0
  17. {haliax-1.4.dev343 → haliax-1.4.dev345}/.pre-commit-config.yaml +0 -0
  18. {haliax-1.4.dev343 → haliax-1.4.dev345}/.readthedocs.yaml +0 -0
  19. {haliax-1.4.dev343 → haliax-1.4.dev345}/CONTRIBUTING.md +0 -0
  20. {haliax-1.4.dev343 → haliax-1.4.dev345}/LICENSE +0 -0
  21. {haliax-1.4.dev343 → haliax-1.4.dev345}/README.md +0 -0
  22. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/api.md +0 -0
  23. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/broadcasting.md +0 -0
  24. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/cheatsheet.md +0 -0
  25. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/css/material.css +0 -0
  26. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/css/mkdocstrings.css +0 -0
  27. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/faq.md +0 -0
  28. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/data_parallel_mesh.png +0 -0
  29. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  30. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/device_mesh_1d.png +0 -0
  31. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/device_mesh_1d_zero.png +0 -0
  32. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/device_mesh_2d.png +0 -0
  33. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  34. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  35. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  36. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  37. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/figures/device_mesh_2d_zero.png +0 -0
  38. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/fp8.md +0 -0
  39. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/index.md +0 -0
  40. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/indexing.md +0 -0
  41. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/matmul.md +0 -0
  42. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/nn.md +0 -0
  43. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/partitioning.md +0 -0
  44. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/rearrange.ipynb +0 -0
  45. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/rearrange.md +0 -0
  46. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/requirements.txt +0 -0
  47. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/state-dict.md +0 -0
  48. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/tutorial.md +0 -0
  49. {haliax-1.4.dev343 → haliax-1.4.dev345}/docs/vmap.md +0 -0
  50. {haliax-1.4.dev343 → haliax-1.4.dev345}/mkdocs.yml +0 -0
  51. {haliax-1.4.dev343 → haliax-1.4.dev345}/pyproject.toml +0 -0
  52. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/__init__.py +0 -0
  53. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/compile_utils.py +0 -0
  54. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/dot.py +0 -0
  55. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/einsum.py +0 -0
  56. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/fp8.py +0 -0
  57. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/parsing.py +0 -0
  58. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/rearrange.py +0 -0
  59. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/state_dict.py +0 -0
  60. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/_src/util.py +0 -0
  61. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/axis.py +0 -0
  62. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/core.py +0 -0
  63. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/debug.py +0 -0
  64. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/hof.py +0 -0
  65. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/__init__.py +0 -0
  66. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/activations.py +0 -0
  67. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/attention.py +0 -0
  68. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/conv.py +0 -0
  69. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/dropout.py +0 -0
  70. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/embedding.py +0 -0
  71. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/linear.py +0 -0
  72. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/loss.py +0 -0
  73. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/mlp.py +0 -0
  74. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/normalization.py +0 -0
  75. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/pool.py +0 -0
  76. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/nn/scan.py +0 -0
  77. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/ops.py +0 -0
  78. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/partitioning.py +0 -0
  79. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/quantization.py +0 -0
  80. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/specialized_fns.py +0 -0
  81. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/state_dict.py +0 -0
  82. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/tree_util.py +0 -0
  83. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/types.py +0 -0
  84. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/util.py +0 -0
  85. {haliax-1.4.dev343 → haliax-1.4.dev345}/src/haliax/wrap.py +0 -0
  86. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/core_test.py +0 -0
  87. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_attention.py +0 -0
  88. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_axis.py +0 -0
  89. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_conv.py +0 -0
  90. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_debug.py +0 -0
  91. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_dot.py +0 -0
  92. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_einsum.py +0 -0
  93. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_fp8.py +0 -0
  94. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_hof.py +0 -0
  95. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_int8.py +0 -0
  96. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_nn.py +0 -0
  97. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_ops.py +0 -0
  98. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_parsing.py +0 -0
  99. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_partitioning.py +0 -0
  100. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_pool.py +0 -0
  101. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_random.py +0 -0
  102. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_rearrange.py +0 -0
  103. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_specialized_fns.py +0 -0
  104. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_state_dict.py +0 -0
  105. {haliax-1.4.dev343 → haliax-1.4.dev345}/tests/test_tree_util.py +0 -0
  106. {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.dev343
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: Axis, *, start=0, step=1, dtype: Optional[DTypeLike] = None) -> NamedArray:
125
- """Version of jnp.arange that returns a NamedArray"""
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=axis.size) * step + start
130
- return NamedArray(arr, (axis,))
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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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 = _to_jax_shape(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
- ("disabled", ScanCheckpointPolicy(disable=True), [(E.size,), (Block.size, E.size), (Block.size, E.size)]),
177
- ("carry_true", True, [(E.size,), (Block.size, E.size)]),
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