haliax 1.4.dev354__tar.gz → 1.4.dev356__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 (107) hide show
  1. {haliax-1.4.dev354 → haliax-1.4.dev356}/PKG-INFO +1 -1
  2. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/indexing.md +39 -0
  3. haliax-1.4.dev356/src/haliax/__about__.py +1 -0
  4. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/__init__.py +2 -0
  5. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/activations.py +10 -0
  6. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_nn.py +52 -0
  7. haliax-1.4.dev354/src/haliax/__about__.py +0 -1
  8. {haliax-1.4.dev354 → haliax-1.4.dev356}/.coveragerc +0 -0
  9. {haliax-1.4.dev354 → haliax-1.4.dev356}/.flake8 +0 -0
  10. {haliax-1.4.dev354 → haliax-1.4.dev356}/.github/workflows/publish_dev.yaml +0 -0
  11. {haliax-1.4.dev354 → haliax-1.4.dev356}/.github/workflows/run_pre_commit.yaml +0 -0
  12. {haliax-1.4.dev354 → haliax-1.4.dev356}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  13. {haliax-1.4.dev354 → haliax-1.4.dev356}/.github/workflows/run_tests.yaml +0 -0
  14. {haliax-1.4.dev354 → haliax-1.4.dev356}/.gitignore +0 -0
  15. {haliax-1.4.dev354 → haliax-1.4.dev356}/.pre-commit-config.yaml +0 -0
  16. {haliax-1.4.dev354 → haliax-1.4.dev356}/.readthedocs.yaml +0 -0
  17. {haliax-1.4.dev354 → haliax-1.4.dev356}/CONTRIBUTING.md +0 -0
  18. {haliax-1.4.dev354 → haliax-1.4.dev356}/LICENSE +0 -0
  19. {haliax-1.4.dev354 → haliax-1.4.dev356}/README.md +0 -0
  20. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/api.md +0 -0
  21. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/broadcasting.md +0 -0
  22. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/cheatsheet.md +0 -0
  23. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/css/material.css +0 -0
  24. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/css/mkdocstrings.css +0 -0
  25. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/faq.md +0 -0
  26. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/data_parallel_mesh.png +0 -0
  27. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  28. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/device_mesh_1d.png +0 -0
  29. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/device_mesh_1d_zero.png +0 -0
  30. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/device_mesh_2d.png +0 -0
  31. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  32. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  33. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  34. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  35. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/device_mesh_2d_zero.png +0 -0
  36. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/fp8.md +0 -0
  37. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/index.md +0 -0
  38. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/matmul.md +0 -0
  39. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/nn.md +0 -0
  40. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/partitioning.md +0 -0
  41. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/rearrange.ipynb +0 -0
  42. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/rearrange.md +0 -0
  43. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/requirements.txt +0 -0
  44. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/scan.md +0 -0
  45. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/state-dict.md +0 -0
  46. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/tutorial.md +0 -0
  47. {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/vmap.md +0 -0
  48. {haliax-1.4.dev354 → haliax-1.4.dev356}/mkdocs.yml +0 -0
  49. {haliax-1.4.dev354 → haliax-1.4.dev356}/pyproject.toml +0 -0
  50. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/__init__.py +0 -0
  51. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/__init__.py +0 -0
  52. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/compile_utils.py +0 -0
  53. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/dot.py +0 -0
  54. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/einsum.py +0 -0
  55. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/fp8.py +0 -0
  56. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/parsing.py +0 -0
  57. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/rearrange.py +0 -0
  58. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/scan.py +0 -0
  59. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/state_dict.py +0 -0
  60. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/util.py +0 -0
  61. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/axis.py +0 -0
  62. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/core.py +0 -0
  63. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/debug.py +0 -0
  64. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/hof.py +0 -0
  65. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/jax_utils.py +0 -0
  66. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/attention.py +0 -0
  67. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/conv.py +0 -0
  68. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/dropout.py +0 -0
  69. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/embedding.py +0 -0
  70. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/linear.py +0 -0
  71. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/loss.py +0 -0
  72. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/mlp.py +0 -0
  73. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/normalization.py +0 -0
  74. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/pool.py +0 -0
  75. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/scan.py +0 -0
  76. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/ops.py +0 -0
  77. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/partitioning.py +0 -0
  78. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/quantization.py +0 -0
  79. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/random.py +0 -0
  80. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/specialized_fns.py +0 -0
  81. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/state_dict.py +0 -0
  82. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/tree_util.py +0 -0
  83. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/types.py +0 -0
  84. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/util.py +0 -0
  85. {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/wrap.py +0 -0
  86. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/core_test.py +0 -0
  87. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_attention.py +0 -0
  88. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_axis.py +0 -0
  89. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_conv.py +0 -0
  90. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_debug.py +0 -0
  91. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_dot.py +0 -0
  92. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_einsum.py +0 -0
  93. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_fp8.py +0 -0
  94. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_hof.py +0 -0
  95. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_int8.py +0 -0
  96. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_ops.py +0 -0
  97. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_parsing.py +0 -0
  98. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_partitioning.py +0 -0
  99. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_pool.py +0 -0
  100. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_random.py +0 -0
  101. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_rearrange.py +0 -0
  102. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_scan.py +0 -0
  103. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_scatter_gather.py +0 -0
  104. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_specialized_fns.py +0 -0
  105. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_state_dict.py +0 -0
  106. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_tree_util.py +0 -0
  107. {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_utils.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev354
3
+ Version: 1.4.dev356
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/
@@ -287,3 +287,42 @@ operation more effectively.)
287
287
 
288
288
  It's worth emphasizing that these functions are typically compiled to scatter-add and friends (as appropriate).
289
289
  This is the preferred way to do scatter/gather operations in JAX, as well as in Haliax.
290
+
291
+ ## Scatter/Gather
292
+
293
+ Haliax supports scatter/gather semantics in its indexing operations. When an axis
294
+ is indexed by another NamedArray (or a 1-D JAX array), the values of that axis
295
+ are gathered according to the index array and the axes of the indexer are
296
+ inserted into the result.
297
+
298
+ ```python
299
+ import haliax as hax
300
+ import jax.numpy as jnp
301
+
302
+ B, S, V = Axis("batch", 4), Axis("seq", 3), Axis("vocab", 7)
303
+ x = hax.arange((B, S, V))
304
+ idx = hax.arange((B, S), dtype=jnp.int32) % V.size
305
+
306
+ out = x["vocab", idx]
307
+ ```
308
+
309
+ Here `out` has axes `(B, S)` and its values match `jax.numpy.take_along_axis`
310
+ on the underlying ndarray.
311
+
312
+ For scatter-style updates where each batch writes to a different position, use
313
+ [`updated_slice`][haliax.updated_slice]:
314
+
315
+ ```python
316
+ Batch = hax.Axis("batch", 2)
317
+ Seq = hax.Axis("seq", 5)
318
+ New = hax.Axis("seq", 2)
319
+
320
+ cache = hax.zeros((Batch, Seq), dtype=int)
321
+ lengths = hax.named([1, 3], axis=Batch)
322
+ kv = hax.named([[1, 2], [3, 4]], axis=(Batch, New))
323
+
324
+ result = updated_slice(cache, {"seq": lengths}, kv)
325
+ ```
326
+
327
+ This inserts `[1, 2]` starting at position `1` in batch `0` and `[3, 4]` starting
328
+ at position `3` in batch `1`.
@@ -0,0 +1 @@
1
+ __version__ = "1.4.dev356"
@@ -23,6 +23,7 @@ from .activations import (
23
23
  quick_gelu,
24
24
  relu,
25
25
  relu6,
26
+ relu_squared,
26
27
  selu,
27
28
  sigmoid,
28
29
  silu,
@@ -94,6 +95,7 @@ __all__ = [
94
95
  "quick_gelu",
95
96
  "glu",
96
97
  "relu6",
98
+ "relu_squared",
97
99
  "sigmoid",
98
100
  "soft_sign",
99
101
  "softplus",
@@ -87,3 +87,13 @@ def glu(x: NamedArray, axis: Axis) -> NamedArray:
87
87
 
88
88
  def quick_gelu(x):
89
89
  return x * sigmoid(1.702 * x)
90
+
91
+
92
+
93
+ def relu_squared(x: A) -> A:
94
+ """ReLU squared activation function. jnp.square(jnp.maximum(0, x))"""
95
+
96
+ def _fn(a):
97
+ return jnp.square(jnn.relu(a))
98
+
99
+ return typing.cast(A, wrap_elemwise_unary(_fn, x))
@@ -147,3 +147,55 @@ def test_linear_has_no_function_leaves_by_default():
147
147
 
148
148
  hax_linear = hax.nn.Linear.init((H, C, W), E, key=jrandom.PRNGKey(0))
149
149
  assert all(not isinstance(v, Callable) for v in jax.tree_util.tree_leaves(hax_linear)) # type: ignore
150
+
151
+
152
+ @pytest.mark.parametrize(
153
+ "input_data, axes",
154
+ [
155
+ (jnp.array([-2.0, -1.0, 0.0, 1.0, 2.0]), (hax.Axis("X", 5),)),
156
+ (jnp.array([[1.0, -1.0], [0.0, 2.0]]), (hax.Axis("Y", 2), hax.Axis("Z", 2))),
157
+ (jnp.array([jnp.nan, 1.0, -1.0]), (hax.Axis("A", 3),)),
158
+ (jnp.array([jnp.inf, -jnp.inf, 0.0]), (hax.Axis("B", 3),)),
159
+ ],
160
+ )
161
+ @pytest.mark.parametrize("dtype", [jnp.float16, jnp.float32, jnp.bfloat16])
162
+ @pytest.mark.parametrize("use_jit", [False, True])
163
+ def test_relu_squared_robust(input_data, axes, dtype, use_jit):
164
+ input_data = input_data.astype(dtype)
165
+ x = hax.named(input_data, axes)
166
+
167
+ # Manually compute the expected output using the base JAX functions
168
+ expected_raw = jnp.square(jax.nn.relu(input_data))
169
+ expected = hax.named(expected_raw, axes)
170
+
171
+ f = hax.nn.relu_squared
172
+ if use_jit:
173
+ f = hax.named_jit(f)
174
+
175
+ # Apply the relu_squared function
176
+ actual = f(x)
177
+
178
+ # Check that the output is a NamedArray with the correct axes and dtype
179
+ assert isinstance(actual, hax.NamedArray)
180
+ assert actual.axes == expected.axes
181
+ assert actual.dtype == expected.dtype
182
+
183
+ # Check that the values are correct, handling NaNs correctly
184
+ assert jnp.allclose(actual.array, expected.array, equal_nan=True)
185
+
186
+
187
+ @pytest.mark.parametrize("use_jit", [False, True])
188
+ def test_relu_squared_scalar(use_jit):
189
+ f = hax.nn.relu_squared
190
+ if use_jit:
191
+ f = jax.jit(f)
192
+
193
+ x = 5.0
194
+ expected = 25.0
195
+ actual = f(x)
196
+ assert jnp.allclose(actual, expected)
197
+
198
+ x_neg = -5.0
199
+ expected_neg = 0.0
200
+ actual_neg = f(x_neg)
201
+ assert jnp.allclose(actual_neg, expected_neg)
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev354"
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