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