haliax 1.4.dev315__tar.gz → 1.4.dev317__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 (99) hide show
  1. {haliax-1.4.dev315 → haliax-1.4.dev317}/PKG-INFO +1 -1
  2. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/api.md +1 -0
  3. haliax-1.4.dev317/src/haliax/__about__.py +1 -0
  4. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/__init__.py +4 -0
  5. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/axis.py +11 -0
  6. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/core.py +2 -3
  7. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/mlp.py +20 -21
  8. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/scan.py +1 -1
  9. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/core_test.py +26 -394
  10. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_attention.py +2 -6
  11. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_axis.py +2 -4
  12. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_einsum.py +1 -6
  13. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_nn.py +5 -15
  14. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_rearrange.py +1 -5
  15. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_specialized_fns.py +3 -4
  16. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_tree_util.py +1 -4
  17. haliax-1.4.dev315/src/haliax/__about__.py +0 -1
  18. {haliax-1.4.dev315 → haliax-1.4.dev317}/.coveragerc +0 -0
  19. {haliax-1.4.dev315 → haliax-1.4.dev317}/.flake8 +0 -0
  20. {haliax-1.4.dev315 → haliax-1.4.dev317}/.github/workflows/publish_dev.yaml +0 -0
  21. {haliax-1.4.dev315 → haliax-1.4.dev317}/.github/workflows/run_pre_commit.yaml +0 -0
  22. {haliax-1.4.dev315 → haliax-1.4.dev317}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  23. {haliax-1.4.dev315 → haliax-1.4.dev317}/.github/workflows/run_tests.yaml +0 -0
  24. {haliax-1.4.dev315 → haliax-1.4.dev317}/.gitignore +0 -0
  25. {haliax-1.4.dev315 → haliax-1.4.dev317}/.pre-commit-config.yaml +0 -0
  26. {haliax-1.4.dev315 → haliax-1.4.dev317}/.readthedocs.yaml +0 -0
  27. {haliax-1.4.dev315 → haliax-1.4.dev317}/CONTRIBUTING.md +0 -0
  28. {haliax-1.4.dev315 → haliax-1.4.dev317}/LICENSE +0 -0
  29. {haliax-1.4.dev315 → haliax-1.4.dev317}/README.md +0 -0
  30. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/broadcasting.md +0 -0
  31. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/cheatsheet.md +0 -0
  32. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/css/material.css +0 -0
  33. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/css/mkdocstrings.css +0 -0
  34. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/faq.md +0 -0
  35. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/data_parallel_mesh.png +0 -0
  36. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  37. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/device_mesh_1d.png +0 -0
  38. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/device_mesh_1d_zero.png +0 -0
  39. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/device_mesh_2d.png +0 -0
  40. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  41. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  42. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  43. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  44. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/device_mesh_2d_zero.png +0 -0
  45. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/fp8.md +0 -0
  46. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/hof.md +0 -0
  47. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/index.md +0 -0
  48. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/indexing.md +0 -0
  49. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/matmul.md +0 -0
  50. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/nn.md +0 -0
  51. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/partitioning.md +0 -0
  52. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/rearrange.ipynb +0 -0
  53. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/rearrange.md +0 -0
  54. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/requirements.txt +0 -0
  55. {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/tutorial.md +0 -0
  56. {haliax-1.4.dev315 → haliax-1.4.dev317}/mkdocs.yml +0 -0
  57. {haliax-1.4.dev315 → haliax-1.4.dev317}/pyproject.toml +0 -0
  58. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/_src/__init__.py +0 -0
  59. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/_src/compile_utils.py +0 -0
  60. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/_src/dot.py +0 -0
  61. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/_src/einsum.py +0 -0
  62. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/_src/fp8.py +0 -0
  63. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/_src/parsing.py +0 -0
  64. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/_src/rearrange.py +0 -0
  65. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/_src/util.py +0 -0
  66. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/debug.py +0 -0
  67. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/hof.py +0 -0
  68. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/jax_utils.py +0 -0
  69. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/__init__.py +0 -0
  70. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/activations.py +0 -0
  71. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/attention.py +0 -0
  72. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/conv.py +0 -0
  73. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/dropout.py +0 -0
  74. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/embedding.py +0 -0
  75. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/linear.py +0 -0
  76. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/loss.py +0 -0
  77. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/normalization.py +0 -0
  78. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/pool.py +0 -0
  79. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/ops.py +0 -0
  80. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/partitioning.py +0 -0
  81. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/quantization.py +0 -0
  82. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/random.py +0 -0
  83. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/specialized_fns.py +0 -0
  84. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/tree_util.py +0 -0
  85. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/types.py +0 -0
  86. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/util.py +0 -0
  87. {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/wrap.py +0 -0
  88. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_conv.py +0 -0
  89. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_debug.py +0 -0
  90. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_dot.py +0 -0
  91. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_fp8.py +0 -0
  92. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_hof.py +0 -0
  93. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_ops.py +0 -0
  94. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_parsing.py +0 -0
  95. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_partitioning.py +0 -0
  96. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_pool.py +0 -0
  97. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_random.py +0 -0
  98. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_scan.py +0 -0
  99. {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_utils.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.3
2
2
  Name: haliax
3
- Version: 1.4.dev315
3
+ Version: 1.4.dev317
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/
@@ -27,6 +27,7 @@ Occasionally, an axis size can be inferred in some circumstances but not others.
27
27
 
28
28
  ### Axis Manipulation
29
29
 
30
+ ::: haliax.make_axes
30
31
  ::: haliax.axis.axis_name
31
32
  ::: haliax.axis.concat_axes
32
33
  ::: haliax.axis.union_axes
@@ -0,0 +1 @@
1
+ __version__ = "1.4.dev317"
@@ -32,6 +32,7 @@ from .axis import (
32
32
  ds,
33
33
  dslice,
34
34
  eliminate_axes,
35
+ make_axes,
35
36
  selects_axis,
36
37
  )
37
38
  from .core import (
@@ -893,6 +894,9 @@ __all__ = [
893
894
  "AxisSpec",
894
895
  "AxisSelection",
895
896
  "AxisSelector",
897
+ "make_axes",
898
+ "axis_name",
899
+ "axis_size",
896
900
  "NamedArray",
897
901
  "broadcast_to",
898
902
  "broadcast_axis",
@@ -28,6 +28,17 @@ class Axis:
28
28
  return f"{self.name}({self.size})"
29
29
 
30
30
 
31
+ def make_axes(**kwargs: int) -> Tuple[Axis, ...]:
32
+ """
33
+ Convenience function for creating a tuple of Axis objects.
34
+
35
+ Example:
36
+ ```
37
+ X, Y = axes(X=10, Y=20)
38
+ """
39
+ return tuple(Axis(name, size) for name, size in kwargs.items())
40
+
41
+
31
42
  AxisSelector = Union[Axis, str]
32
43
  """AxisSelector is a type that can be used to select a single axis from an array. str or Axis"""
33
44
  AxisSelection = Union[AxisSelector, Sequence[AxisSelector]]
@@ -373,11 +373,10 @@ class NamedArray:
373
373
 
374
374
  Supports indexing like:
375
375
 
376
- >>> X = Axis("x", 10)
377
- >>> Y = Axis("y", 20)
376
+ >>> X, Y = haliax.make_axes(X=10, Y=20)
378
377
  >>> arr = haliax.random.randint(jax.random.PRNGKey(0), (X, Y), 0, X.size)
379
378
  # slice with ints or slices
380
- >>> arr[{"x": 1, "y": slice(0,10,new_axis=2)}]
379
+ >>> arr[{"x": 1, "y": slice(0,10,2)}]
381
380
  >>> Z = Axis("z", 3)
382
381
  # so-called "advanced indexing" with NamedArrays.
383
382
  >>> index_arr = NamedArray(np.array([1, 2, 3]), Z)
@@ -44,6 +44,7 @@ class MLP(eqx.Module):
44
44
  depth: int,
45
45
  activation: Callable = relu,
46
46
  *,
47
+ out_first: bool = True,
47
48
  use_bias: bool = True,
48
49
  use_final_bias: bool = True,
49
50
  key: PRNGKeyArray,
@@ -57,36 +58,34 @@ class MLP(eqx.Module):
57
58
 
58
59
  layers = []
59
60
 
61
+ kwargs: dict = {
62
+ "use_bias": use_bias,
63
+ "dot_general": dot_general,
64
+ "init_scale": init_scale,
65
+ "out_first": out_first,
66
+ }
67
+
68
+ last_kwargs: dict = {
69
+ "use_bias": use_final_bias,
70
+ "dot_general": dot_general,
71
+ "init_scale": init_scale,
72
+ "out_first": out_first,
73
+ }
74
+
60
75
  if depth == 0:
61
76
  # special case: no hidden layers
62
- layers.append(
63
- Linear.init(
64
- Input, Output, use_bias=use_final_bias, key=keys[0], dot_general=dot_general, init_scale=init_scale
65
- )
66
- )
77
+ layers.append(Linear.init(Input, Output, key=keys[0], **last_kwargs))
67
78
  else:
68
79
  # first hidden layer
69
- layers.append(
70
- Linear.init(
71
- Input, Width, use_bias=use_bias, key=keys[0], dot_general=dot_general, init_scale=init_scale
72
- )
73
- )
80
+ layers.append(Linear.init(Input, Width, key=keys[0], **kwargs))
74
81
  # middle hidden layers
75
82
  cur = Width
76
83
  next = Width2
77
84
  for i in range(1, depth):
78
- layers.append(
79
- Linear.init(
80
- cur, next, use_bias=use_bias, key=keys[i], dot_general=dot_general, init_scale=init_scale
81
- )
82
- )
85
+ layers.append(Linear.init(cur, next, key=keys[i], **kwargs))
83
86
  cur, next = next, cur
84
- # final hidden layer
85
- layers.append(
86
- Linear.init(
87
- cur, Output, use_bias=use_final_bias, key=keys[-1], dot_general=dot_general, init_scale=init_scale
88
- )
89
- )
87
+ # final layer
88
+ layers.append(Linear.init(cur, Output, key=keys[-1], **last_kwargs))
90
89
 
91
90
  return MLP(
92
91
  layers=tuple(layers),
@@ -300,7 +300,7 @@ class Stacked(eqx.Module, Generic[M]):
300
300
  """
301
301
 
302
302
  def unbatch_leaf(x):
303
- if haliax.is_named_array(x):
303
+ if isinstance(x, haliax.core.NamedArray):
304
304
  if haliax.selects_axis(x.axes, self.Block):
305
305
  return haliax.unbind(x, self.Block)
306
306
  else:
@@ -9,9 +9,7 @@ from haliax import Axis, NamedArray
9
9
 
10
10
 
11
11
  def test_unary_np_functions():
12
- Height = Axis("Height", 2)
13
- Width = Axis("Width", 3)
14
- Depth = Axis("Depth", 4)
12
+ Height, Width, Depth = hax.make_axes(Height=2, Width=3, Depth=4)
15
13
 
16
14
  m1 = NamedArray(jnp.ones((Height.size, Width.size, Depth.size)), (Height, Width, Depth))
17
15
 
@@ -29,9 +27,7 @@ def test_unary_np_functions():
29
27
 
30
28
 
31
29
  def test_reduction_functions():
32
- Height = Axis("Height", 2)
33
- Width = Axis("Width", 3)
34
- Depth = Axis("Depth", 4)
30
+ Height, Width, Depth = hax.make_axes(Height=2, Width=3, Depth=4)
35
31
 
36
32
  rand_m = jax.random.uniform(PRNGKey(0), (Height.size, Width.size, Depth.size))
37
33
 
@@ -69,9 +65,7 @@ def test_reduction_functions():
69
65
 
70
66
 
71
67
  def test_reduction_functions_with_where():
72
- H = Axis("H", 2)
73
- W = Axis("W", 3)
74
- D = Axis("D", 4)
68
+ H, W, D = hax.make_axes(H=2, W=3, D=4)
75
69
 
76
70
  rand_m = jax.random.uniform(PRNGKey(0), (H.size, W.size, D.size))
77
71
 
@@ -118,9 +112,7 @@ def test_reduction_functions_with_where():
118
112
 
119
113
 
120
114
  def test_split():
121
- Height = Axis("Height", 2)
122
- Width = Axis("Width", 3)
123
- Depth = Axis("Depth", 4)
115
+ Height, Width, Depth = hax.make_axes(Height=2, Width=3, Depth=4)
124
116
 
125
117
  D10 = Axis("Depth", Depth.size * 10)
126
118
 
@@ -144,9 +136,7 @@ def test_split():
144
136
 
145
137
 
146
138
  def test_take():
147
- Height = Axis("Height", 2)
148
- Width = Axis("Width", 3)
149
- Depth = Axis("Depth", 4)
139
+ Height, Width, Depth = hax.make_axes(Height=2, Width=3, Depth=4)
150
140
  named1 = hax.random.uniform(PRNGKey(0), (Height, Width, Depth))
151
141
 
152
142
  assert jnp.all(jnp.equal(hax.take(named1, Height, 0).array, named1.array[0]))
@@ -178,9 +168,7 @@ def test_take():
178
168
 
179
169
 
180
170
  def test_take_overlapping_names():
181
- Height = Axis("Height", 20)
182
- Width = Axis("Width", 30)
183
- Depth = Axis("Depth", 40)
171
+ Height, Width, Depth = hax.make_axes(Height=20, Width=30, Depth=40)
184
172
  named1 = hax.random.uniform(PRNGKey(0), (Height, Width, Depth))
185
173
 
186
174
  Height2 = Axis("Height", 10)
@@ -198,9 +186,7 @@ def test_take_overlapping_2():
198
186
  def cross_entropy(logits: hax.NamedArray, labels: hax.NamedArray) -> hax.NamedArray:
199
187
  return hax.take(logits, Embed, labels) # extract log probability of the correct token
200
188
 
201
- Embed = Axis("Embed", 10)
202
- Block = Axis("Block", 20)
203
- Batch = Axis("Batch", 30)
189
+ Embed, Block, Batch = hax.make_axes(Embed=10, Block=20, Batch=30)
204
190
  logits = hax.random.uniform(PRNGKey(0), (Batch, Block, Embed))
205
191
  labels = hax.random.randint(PRNGKey(0), (Batch, Block), 0, Embed.size)
206
192
 
@@ -223,9 +209,7 @@ def test_take_overlapping_2():
223
209
 
224
210
 
225
211
  def test_cumsum_etc():
226
- Height = Axis("Height", 2)
227
- Width = Axis("Width", 3)
228
- Depth = Axis("Depth", 4)
212
+ Height, Width, Depth = hax.make_axes(Height=2, Width=3, Depth=4)
229
213
 
230
214
  named1 = hax.random.uniform(PRNGKey(0), (Height, Width, Depth))
231
215
 
@@ -255,10 +239,7 @@ def test_cumsum_etc():
255
239
 
256
240
 
257
241
  def test_rearrange():
258
- H = Axis("H", 2)
259
- W = Axis("W", 3)
260
- D = Axis("D", 4)
261
- C = Axis("C", 5)
242
+ H, W, D, C = hax.make_axes(H=2, W=3, D=4, C=5)
262
243
 
263
244
  named1 = hax.random.uniform(PRNGKey(0), (H, W, D, C))
264
245
 
@@ -304,9 +285,7 @@ def test_rearrange():
304
285
 
305
286
  def test_rearrange_unused_ellipsis():
306
287
  # Make sure we just ignore the ellipsis if all axes are specified in addition
307
- H = Axis("Height", 2)
308
- W = Axis("Width", 3)
309
- D = Axis("Depth", 4)
288
+ H, W, D = hax.make_axes(Height=2, Width=3, Depth=4)
310
289
 
311
290
  named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
312
291
 
@@ -321,7 +300,7 @@ def test_rearrange_unused_ellipsis():
321
300
 
322
301
 
323
302
  def test_arange():
324
- H = Axis("Height", 10)
303
+ (H,) = hax.make_axes(H=10)
325
304
 
326
305
  assert jnp.all(jnp.equal(hax.arange(H).array, jnp.arange(10)))
327
306
  assert hax.arange(H).axes == (H,)
@@ -334,8 +313,7 @@ def test_arange():
334
313
 
335
314
 
336
315
  def test_stack():
337
- H = Axis("H", 4)
338
- W = Axis("W", 3)
316
+ H, W = hax.make_axes(H=4, W=3)
339
317
 
340
318
  named1 = hax.random.uniform(PRNGKey(0), (H, W))
341
319
  named2 = hax.random.uniform(PRNGKey(1), (H, W))
@@ -350,9 +328,8 @@ def test_stack():
350
328
 
351
329
 
352
330
  def test_concatenate():
353
- H1 = Axis("H", 4)
354
- H2 = Axis("H", 3)
355
- W = Axis("W", 3)
331
+ H1, W = hax.make_axes(H=4, W=3)
332
+ H2 = H1.resize(3)
356
333
 
357
334
  named1 = hax.random.uniform(PRNGKey(0), (H1, W))
358
335
  named2 = hax.random.uniform(PRNGKey(1), (H2, W))
@@ -392,8 +369,7 @@ def test_repeat():
392
369
  # [3, 4],
393
370
  # [3, 4]])
394
371
 
395
- H = Axis("H", 2)
396
- W = Axis("W", 2)
372
+ H, W = hax.make_axes(H=2, W=2)
397
373
 
398
374
  named1 = hax.named([[1, 2], [3, 4]], (H, W))
399
375
 
@@ -459,9 +435,7 @@ def test_tile():
459
435
 
460
436
 
461
437
  def test_unflatten_axis():
462
- H = Axis("Height", 2)
463
- W = Axis("Width", 3)
464
- D = Axis("Depth", 4)
438
+ H, W, D = hax.make_axes(Height=2, Width=3, Depth=4)
465
439
 
466
440
  named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
467
441
  flattened_HW = named1.flatten_axes((H, W), "Z")
@@ -483,9 +457,7 @@ def test_unflatten_axis():
483
457
 
484
458
 
485
459
  def test_ravel():
486
- H = Axis("Height", 2)
487
- W = Axis("Width", 3)
488
- D = Axis("Depth", 4)
460
+ H, W, D = hax.make_axes(Height=2, Width=3, Depth=4)
489
461
 
490
462
  named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
491
463
  raveled = named1.ravel("Z")
@@ -496,13 +468,8 @@ def test_ravel():
496
468
 
497
469
 
498
470
  def test_rename():
499
- H = Axis("H", 2)
500
- W = Axis("W", 3)
501
- D = Axis("D", 4)
502
-
503
- H2 = Axis("H2", 2)
504
- W2 = Axis("W2", 3)
505
- D2 = Axis("D2", 4)
471
+ H, W, D = hax.make_axes(H=2, W=3, D=4)
472
+ H2, W2, D2 = hax.make_axes(H2=2, W2=3, D2=4)
506
473
 
507
474
  named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
508
475
 
@@ -517,9 +484,7 @@ def test_rename():
517
484
 
518
485
 
519
486
  def test_index():
520
- H = Axis("H", 20)
521
- W = Axis("W", 30)
522
- D = Axis("D", 40)
487
+ H, W, D = hax.make_axes(H=20, W=30, D=40)
523
488
 
524
489
  named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
525
490
 
@@ -544,9 +509,7 @@ def test_index():
544
509
 
545
510
 
546
511
  def test_index_with_tracer():
547
- H = Axis("H", 20)
548
- W = Axis("W", 30)
549
- D = Axis("D", 40)
512
+ H, W, D = hax.make_axes(H=20, W=30, D=40)
550
513
  named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
551
514
 
552
515
  @jax.jit
@@ -562,12 +525,7 @@ def test_index_with_tracer():
562
525
 
563
526
  def test_index_array_slices():
564
527
  # fancier tests with array slices with named array args
565
- H = Axis("H", 10)
566
- W = Axis("W", 20)
567
- D = Axis("D", 30)
568
- C = Axis("C", 40)
569
- Q = Axis("Q", 50)
570
- I0 = Axis("I0", 10)
528
+ H, W, D, C, Q, I0 = hax.make_axes(H=10, W=20, D=30, C=40, Q=50, I0=10)
571
529
 
572
530
  named1 = hax.random.uniform(PRNGKey(0), (H, W, D, C, Q))
573
531
  index_1 = hax.random.randint(PRNGKey(0), (I0,), 0, H.size)
@@ -586,9 +544,7 @@ def test_index_array_slices():
586
544
  # https://numpy.org/doc/stable/user/basics.indexing.html#combining-advanced-and-basic-indexing
587
545
  # Example
588
546
  # Let x.shape be (10, 20, 30, 40, 50) and suppose ind_1 and ind_2 can be broadcast to the shape (2, 3, 4).
589
- I1 = Axis("I1", 2)
590
- I2 = Axis("I2", 3)
591
- I3 = Axis("I3", 4)
547
+ I1, I2, I3 = hax.make_axes(I1=2, I2=3, I3=4)
592
548
 
593
549
  ind_1 = hax.random.randint(PRNGKey(0), (I2, I3), 0, W.size)
594
550
  ind_2 = hax.random.randint(PRNGKey(0), (I1, I3), 0, D.size)
@@ -615,9 +571,7 @@ def test_index_array_slices():
615
571
  def test_slice_nd_shorthand_syntax():
616
572
  # syntax like arr["X", 0:10, "Y", 0:10] is supported
617
573
 
618
- H = Axis("H", 10)
619
- W = Axis("W", 20)
620
- D = Axis("D", 30)
574
+ H, W, D = hax.make_axes(H=10, W=20, D=30)
621
575
 
622
576
  named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
623
577
 
@@ -625,9 +579,7 @@ def test_slice_nd_shorthand_syntax():
625
579
 
626
580
 
627
581
  def test_slice_nd_dslice():
628
- H = Axis("H", 10)
629
- W = Axis("W", 20)
630
- D = Axis("D", 30)
582
+ H, W, D = hax.make_axes(H=10, W=20, D=30)
631
583
 
632
584
  named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
633
585
  from haliax import ds
@@ -641,9 +593,7 @@ def test_slice_nd_dslice():
641
593
 
642
594
  def test_slice_nd_array_present_dims():
643
595
  # tests slicing with arrays that are already present in the named array, which is sometimes ok
644
- H = Axis("H", 10)
645
- W = Axis("W", 20)
646
- D = Axis("D", 30)
596
+ H, W, D = hax.make_axes(H=10, W=20, D=30)
647
597
 
648
598
  named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
649
599
 
@@ -653,321 +603,3 @@ def test_slice_nd_array_present_dims():
653
603
  assert jnp.all(jnp.equal(named1[{"H": index1}].array, named1.array[index1.array, :, :]))
654
604
 
655
605
  # this is not ok, since the H would not be eliminated
656
- with pytest.raises(ValueError):
657
- named1[{W: index1}]
658
-
659
- # this is not ok, but is trickier because the H has a different size
660
- H2 = H.resize(5)
661
- index2 = hax.random.randint(PRNGKey(0), (H2,), 0, H.size)
662
- with pytest.raises(ValueError):
663
- named1[{W: index2}]
664
-
665
- # this is ok, since the H would be eliminated anyway
666
- assert jnp.all(jnp.equal(named1[{"H": index2}].array, named1.array[index2.array, :, :]))
667
-
668
-
669
- def test_slice_nd_array_unnamed_slice():
670
- # tests slicing with arrays that are already present in the named array, which is sometimes ok
671
- H = Axis("H", 10)
672
- W = Axis("W", 20)
673
- D = Axis("D", 30)
674
-
675
- named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
676
-
677
- index1 = jax.random.randint(PRNGKey(1), (4,), 0, H.size)
678
- assert jnp.all(jnp.equal(named1[{"H": index1}].array, named1.array[index1, :, :]))
679
-
680
- # hidden behavior: if we also pass in an H index to e.g. D, it is zipped together
681
- index2 = hax.random.randint(PRNGKey(2), Axis("H", 4), 0, D.size)
682
- assert jnp.all(jnp.equal(named1[{"H": index1, "D": index2}].array, named1.array[index1, :, index2.array]))
683
-
684
- # this is different though:
685
- index2r = index2.array
686
- assert jnp.all(
687
- jnp.equal(
688
- named1[{"H": index1, "D": index2r}].array, named1.array[index1.reshape(1, -1), :, index2r.reshape(-1, 1)]
689
- )
690
- )
691
- assert named1[{"H": index1, "D": index2r}].shape != named1[{"H": index1, "D": index2}].shape
692
-
693
- index1 = list(index1)
694
- assert jnp.all(jnp.equal(named1[{"H": index1}].array, named1.array[index1, :, :]))
695
-
696
-
697
- def test_full_indexing_returns_named_array():
698
- H = Axis("H", 10)
699
- W = Axis("W", 20)
700
- D = Axis("D", 30)
701
-
702
- named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
703
- sliced = named1[{"H": 0, "W": 0, "D": 0}]
704
-
705
- assert isinstance(sliced, NamedArray)
706
- assert sliced.shape == {}
707
-
708
-
709
- def test_indexing_bug_from_docs():
710
- X = hax.Axis("X", 10)
711
- Y = hax.Axis("Y", 20)
712
- Z = hax.Axis("Z", 30)
713
-
714
- a = hax.random.uniform(jax.random.PRNGKey(0), (X, Y, Z))
715
-
716
- I1 = hax.Axis("I1", 5)
717
- I2 = hax.Axis("I2", 5)
718
- I3 = hax.Axis("I3", 5)
719
- ind1 = hax.random.randint(jax.random.PRNGKey(0), (I1,), 0, 10)
720
- ind2 = hax.random.randint(jax.random.PRNGKey(0), (I2, I3), 0, 20)
721
-
722
- # assert a[{"X": ind1, "Y": ind2}].axes == (I1, I2, I3, Z)
723
- assert a[{"X": ind1, "Y": ind2, "Z": 3}].axes == (I1, I2, I3)
724
-
725
-
726
- def test_duplicate_axis_names_in_slicing():
727
- X = hax.Axis("X", 10)
728
- Y = hax.Axis("Y", 20)
729
- Z = hax.Axis("Z", 30)
730
-
731
- X2 = hax.Axis("X", 5)
732
- Y2 = hax.Axis("Y", 5)
733
-
734
- a = hax.random.uniform(jax.random.PRNGKey(0), (X, Y, Z))
735
- ind1 = hax.random.randint(jax.random.PRNGKey(0), (X2,), 0, 10)
736
- ind2 = hax.random.randint(jax.random.PRNGKey(0), (Y2,), 0, 10)
737
-
738
- a[{"X": ind1, "Y": ind2}] # returns a NamedArray with axes = Axis("X", 5), Axis("Y", 5), Axis("Z", 30)
739
-
740
- with pytest.raises(ValueError):
741
- a[{"Y": ind1}] # error, "X" is not eliminated by the indexing operation
742
-
743
- a[{"X": ind2, "Y": ind1}] # ok, because X and Y are eliminated by the indexing operation
744
-
745
-
746
- def test_slice_old_style():
747
- H = Axis("H", 10)
748
- W = Axis("W", 20)
749
- D = Axis("D", 30)
750
-
751
- named1 = hax.random.randint(PRNGKey(0), (H, W, D), minval=0, maxval=10)
752
-
753
- assert jnp.all(named1.slice("H", start=4, length=2).array == named1.array[4:6, :, :])
754
- assert jnp.all(named1.slice("W", start=4, length=2).array == named1.array[:, 4:6, :])
755
- assert jnp.all(named1.slice("D", start=4, length=2).array == named1.array[:, :, 4:6])
756
-
757
- H2 = Axis("H2", 5)
758
- W2 = Axis("W2", 10)
759
- D2 = Axis("D2", 15)
760
-
761
- assert jnp.all(named1.slice("H", H2, start=4).array == named1.array[4 : 4 + H2.size, :, :])
762
- assert jnp.all(named1.slice("W", W2, start=4).array == named1.array[:, 4 : 4 + W2.size, :])
763
- assert jnp.all(named1.slice("D", D2, start=4).array == named1.array[:, :, 4 : 4 + D2.size])
764
-
765
-
766
- def test_slice_new_style():
767
- H = Axis("H", 10)
768
- W = Axis("W", 20)
769
- D = Axis("D", 30)
770
-
771
- named1 = hax.random.randint(PRNGKey(0), (H, W, D), minval=0, maxval=10)
772
-
773
- x1 = named1.slice({"H": 4, "W": 5, "D": 7}, length={"H": 2, "W": 3, "D": 4})
774
- assert jnp.all(x1.array == named1.array[4:6, 5:8, 7:11])
775
-
776
- with pytest.raises(TypeError):
777
- named1.slice({"H": 4, "W": 5, "D": 7}, length={"H": 2, "W": 3, "D": 4}, start={"H": 1, "W": 2, "D": 3})
778
-
779
- with pytest.raises(ValueError):
780
- named1.slice({"H": 4, "W": 5, "D": 7}, length={"H": 2, "W": 3})
781
-
782
- H2 = Axis("H2", 5)
783
- W2 = Axis("W2", 10)
784
- D2 = Axis("D2", 15)
785
-
786
- x2 = named1.slice({"H": 4, "W": 5, "D": 7}, length={"H": H2, "W": W2, "D": D2})
787
- assert jnp.all(x2.array == named1.array[4 : 4 + H2.size, 5 : 5 + W2.size, 7 : 7 + D2.size])
788
-
789
-
790
- def test_updated_slice():
791
- H = Axis("H", 10)
792
- W = Axis("W", 20)
793
- D = Axis("D", 30)
794
-
795
- H2 = H.resize(5)
796
- W2 = W.resize(10)
797
- D2 = D.resize(15)
798
-
799
- named1 = hax.random.randint(PRNGKey(0), (H, W, D), minval=0, maxval=10)
800
- named2 = hax.random.randint(PRNGKey(0), (H2, W2, D2), minval=10, maxval=30)
801
-
802
- named1_updated = named1.updated_slice({"H": 0, "W": 0, "D": 0}, named2)
803
-
804
- assert named1_updated.axes == named1.axes
805
- assert jnp.all(named1_updated["H", 0 : H2.size, "W", 0 : W2.size, "D", 0 : D2.size].array == named2.array)
806
-
807
- # test broadcasting
808
- for pair in [(H2, D2), (H2, W2), (W2, D2), (D2, H2), (D2, W2), (W2, H2)]:
809
- n3 = hax.random.randint(PRNGKey(0), pair, minval=10, maxval=30)
810
- named1_updated = named1.updated_slice({ax.name: 0 for ax in pair}, n3)
811
- assert named1_updated.axes == named1.axes
812
- assert jnp.all((named1_updated[{ax.name: slice(0, ax.size) for ax in pair}] == n3).array)
813
- # check that the array outside the slice is unchanged
814
- assert jnp.all(
815
- (
816
- named1_updated[{ax.name: slice(ax.size, None) for ax in pair}]
817
- == named1[{ax.name: slice(ax.size, None) for ax in pair}]
818
- ).array
819
- )
820
-
821
-
822
- def test_updated_slice_extra_update_axis_errors():
823
- H = Axis("H", 10)
824
- W = Axis("W", 20)
825
- D = Axis("D", 30)
826
-
827
- named1 = hax.random.randint(PRNGKey(0), (H, W, D), minval=0, maxval=10)
828
- named2 = hax.random.randint(PRNGKey(0), (H, W, D), minval=10, maxval=30)
829
-
830
- with pytest.raises(ValueError):
831
- named1.updated_slice({"H": 0, "W": 0, "D": 0, "extra": 0}, named2)
832
-
833
- with pytest.raises(ValueError):
834
- named3 = hax.random.randint(PRNGKey(0), (H, W), minval=10, maxval=30)
835
- named3.updated_slice({"H": 0, "W": 0}, named2)
836
-
837
-
838
- def test_order_of_transpose_add():
839
- H = Axis("H", 10)
840
- W = Axis("W", 20)
841
-
842
- named1 = hax.random.randint(PRNGKey(0), (H, W), minval=0, maxval=10)
843
- named2 = hax.random.randint(PRNGKey(0), (W, H), minval=10, maxval=30)
844
-
845
- assert (named1 + named2).axes == (H, W)
846
- assert jnp.all((named1 + named2).array == named1.array + named2.array.T)
847
-
848
-
849
- def test_nice_short_string_in_named_array():
850
- H = Axis("H", 10)
851
- W = Axis("W", 20)
852
-
853
- named1 = hax.random.randint(PRNGKey(0), (H, W), minval=0, maxval=10)
854
-
855
- assert str(named1).startswith("NamedArray(int32{'H': 10, 'W': 20}")
856
-
857
-
858
- def test_nice_short_string_in_named_array_in_eqx_module():
859
- H = Axis("H", 10)
860
- W = Axis("W", 20)
861
-
862
- named1 = hax.random.randint(PRNGKey(0), (H, W), minval=0, maxval=10)
863
-
864
- class TestModule(eqx.Module):
865
- named1: NamedArray
866
-
867
- mod = TestModule(named1)
868
-
869
- assert str(mod).startswith("TestModule(named1=Named(int32{'H': 10, 'W': 20}))")
870
-
871
-
872
- def test_named_arrays_work_in_eqxi_while_loop():
873
- H = Axis("H", 10)
874
- W = Axis("W", 20)
875
-
876
- named1 = hax.random.uniform(PRNGKey(0), (H, W))
877
-
878
- import equinox.internal as eqxi
879
-
880
- def body_fun(t):
881
- i, named1 = t
882
- return i + 1, named1 + named1
883
-
884
- def cond_fun(t):
885
- i, named1 = t
886
- return i < 10
887
-
888
- def loss_fun(named1):
889
- i, named1 = eqxi.while_loop(cond_fun, body_fun, (0, named1), kind="checkpointed", max_steps=10)
890
- return named1.sum().scalar()
891
-
892
- grad_fun = eqx.filter_value_and_grad(loss_fun)
893
-
894
- grad_fun(named1)
895
-
896
-
897
- def test_at_for_in_placeish():
898
- H = Axis("H", 10)
899
- W = Axis("W", 20)
900
-
901
- named1 = hax.random.uniform(PRNGKey(0), (H, W))
902
-
903
- named1_at = named1.at[H, 0].set(0)
904
-
905
- assert jnp.all(jnp.equal(named1_at[H, 0].array, 0))
906
- assert jnp.all(named1_at[H, 1:].array == named1[H, 1:].array)
907
-
908
- # test add, multiply, power, etc.
909
- named1_at = named1.at[H, 0].add(1)
910
- assert jnp.all(named1_at.array == named1.array.at[0].add(1))
911
-
912
- named1_at = named1.at[H, 0].multiply(2)
913
- assert jnp.all(named1_at.array == named1.array.at[0].multiply(2))
914
-
915
- named1_at = named1.at[H, 0].power(2)
916
- assert jnp.all(named1_at.array == named1.array.at[0].power(2))
917
-
918
- named1_at = named1.at[H, 0].divide(2)
919
- assert jnp.all(named1_at.array == named1.array.at[0].divide(2))
920
-
921
- named1_at = named1.at[H, 0].apply(hax.square)
922
- assert jnp.all(named1_at.array == named1.array.at[0].apply(jnp.square))
923
-
924
- named1_at = named1.at[H, 0].max(0.5)
925
- assert jnp.all(named1_at.array == named1.array.at[0].max(0.5))
926
-
927
- named1_at = named1.at[H, 0].min(0.5)
928
- assert jnp.all(named1_at.array == named1.array.at[0].min(0.5))
929
-
930
-
931
- def test_at_with_fancy_indexing():
932
- H = Axis("H", 10)
933
- W = Axis("W", 20)
934
- I0 = Axis("I0", 5)
935
- I1 = Axis("I1", 5)
936
-
937
- named1 = hax.random.uniform(PRNGKey(0), (H, W))
938
- ind1 = hax.random.randint(PRNGKey(0), (I0,), 0, H.size)
939
- ind2 = hax.random.randint(PRNGKey(0), (I1,), 0, W.size)
940
-
941
- named1_at = named1.at[H, ind1].set(0)
942
- assert jnp.all(named1_at.array == named1.array.at[ind1.array].set(0))
943
-
944
- named1_at = named1.at[H, ind1].add(1, mode="clip")
945
- assert jnp.all(named1_at.array == named1.array.at[ind1.array].add(1, mode="clip"))
946
-
947
- named1_at = named1.at[H, ind1, W, ind2].set(0)
948
- assert jnp.all(named1_at.array == named1.array.at[ind1.array.reshape(-1, 1), ind2.array.reshape(1, -1)].set(0))
949
-
950
- # dslices
951
- from haliax import ds
952
-
953
- named1_at = named1.at[H, ds(3, 5)].set(0)
954
- assert jnp.all(named1_at.array == named1.array.at[3:8].set(0))
955
-
956
- named1_at = named1.at[H, ds(3, 5), W, ind2].power(2)
957
- assert jnp.all(named1_at.array == named1.array.at[3:8, ind2.array].power(2))
958
-
959
-
960
- def test_slice_dslice_and_array():
961
- H = Axis("H", 10)
962
- W = Axis("W", 20)
963
- I0 = Axis("I0", 5)
964
-
965
- named1 = hax.random.uniform(PRNGKey(0), (H, W))
966
- ind2 = hax.random.randint(PRNGKey(0), (I0,), 0, W.size)
967
-
968
- from haliax import ds
969
-
970
- named1.array.at[3:8, ind2.array].add(jnp.full((5, 5), 2))
971
-
972
- named1_at = named1.at[H, ds(3, 5), W, ind2].add(2)
973
- assert jnp.all(named1_at.array == named1.array.at[3:8, ind2.array].add(2))
@@ -93,8 +93,7 @@ def test_alibi_attention_compared_to_hf():
93
93
  import torch
94
94
  from transformers.models.bloom.modeling_bloom import build_alibi_tensor
95
95
 
96
- L = hax.Axis("L", 128)
97
- H = hax.Axis("NumHeads", 16)
96
+ L, H = hax.make_axes(L=1, H=16)
98
97
 
99
98
  # Returns tensor shaped (batch_size * num_heads, 1, max_seq_len)
100
99
  torch_tensor = (
@@ -107,7 +106,7 @@ def test_alibi_attention_compared_to_hf():
107
106
 
108
107
 
109
108
  def test_fcm_attention_mask():
110
- KeyPos = hax.Axis("KeyPos", 20)
109
+ KeyPos, QueryPos, Head = hax.make_axes(KeyPos=20, QueryPos=10, Head=8)
111
110
 
112
111
  mask = forgetful_causal_mask(KeyPos, mask_prob=0.6, sample_prob=False, key=PRNGKey(0))
113
112
 
@@ -116,9 +115,6 @@ def test_fcm_attention_mask():
116
115
 
117
116
  assert mask.astype(float).sum().item() <= KeyPos.size
118
117
 
119
- QueryPos = hax.Axis("QueryPos", 10)
120
- Head = hax.Axis("Head", 8)
121
-
122
118
  query = hax.arange(QueryPos).broadcast_axis(Head)
123
119
  key = hax.arange(KeyPos).broadcast_axis(Head)
124
120
 
@@ -1,12 +1,10 @@
1
1
  import pytest
2
2
 
3
- from haliax.axis import Axis, eliminate_axes, rearrange_for_partial_order
3
+ from haliax.axis import Axis, eliminate_axes, make_axes, rearrange_for_partial_order
4
4
 
5
5
 
6
6
  def test_eliminate_axes():
7
- H = Axis("H", 3)
8
- W = Axis("W", 4)
9
- C = Axis("C", 5)
7
+ H, W, C = make_axes(H=3, W=4, C=5)
10
8
 
11
9
  assert eliminate_axes((H, W), (H,)) == (W,)
12
10
  assert eliminate_axes((H, W), (W,)) == (H,)
@@ -273,12 +273,7 @@ def test_einsum_various_errors():
273
273
 
274
274
 
275
275
  def test_einsum_examples():
276
-
277
- Batch = hax.Axis("batch", 32)
278
- Embed = hax.Axis("embed", 64)
279
- H = hax.Axis("h", 16)
280
- W = hax.Axis("w", 16)
281
- C = hax.Axis("c", 3)
276
+ Batch, Embed, H, W, C = hax.make_axes(batch=32, embed=64, h=16, w=16, c=3)
282
277
 
283
278
  # for jax
284
279
  im = jnp.zeros((32, 16, 16, 3))
@@ -46,8 +46,7 @@ def test_dropout():
46
46
 
47
47
 
48
48
  def test_one_hot():
49
- i = Axis("i", 3)
50
- c = Axis("c", 3)
49
+ i, c = hax.make_axes(i=3, c=3)
51
50
  actual = hax.nn.one_hot(hax.NamedArray(jnp.array([0, 1, 2]), (i,)), c)
52
51
  expected = jnp.array([[1.0, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 1.0]])
53
52
 
@@ -61,16 +60,14 @@ def test_one_hot():
61
60
 
62
61
 
63
62
  def test_one_hot_out_of_bound():
64
- i = Axis("i", 2)
65
- c = Axis("c", 3)
63
+ i, c = hax.make_axes(i=2, c=3)
66
64
  actual = hax.nn.one_hot(hax.NamedArray(jnp.array([-1, 3]), (i,)), c)
67
65
  expected = jnp.array([[0.0, 0.0, 0.0], [0.0, 0.0, 0.0]])
68
66
  assert jnp.all(jnp.isclose(actual.array, expected))
69
67
 
70
68
 
71
69
  def test_standardize():
72
- b = Axis("b", 2)
73
- c = Axis("c", 3)
70
+ b, c = hax.make_axes(b=2, c=3)
74
71
  actual = hax.nn.standardize(hax.NamedArray(jnp.array([0, 1, 2]), (c,)), c)
75
72
  expected = jax.nn.standardize(jnp.array([0, 1, 2]), axis=0)
76
73
 
@@ -113,11 +110,7 @@ def test_standardize():
113
110
  @pytest.mark.parametrize("depth", [0, 1, 2, 3, 4, 5])
114
111
  def test_mlp(depth):
115
112
  key = jrandom.PRNGKey(0)
116
- H = Axis("H", 10)
117
- C = Axis("C", 12)
118
- W = Axis("W", 14)
119
-
120
- E = Axis("E", 16)
113
+ H, C, W, E = hax.make_axes(H=10, C=12, W=14, E=16)
121
114
 
122
115
  hax_mlp = hax.nn.MLP.init((H, C, W), E, width=8, depth=depth, key=key)
123
116
  x = hax.random.uniform(key, (H, C, W))
@@ -150,10 +143,7 @@ def test_mlp(depth):
150
143
 
151
144
 
152
145
  def test_linear_has_no_function_leaves_by_default():
153
- H = Axis("H", 10)
154
- C = Axis("C", 12)
155
- W = Axis("W", 14)
156
- E = Axis("E", 16)
146
+ H, C, W, E = hax.make_axes(H=10, C=12, W=14, E=16)
157
147
 
158
148
  hax_linear = hax.nn.Linear.init((H, C, W), E, key=jrandom.PRNGKey(0))
159
149
  assert all(not isinstance(v, Callable) for v in jax.tree_util.tree_leaves(hax_linear)) # type: ignore
@@ -7,11 +7,7 @@ from haliax._src.rearrange import einops_rearrange
7
7
 
8
8
 
9
9
  # some axes
10
- W = Axis("W", 4)
11
- H = Axis("H", 6)
12
- C = Axis("C", 3)
13
- D = Axis("D", 2)
14
- B = Axis("B", 5)
10
+ H, W, C, D, B = hax.make_axes(H=6, W=4, C=3, D=2, B=5)
15
11
  Q = hax.Axis("Q", B.size * H.size)
16
12
  E = Axis("E", H.size * W.size * D.size)
17
13
 
@@ -1,14 +1,13 @@
1
1
  import jax
2
2
  import jax.numpy as jnp
3
3
 
4
+ import haliax
4
5
  import haliax.specialized_fns as hfns
5
- from haliax import Axis, NamedArray
6
+ from haliax import NamedArray
6
7
 
7
8
 
8
9
  def test_top_k():
9
- H = Axis("H", 5)
10
- W = Axis("W", 6)
11
- D = Axis("D", 7)
10
+ H, W, D = haliax.make_axes(H=3, W=4, D=5)
12
11
 
13
12
  rand = jax.random.uniform(jax.random.PRNGKey(0), (H.size, W.size, D.size))
14
13
  n_rand = NamedArray(rand, (H, W, D))
@@ -11,10 +11,7 @@ from haliax import Axis
11
11
 
12
12
 
13
13
  def test_resize_axis():
14
-
15
- A = hax.Axis("A", 10)
16
- B = hax.Axis("B", 20)
17
- C = hax.Axis("C", 30)
14
+ A, B, C = hax.make_axes(A=10, B=20, C=30)
18
15
 
19
16
  class Module(eqx.Module):
20
17
  name1: hax.NamedArray
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev315"
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