haliax 1.4.dev325__tar.gz → 1.4.dev327__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.dev325 → haliax-1.4.dev327}/PKG-INFO +3 -2
- haliax-1.4.dev327/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/__init__.py +4 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/_src/state_dict.py +10 -11
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/axis.py +47 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/core.py +90 -1
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/partitioning.py +6 -6
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/quantization.py +1 -2
- haliax-1.4.dev325/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev325 → haliax-1.4.dev327}/.coveragerc +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/.flake8 +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/.gitignore +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/LICENSE +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/README.md +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/api.md +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/css/material.css +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/faq.md +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/fp8.md +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/hof.md +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/index.md +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/indexing.md +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/matmul.md +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/nn.md +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/partitioning.md +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/rearrange.md +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/requirements.txt +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/state-dict.md +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/tutorial.md +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/mkdocs.yml +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/pyproject.toml +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/random.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/types.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/util.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/core_test.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_attention.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_axis.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_conv.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_debug.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_dot.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_hof.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_nn.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_ops.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_pool.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_random.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_scan.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_utils.py +0 -0
|
@@ -1,11 +1,12 @@
|
|
|
1
|
-
Metadata-Version: 2.
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev327
|
|
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/
|
|
7
7
|
Project-URL: Documentation, https://haliax.readthedocs.io/en/latest/
|
|
8
8
|
Author-email: David Hall <dlwh@cs.stanford.edu>
|
|
9
|
+
License-File: LICENSE
|
|
9
10
|
Classifier: Development Status :: 4 - Beta
|
|
10
11
|
Classifier: Intended Audience :: Science/Research
|
|
11
12
|
Classifier: License :: OSI Approved :: Apache Software License
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "1.4.dev327"
|
|
@@ -34,6 +34,8 @@ from .axis import (
|
|
|
34
34
|
dslice,
|
|
35
35
|
eliminate_axes,
|
|
36
36
|
make_axes,
|
|
37
|
+
replace_axis,
|
|
38
|
+
resolve_axis,
|
|
37
39
|
selects_axis,
|
|
38
40
|
)
|
|
39
41
|
from .core import (
|
|
@@ -1060,6 +1062,8 @@ __all__ = [
|
|
|
1060
1062
|
"stack",
|
|
1061
1063
|
"concatenate",
|
|
1062
1064
|
"eliminate_axes",
|
|
1065
|
+
"resolve_axis",
|
|
1066
|
+
"replace_axis",
|
|
1063
1067
|
"selects_axis",
|
|
1064
1068
|
"concat_axes",
|
|
1065
1069
|
"concat_axis_specs",
|
|
@@ -15,7 +15,8 @@ from jaxtyping import PyTree
|
|
|
15
15
|
|
|
16
16
|
import haliax.partitioning as partitioning
|
|
17
17
|
from haliax._src.util import index_where
|
|
18
|
-
from haliax.
|
|
18
|
+
from haliax.axis import Axis
|
|
19
|
+
from haliax.core import NamedArray, flatten_axes, named
|
|
19
20
|
from haliax.jax_utils import is_jax_array_like, is_scalarish
|
|
20
21
|
|
|
21
22
|
|
|
@@ -390,24 +391,22 @@ def flatten_linear_layers(tree: T) -> T:
|
|
|
390
391
|
weight = layer.weight
|
|
391
392
|
bias = layer.bias
|
|
392
393
|
|
|
394
|
+
new_Out: Axis = flatten_axes(layer.Out, "__OUT__")
|
|
395
|
+
new_In: Axis = flatten_axes(layer.In, "__IN__")
|
|
396
|
+
|
|
393
397
|
if weight.array is not None:
|
|
394
398
|
out_first = layer.out_first
|
|
395
|
-
weight = weight.flatten_axes(layer.Out,
|
|
399
|
+
weight = weight.flatten_axes(layer.Out, new_Out).flatten_axes(layer.In, new_In)
|
|
396
400
|
|
|
397
401
|
if out_first:
|
|
398
402
|
weight = weight.rearrange((..., "__OUT__", "__IN__"))
|
|
399
403
|
else:
|
|
400
404
|
weight = weight.rearrange((..., "__IN__", "__OUT__"))
|
|
401
405
|
|
|
402
|
-
|
|
403
|
-
|
|
404
|
-
|
|
405
|
-
In = weight.resolve_axis("__IN__")
|
|
406
|
-
Out = weight.resolve_axis("__OUT__")
|
|
406
|
+
if isinstance(bias, NamedArray):
|
|
407
|
+
bias = bias.flatten_axes(layer.Out, new_Out)
|
|
407
408
|
|
|
408
|
-
|
|
409
|
-
else:
|
|
410
|
-
return layer
|
|
409
|
+
return dataclasses.replace(layer, weight=weight, bias=bias, In=new_In, Out=new_Out) # type: ignore
|
|
411
410
|
|
|
412
411
|
return jax.tree.map(_flatten_linear, tree, is_leaf=lambda x: isinstance(x, Linear))
|
|
413
412
|
|
|
@@ -438,7 +437,7 @@ def unflatten_linear_layers(template: T, tree_with_flattened_linears: T) -> T:
|
|
|
438
437
|
weight = weight.unflatten_axis("__OUT__", template.Out).unflatten_axis("__IN__", template.In)
|
|
439
438
|
weight = weight.rearrange(template.weight.axes)
|
|
440
439
|
|
|
441
|
-
if bias
|
|
440
|
+
if isinstance(bias, NamedArray):
|
|
442
441
|
bias = bias.unflatten_axis("__OUT__", template.Out)
|
|
443
442
|
assert template.bias is not None, "Flattened bias but template has no bias"
|
|
444
443
|
bias = bias.rearrange(template.bias.axes)
|
|
@@ -394,6 +394,52 @@ def axis_size(ax: AxisSpec) -> int:
|
|
|
394
394
|
return prod(axis.size for axis in ensure_tuple(ax)) # type: ignore
|
|
395
395
|
|
|
396
396
|
|
|
397
|
+
@typing.overload
|
|
398
|
+
def resolve_axis(axis_spec: AxisSpec, axis_selection: AxisSelector) -> Axis:
|
|
399
|
+
...
|
|
400
|
+
|
|
401
|
+
|
|
402
|
+
@typing.overload
|
|
403
|
+
def resolve_axis(axis_spec: AxisSpec, axis_selection: AxisSelection) -> AxisSpec:
|
|
404
|
+
...
|
|
405
|
+
|
|
406
|
+
|
|
407
|
+
def resolve_axis(axis_spec: AxisSpec, axis_selection: AxisSelection) -> AxisSpec:
|
|
408
|
+
"""
|
|
409
|
+
Returns the axis or axes in axis_spec that match the name of axis_selection.
|
|
410
|
+
|
|
411
|
+
If axis_selection is a str or axis, returns a single Axis. If it is a sequence, returns a sequence of Axes.
|
|
412
|
+
|
|
413
|
+
If an axis is present with a different size, raises ValueError.
|
|
414
|
+
"""
|
|
415
|
+
ax: Axis
|
|
416
|
+
if isinstance(axis_selection, str | Axis):
|
|
417
|
+
name = axis_name(axis_selection)
|
|
418
|
+
for ax in ensure_tuple(axis_spec):
|
|
419
|
+
if axis_name(ax) == name:
|
|
420
|
+
if isinstance(axis_selection, Axis) and axis_size(ax) != axis_size(axis_selection):
|
|
421
|
+
raise ValueError(f"Axis {name} has different sizes in {axis_spec} and {axis_selection}")
|
|
422
|
+
return ax
|
|
423
|
+
raise ValueError(f"Axis {name} not found in {axis_spec}")
|
|
424
|
+
else:
|
|
425
|
+
as_map = axis_spec_to_shape_dict(axis_spec)
|
|
426
|
+
out: list[Axis] = []
|
|
427
|
+
|
|
428
|
+
for ax in ensure_tuple(axis_selection): # type: ignore
|
|
429
|
+
name = axis_name(ax)
|
|
430
|
+
if name not in as_map:
|
|
431
|
+
raise ValueError(f"Axis {name} not found in {axis_spec}")
|
|
432
|
+
if isinstance(ax, Axis):
|
|
433
|
+
if as_map[name] != ax.size: # type: ignore
|
|
434
|
+
raise ValueError(f"Axis {name} has different sizes in {axis_spec} and {axis_selection}")
|
|
435
|
+
else:
|
|
436
|
+
out.append(ax)
|
|
437
|
+
else:
|
|
438
|
+
out.append(Axis(name, as_map[name]))
|
|
439
|
+
|
|
440
|
+
return tuple(out)
|
|
441
|
+
|
|
442
|
+
|
|
397
443
|
class dslice(eqx.Module):
|
|
398
444
|
"""
|
|
399
445
|
Dynamic slice, comprising a (start, length) pair. Also aliased as ds.
|
|
@@ -576,6 +622,7 @@ __all__ = [
|
|
|
576
622
|
"is_axis_compatible",
|
|
577
623
|
"overlapping_axes",
|
|
578
624
|
"replace_axis",
|
|
625
|
+
"resolve_axis",
|
|
579
626
|
"selects_axis",
|
|
580
627
|
"union_axes",
|
|
581
628
|
"without_axes",
|
|
@@ -1143,12 +1143,65 @@ def rename(array: NamedArray, renames: Mapping[AxisSelector, AxisSelector]) -> N
|
|
|
1143
1143
|
return NamedArray(array.array, new_axes)
|
|
1144
1144
|
|
|
1145
1145
|
|
|
1146
|
+
@typing.overload
|
|
1147
|
+
def flatten_axes(axis: Axis, old_axes: Axis, new_axis: AxisSelector) -> Axis:
|
|
1148
|
+
pass
|
|
1149
|
+
|
|
1150
|
+
|
|
1151
|
+
@typing.overload
|
|
1152
|
+
def flatten_axes(axis: AxisSpec, new_axis: AxisSelector) -> Axis:
|
|
1153
|
+
pass
|
|
1154
|
+
|
|
1155
|
+
|
|
1156
|
+
@typing.overload
|
|
1157
|
+
def flatten_axes(axis: AxisSpec, old_axes: AxisSelection, new_axis: AxisSelector) -> AxisSpec:
|
|
1158
|
+
pass
|
|
1159
|
+
|
|
1160
|
+
|
|
1161
|
+
@typing.overload
|
|
1146
1162
|
def flatten_axes(array: NamedArray, old_axes: AxisSelection, new_axis: AxisSelector) -> NamedArray:
|
|
1163
|
+
pass
|
|
1164
|
+
|
|
1165
|
+
|
|
1166
|
+
def flatten_axes( # type: ignore
|
|
1167
|
+
# array: NamedArray | AxisSpec, old_axes: AxisSelection, new_axis: AxisSelector
|
|
1168
|
+
*args,
|
|
1169
|
+
**kwargs,
|
|
1170
|
+
) -> NamedArray | AxisSpec:
|
|
1147
1171
|
"""
|
|
1148
1172
|
Merge a sequence of axes into a single axis. The new axis must have the same size as the product of the old axes.
|
|
1149
1173
|
|
|
1150
|
-
The new axis is always inserted starting at the index of the first old axis in
|
|
1174
|
+
The new axis is always inserted starting at the index of the first old axis in the underlying array.
|
|
1175
|
+
|
|
1176
|
+
This function can be used in two ways:
|
|
1177
|
+
|
|
1178
|
+
* `flatten_axes(array, old_axes, new_axis)`: merge the old axes of the array into a new axis
|
|
1179
|
+
* `flatten_axes(axes, old_axes, new_axis)`: merge the old axes into a new axis
|
|
1151
1180
|
"""
|
|
1181
|
+
|
|
1182
|
+
if len(args) + len(kwargs) == 2:
|
|
1183
|
+
return _simple_flatten(*args, **kwargs)
|
|
1184
|
+
else:
|
|
1185
|
+
return _full_flatten(*args, **kwargs)
|
|
1186
|
+
|
|
1187
|
+
|
|
1188
|
+
def _simple_flatten(axis: AxisSpec, new_axis: AxisSelector) -> Axis:
|
|
1189
|
+
size = haliax.axis_size(axis)
|
|
1190
|
+
if isinstance(new_axis, Axis):
|
|
1191
|
+
if new_axis.size != size:
|
|
1192
|
+
raise ValueError(f"Cannot merge {axis} into {new_axis}: size mismatch")
|
|
1193
|
+
return new_axis
|
|
1194
|
+
|
|
1195
|
+
assert isinstance(new_axis, str)
|
|
1196
|
+
return Axis(new_axis, size)
|
|
1197
|
+
|
|
1198
|
+
|
|
1199
|
+
def _full_flatten(
|
|
1200
|
+
array: NamedArray | AxisSpec, old_axes: AxisSelection, new_axis: AxisSelector
|
|
1201
|
+
) -> NamedArray | AxisSpec:
|
|
1202
|
+
if isinstance(array, str | Axis | Sequence):
|
|
1203
|
+
return _flatten_axis_spec(array, old_axes, new_axis)
|
|
1204
|
+
|
|
1152
1205
|
old_axes = ensure_tuple(old_axes)
|
|
1153
1206
|
old_axes = array.resolve_axis(old_axes)
|
|
1154
1207
|
total_axis_size = haliax.axis_size(old_axes)
|
|
@@ -1187,6 +1240,42 @@ def flatten_axes(array: NamedArray, old_axes: AxisSelection, new_axis: AxisSelec
|
|
|
1187
1240
|
return NamedArray(raw_array, tuple(new_axes))
|
|
1188
1241
|
|
|
1189
1242
|
|
|
1243
|
+
def _flatten_axis_spec(axes: AxisSpec, old_axes: AxisSelection, new_axis: AxisSelector) -> AxisSpec:
|
|
1244
|
+
axes = ensure_tuple(axes)
|
|
1245
|
+
old_axes = ensure_tuple(old_axes)
|
|
1246
|
+
old_axes = haliax.axis.resolve_axis(axes, old_axes)
|
|
1247
|
+
total_axis_size = haliax.axis_size(old_axes)
|
|
1248
|
+
|
|
1249
|
+
if isinstance(new_axis, Axis):
|
|
1250
|
+
if new_axis.size != total_axis_size:
|
|
1251
|
+
raise ValueError(f"Cannot merge {old_axes} into {new_axis}: size mismatch")
|
|
1252
|
+
else:
|
|
1253
|
+
assert isinstance(new_axis, str)
|
|
1254
|
+
new_axis = Axis(new_axis, total_axis_size)
|
|
1255
|
+
|
|
1256
|
+
if len(old_axes) == 0: # type: ignore
|
|
1257
|
+
return (new_axis,) + axes
|
|
1258
|
+
|
|
1259
|
+
# ensure that the old_axes are contiguous
|
|
1260
|
+
# we basically ensure that the old_axes occur after the index of the first old_axis
|
|
1261
|
+
intermediate_axes: List[Axis] = []
|
|
1262
|
+
new_axes: List[Axis] = []
|
|
1263
|
+
index_of_first_old_axis = None
|
|
1264
|
+
for i, ax in enumerate(axes):
|
|
1265
|
+
if ax in old_axes: # type: ignore
|
|
1266
|
+
if index_of_first_old_axis is None:
|
|
1267
|
+
index_of_first_old_axis = i
|
|
1268
|
+
intermediate_axes.extend(old_axes) # type: ignore
|
|
1269
|
+
new_axes.append(new_axis)
|
|
1270
|
+
else:
|
|
1271
|
+
continue
|
|
1272
|
+
else:
|
|
1273
|
+
intermediate_axes.append(ax)
|
|
1274
|
+
new_axes.append(ax)
|
|
1275
|
+
|
|
1276
|
+
return tuple(new_axes)
|
|
1277
|
+
|
|
1278
|
+
|
|
1190
1279
|
def unflatten_axis(array: NamedArray, axis: AxisSelector, new_axes: AxisSpec) -> NamedArray:
|
|
1191
1280
|
"""
|
|
1192
1281
|
Split an axis into a sequence of axes. The old axis must have the same size as the product of the new axes.
|
|
@@ -8,7 +8,7 @@ from typing import Callable, ContextManager, Mapping, Optional, ParamSpec, Seque
|
|
|
8
8
|
|
|
9
9
|
import equinox as eqx
|
|
10
10
|
import jax
|
|
11
|
-
from equinox import module_update_wrapper
|
|
11
|
+
from equinox import is_array, module_update_wrapper
|
|
12
12
|
from jax.lax import with_sharding_constraint
|
|
13
13
|
from jax.sharding import Mesh, NamedSharding, PartitionSpec, SingleDeviceSharding
|
|
14
14
|
from jaxtyping import PyTree
|
|
@@ -20,7 +20,7 @@ from .axis import Axis, AxisSelection, AxisSelector
|
|
|
20
20
|
from .core import NamedArray
|
|
21
21
|
from .jax_utils import Static, is_in_jit, is_jax_array_like, is_on_mac_metal
|
|
22
22
|
from .tree_util import hashable_combine, hashable_partition
|
|
23
|
-
from .util import StringHolderEnum, ensure_tuple
|
|
23
|
+
from .util import StringHolderEnum, ensure_tuple
|
|
24
24
|
|
|
25
25
|
|
|
26
26
|
PhysicalAxisSpec = Union[(str), Sequence[str]]
|
|
@@ -274,7 +274,7 @@ class _NamedJitWrapper(eqx.Module):
|
|
|
274
274
|
if out_axis_resources is None:
|
|
275
275
|
out_axis_resources = axis_resources
|
|
276
276
|
|
|
277
|
-
dynamic_argspec, static_argspec = hashable_partition((args, kwargs),
|
|
277
|
+
dynamic_argspec, static_argspec = hashable_partition((args, kwargs), is_array)
|
|
278
278
|
dynamic = (self._dynamic_fun, dynamic_argspec)
|
|
279
279
|
|
|
280
280
|
donate_args = self._donate_args
|
|
@@ -436,7 +436,7 @@ def named_jit(
|
|
|
436
436
|
**pjit_args,
|
|
437
437
|
)
|
|
438
438
|
|
|
439
|
-
dynamic_fun, static_fun = hashable_partition(fn,
|
|
439
|
+
dynamic_fun, static_fun = hashable_partition(fn, is_array)
|
|
440
440
|
|
|
441
441
|
wrapper = _NamedJitWrapper(
|
|
442
442
|
fn,
|
|
@@ -514,7 +514,7 @@ def _named_pjit_cache(fun_names, **jitkwargs) -> WrappedCallable:
|
|
|
514
514
|
fun = hashable_combine(dynamic_fun, static_fun)
|
|
515
515
|
args, kwargs = hashable_combine(dynamic_spec, static_spec)
|
|
516
516
|
out = fun(*args, **kwargs)
|
|
517
|
-
out_dynamic, out_static = hashable_partition(out,
|
|
517
|
+
out_dynamic, out_static = hashable_partition(out, is_array)
|
|
518
518
|
return out_dynamic, Static(out_static)
|
|
519
519
|
|
|
520
520
|
fun_name, fun_qualname = fun_names
|
|
@@ -543,7 +543,7 @@ def _cached_filter_eval_shape(fun, *args, **kwargs):
|
|
|
543
543
|
eval_shape is surprisingly expensive, so we cache it. We use this for named_pjit for evaluating resource partitions
|
|
544
544
|
of the output.
|
|
545
545
|
"""
|
|
546
|
-
dynamic, static = hashable_partition((fun, args, kwargs),
|
|
546
|
+
dynamic, static = hashable_partition((fun, args, kwargs), is_array)
|
|
547
547
|
if static not in _eval_shape_cache:
|
|
548
548
|
_eval_shape_cache[static] = eqx.filter_eval_shape(fun, *args, **kwargs)
|
|
549
549
|
|
|
@@ -10,7 +10,6 @@ from typing import Optional, Protocol, TypeVar
|
|
|
10
10
|
import equinox as eqx
|
|
11
11
|
import jax
|
|
12
12
|
from jax import numpy as jnp
|
|
13
|
-
from jax._src.tree_util import BuiltInKeyEntry
|
|
14
13
|
from jax.tree_util import DictKey, FlattenedIndexKey, GetAttrKey, SequenceKey
|
|
15
14
|
from jax.typing import DTypeLike
|
|
16
15
|
|
|
@@ -253,7 +252,7 @@ def _matches_target_fp8(key_path, config: Fp8Config) -> bool:
|
|
|
253
252
|
return re.match(config.targets, key_path_str) is not None
|
|
254
253
|
|
|
255
254
|
|
|
256
|
-
def _key_path_to_str(key_path: tuple
|
|
255
|
+
def _key_path_to_str(key_path: tuple) -> str:
|
|
257
256
|
out = ""
|
|
258
257
|
for k in key_path:
|
|
259
258
|
match k:
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev325"
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|