haliax 1.4.dev318__tar.gz → 1.4.dev319__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.dev318 → haliax-1.4.dev319}/PKG-INFO +1 -1
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/api.md +2 -1
- haliax-1.4.dev319/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/axis.py +22 -4
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/core.py +2 -2
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/test_axis.py +15 -1
- haliax-1.4.dev318/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev318 → haliax-1.4.dev319}/.coveragerc +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/.flake8 +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/.gitignore +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/LICENSE +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/README.md +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/css/material.css +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/faq.md +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/fp8.md +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/hof.md +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/index.md +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/indexing.md +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/matmul.md +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/nn.md +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/partitioning.md +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/rearrange.md +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/requirements.txt +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/docs/tutorial.md +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/mkdocs.yml +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/pyproject.toml +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/random.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/types.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/util.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/core_test.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/test_attention.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/test_conv.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/test_debug.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/test_dot.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/test_hof.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/test_nn.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/test_ops.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/test_pool.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/test_random.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/test_scan.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev318 → haliax-1.4.dev319}/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.dev319
|
|
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/
|
|
@@ -31,10 +31,11 @@ Occasionally, an axis size can be inferred in some circumstances but not others.
|
|
|
31
31
|
::: haliax.axis.axis_name
|
|
32
32
|
::: haliax.axis.concat_axes
|
|
33
33
|
::: haliax.axis.union_axes
|
|
34
|
+
::: haliax.axis.intersect_axes
|
|
34
35
|
::: haliax.axis.eliminate_axes
|
|
35
36
|
::: haliax.axis.without_axes
|
|
36
|
-
::: haliax.axis.overlapping_axes
|
|
37
37
|
::: haliax.axis.selects_axis
|
|
38
|
+
::: haliax.axis.overlapping_axes
|
|
38
39
|
::: haliax.axis.is_axis_compatible
|
|
39
40
|
|
|
40
41
|
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "1.4.dev319"
|
|
@@ -297,21 +297,21 @@ def replace_axis(axis_spec: AxisSelection, old: AxisSelector, new: AxisSelection
|
|
|
297
297
|
|
|
298
298
|
|
|
299
299
|
@overload
|
|
300
|
-
def
|
|
300
|
+
def intersect_axes(ax1: AxisSpec, ax2: AxisSelection) -> Tuple[Axis, ...]:
|
|
301
301
|
...
|
|
302
302
|
|
|
303
303
|
|
|
304
304
|
@overload
|
|
305
|
-
def
|
|
305
|
+
def intersect_axes(ax1: AxisSelection, ax2: AxisSpec) -> Tuple[Axis, ...]:
|
|
306
306
|
...
|
|
307
307
|
|
|
308
308
|
|
|
309
309
|
@overload
|
|
310
|
-
def
|
|
310
|
+
def intersect_axes(ax1: AxisSelection, ax2: AxisSelection) -> Tuple[AxisSelector, ...]:
|
|
311
311
|
...
|
|
312
312
|
|
|
313
313
|
|
|
314
|
-
def
|
|
314
|
+
def intersect_axes(ax1: AxisSelection, ax2: AxisSelection) -> Tuple[AxisSelector, ...]:
|
|
315
315
|
"""Returns a tuple of axes that are present in both ax1 and ax2.
|
|
316
316
|
The returned order is the same as ax1.
|
|
317
317
|
"""
|
|
@@ -339,6 +339,23 @@ def overlapping_axes(ax1: AxisSelection, ax2: AxisSelection) -> Tuple[AxisSelect
|
|
|
339
339
|
return tuple(out)
|
|
340
340
|
|
|
341
341
|
|
|
342
|
+
def overlapping_axes(ax1: AxisSelection, ax2: AxisSelection) -> Tuple[str, ...]:
|
|
343
|
+
"""
|
|
344
|
+
Like intersect_axes, but returns the names instead of the axes themselves.
|
|
345
|
+
Unlike intersect_axes, it does not throw an error if the sizes of a common axis are
|
|
346
|
+
different.
|
|
347
|
+
|
|
348
|
+
The returned order is the same as in ax1.
|
|
349
|
+
"""
|
|
350
|
+
ax1 = ensure_tuple(ax1)
|
|
351
|
+
ax2 = ensure_tuple(ax2)
|
|
352
|
+
ax1_names = map(axis_name, ax1)
|
|
353
|
+
ax2_names = set(map(axis_name, ax2))
|
|
354
|
+
|
|
355
|
+
out = tuple(name for name in ax1_names if name in ax2_names)
|
|
356
|
+
return out
|
|
357
|
+
|
|
358
|
+
|
|
342
359
|
@overload
|
|
343
360
|
def axis_name(ax: AxisSelector) -> str: # type: ignore
|
|
344
361
|
...
|
|
@@ -555,6 +572,7 @@ __all__ = [
|
|
|
555
572
|
"dslice",
|
|
556
573
|
"dblock",
|
|
557
574
|
"eliminate_axes",
|
|
575
|
+
"intersect_axes",
|
|
558
576
|
"is_axis_compatible",
|
|
559
577
|
"overlapping_axes",
|
|
560
578
|
"replace_axis",
|
|
@@ -689,9 +689,9 @@ def take(array: NamedArray, axis: AxisSelector, index: Union[int, NamedArray]) -
|
|
|
689
689
|
remaining_axes = eliminate_axes(array.axes, axis)
|
|
690
690
|
# axis order is generally [array.axes[:axis_index], index.axes, array.axes[axis_index + 1 :]]
|
|
691
691
|
# except that index.axes may overlap with array.axes
|
|
692
|
-
|
|
692
|
+
intersecting_axes: AxisSpec = haliax.axis.intersect_axes(remaining_axes, index.axes)
|
|
693
693
|
|
|
694
|
-
if
|
|
694
|
+
if intersecting_axes:
|
|
695
695
|
# if the eliminated axis is also in the index, we rename it to a dummy axis that we can broadcast over it
|
|
696
696
|
need_to_use_dummy_axis = index._lookup_indices(axis.name) is not None
|
|
697
697
|
if need_to_use_dummy_axis:
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
import pytest
|
|
2
2
|
|
|
3
|
-
from haliax.axis import Axis, eliminate_axes, make_axes, rearrange_for_partial_order
|
|
3
|
+
from haliax.axis import Axis, eliminate_axes, make_axes, overlapping_axes, rearrange_for_partial_order
|
|
4
4
|
|
|
5
5
|
|
|
6
6
|
def test_eliminate_axes():
|
|
@@ -133,3 +133,17 @@ def test_duplicate_elements_errors():
|
|
|
133
133
|
|
|
134
134
|
with pytest.raises(ValueError):
|
|
135
135
|
rearrange_for_partial_order(partial_order, candidates)
|
|
136
|
+
|
|
137
|
+
|
|
138
|
+
def test_overlapping_axes_with_different_sizes():
|
|
139
|
+
A1 = Axis("A", 10)
|
|
140
|
+
A2 = Axis("A", 12)
|
|
141
|
+
B = Axis("B", 14)
|
|
142
|
+
C = Axis("C", 16)
|
|
143
|
+
D = Axis("D", 18)
|
|
144
|
+
|
|
145
|
+
ax1 = (A1, B, C)
|
|
146
|
+
ax2 = (A2, C, D)
|
|
147
|
+
|
|
148
|
+
overlapping_names = overlapping_axes(ax1, ax2) # Should not error
|
|
149
|
+
assert overlapping_names == ("A", "C")
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev318"
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|