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.
- {haliax-1.4.dev354 → haliax-1.4.dev356}/PKG-INFO +1 -1
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/indexing.md +39 -0
- haliax-1.4.dev356/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/__init__.py +2 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/activations.py +10 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_nn.py +52 -0
- haliax-1.4.dev354/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev354 → haliax-1.4.dev356}/.coveragerc +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/.flake8 +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/.gitignore +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/LICENSE +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/README.md +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/api.md +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/css/material.css +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/faq.md +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/fp8.md +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/index.md +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/matmul.md +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/nn.md +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/partitioning.md +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/rearrange.md +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/requirements.txt +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/scan.md +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/state-dict.md +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/tutorial.md +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/docs/vmap.md +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/mkdocs.yml +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/pyproject.toml +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/core.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/random.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/types.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/util.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/core_test.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_attention.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_axis.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_conv.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_debug.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_dot.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_hof.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_int8.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_ops.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_pool.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_random.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_scan.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev354 → haliax-1.4.dev356}/tests/test_tree_util.py +0 -0
- {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.
|
|
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
|
|
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
|