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