haliax 1.4.dev319__tar.gz → 1.4.dev321__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.dev319 → haliax-1.4.dev321}/PKG-INFO +1 -1
- haliax-1.4.dev321/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/_src/einsum.py +14 -2
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_einsum.py +6 -3
- haliax-1.4.dev319/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev319 → haliax-1.4.dev321}/.coveragerc +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/.flake8 +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/.gitignore +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/LICENSE +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/README.md +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/api.md +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/css/material.css +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/faq.md +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/fp8.md +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/hof.md +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/index.md +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/indexing.md +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/matmul.md +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/nn.md +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/partitioning.md +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/rearrange.md +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/requirements.txt +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/docs/tutorial.md +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/mkdocs.yml +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/pyproject.toml +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/core.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/random.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/types.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/util.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/core_test.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_attention.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_axis.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_conv.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_debug.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_dot.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_hof.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_nn.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_ops.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_pool.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_random.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_scan.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev319 → haliax-1.4.dev321}/tests/test_utils.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.3
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev321
|
|
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.dev321"
|
|
@@ -121,10 +121,13 @@ def _output_only_named_einsum(equation, arrays, rhs, axis_aliases):
|
|
|
121
121
|
used_aliases = set()
|
|
122
122
|
|
|
123
123
|
input_axis_names = set(ax.name for ax in _all_input_axes(arrays))
|
|
124
|
+
has_ellipsis = False
|
|
124
125
|
|
|
125
126
|
for capture in rhs.captures:
|
|
126
127
|
if capture is Ellipsis:
|
|
127
|
-
raise_parse_error("Can't use ellipsis on the rhs of an einsum without an lhs", equation, None)
|
|
128
|
+
# raise_parse_error("Can't use ellipsis on the rhs of an einsum without an lhs", equation, None)
|
|
129
|
+
out_axes.append(Ellipsis)
|
|
130
|
+
has_ellipsis = True
|
|
128
131
|
elif capture.binding is None or len(capture.axes) > 1:
|
|
129
132
|
raise_parse_error(
|
|
130
133
|
"Parenthesized axes are not currently supported in the output of an einsum",
|
|
@@ -150,7 +153,6 @@ def _output_only_named_einsum(equation, arrays, rhs, axis_aliases):
|
|
|
150
153
|
)
|
|
151
154
|
|
|
152
155
|
name = ax_name
|
|
153
|
-
used_axes.add(name)
|
|
154
156
|
|
|
155
157
|
if name in out_axes:
|
|
156
158
|
raise_parse_error(
|
|
@@ -160,8 +162,18 @@ def _output_only_named_einsum(equation, arrays, rhs, axis_aliases):
|
|
|
160
162
|
if name not in input_axis_names:
|
|
161
163
|
raise_parse_error(f"Axis {name} not found in any of the input arrays", equation, capture.char_range)
|
|
162
164
|
|
|
165
|
+
used_axes.add(name)
|
|
163
166
|
out_axes.append(name)
|
|
164
167
|
|
|
168
|
+
# if there's an ellipsis, put all unused axes in the ellipsis
|
|
169
|
+
if has_ellipsis:
|
|
170
|
+
all_input_axes = _all_input_axes(arrays)
|
|
171
|
+
unmentioned = [ax.name for ax in all_input_axes if ax.name not in used_axes]
|
|
172
|
+
ellipsis_index = out_axes.index(Ellipsis)
|
|
173
|
+
out_axes = out_axes[:ellipsis_index] + unmentioned + out_axes[ellipsis_index + 1 :]
|
|
174
|
+
|
|
175
|
+
used_axes = set(out_axes)
|
|
176
|
+
|
|
165
177
|
_check_for_unused_aliases(axis_aliases, used_aliases, equation)
|
|
166
178
|
|
|
167
179
|
spec = _make_einsum_spec(arrays, out_axes)
|
|
@@ -250,9 +250,6 @@ def test_einsum_various_errors():
|
|
|
250
250
|
m1 = hax.ones((Height, Hidth, Depth))
|
|
251
251
|
m2 = hax.ones((Depth, Hidth, Height))
|
|
252
252
|
|
|
253
|
-
with pytest.raises(ValueError, match="Can't use ellipsis"):
|
|
254
|
-
einsum("-> ...", m1, m2)
|
|
255
|
-
|
|
256
253
|
with pytest.raises(ValueError, match="multiple times"):
|
|
257
254
|
einsum("-> Height Height", m1, m2)
|
|
258
255
|
|
|
@@ -330,6 +327,12 @@ def test_einsum_output_only_mode():
|
|
|
330
327
|
assert jnp.all(jnp.equal(einsum("-> Height Width", m1, m2).array, jnp.einsum("ijk,kji->ij", m1.array, m2.array)))
|
|
331
328
|
assert jnp.all(jnp.equal(einsum("-> Height", m1).array, jnp.einsum("ijk->i", m1.array)))
|
|
332
329
|
|
|
330
|
+
assert jnp.all(jnp.equal(einsum("-> ...", m1, m2).array, jnp.einsum("ijk,kji->ijk", m1.array, m2.array)))
|
|
331
|
+
assert jnp.all(jnp.equal(einsum("-> ... Width", m1, m2).array, jnp.einsum("ijk,kji->ikj", m1.array, m2.array)))
|
|
332
|
+
assert jnp.all(
|
|
333
|
+
jnp.equal(einsum("-> Depth ... Width", m1, m2).array, jnp.einsum("ijk,kji->kij", m1.array, m2.array))
|
|
334
|
+
)
|
|
335
|
+
|
|
333
336
|
with pytest.raises(ValueError):
|
|
334
337
|
einsum("-> Q Width", m1)
|
|
335
338
|
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev319"
|
|
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
|