haliax 1.4.dev402__tar.gz → 1.4.dev403__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.dev402 → haliax-1.4.dev403}/PKG-INFO +1 -1
- haliax-1.4.dev403/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/partitioning.py +33 -1
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_partitioning.py +28 -0
- haliax-1.4.dev402/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev402 → haliax-1.4.dev403}/.coveragerc +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/.flake8 +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/.gitignore +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/.playbooks/wrap-non-named.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/AGENTS.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/LICENSE +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/README.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/api.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/css/material.css +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/faq.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/fp8.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/index.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/indexing.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/matmul.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/nn.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/partitioning.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/primer.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/rearrange.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/requirements.txt +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/scan.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/state-dict.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/tutorial.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/typing.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/vmap.md +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/mkdocs.yml +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/pyproject.toml +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/core.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/field.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/haxtyping.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/random.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/types.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/util.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/core_test.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_attention.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_axis.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_conv.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_debug.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_dot.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_field.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_hof.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_int8.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_moe_linear.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_nn.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_ops.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_pool.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_random.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_scan.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_utils.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_visualize_sharding.py +0 -0
- {haliax-1.4.dev402 → haliax-1.4.dev403}/uv.lock +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev403
|
|
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.dev403"
|
|
@@ -1,4 +1,5 @@
|
|
|
1
1
|
import contextlib
|
|
2
|
+
import dataclasses
|
|
2
3
|
import functools
|
|
3
4
|
import threading
|
|
4
5
|
import typing
|
|
@@ -210,6 +211,37 @@ def pspec_for(
|
|
|
210
211
|
return None
|
|
211
212
|
else:
|
|
212
213
|
return pspec_for_axis(node.axes, resource_mapping)
|
|
214
|
+
elif isinstance(node, eqx.Module):
|
|
215
|
+
# handle eqx.Module explicitly so that we can look at axis_names metadata
|
|
216
|
+
updates: dict[str, typing.Any] = {}
|
|
217
|
+
for field in dataclasses.fields(node):
|
|
218
|
+
if field.metadata.get("static", False):
|
|
219
|
+
continue
|
|
220
|
+
|
|
221
|
+
value = getattr(node, field.name)
|
|
222
|
+
axis_names = field.metadata.get("axis_names") if field.metadata is not None else None
|
|
223
|
+
if axis_names is not None and is_jax_array_like(value):
|
|
224
|
+
current_sharding = (
|
|
225
|
+
getattr(value, "sharding", None) if preserve_existing_shardings else None
|
|
226
|
+
)
|
|
227
|
+
if current_sharding is not None:
|
|
228
|
+
updates[field.name] = None
|
|
229
|
+
else:
|
|
230
|
+
updates[field.name] = pspec_for_axis(axis_names, resource_mapping)
|
|
231
|
+
else:
|
|
232
|
+
updates[field.name] = htu.tree_map(
|
|
233
|
+
partition_spec, value, is_leaf=lambda x: isinstance(x, eqx.Module)
|
|
234
|
+
)
|
|
235
|
+
|
|
236
|
+
new_node = object.__new__(type(node))
|
|
237
|
+
for field in dataclasses.fields(node):
|
|
238
|
+
object.__setattr__(
|
|
239
|
+
new_node,
|
|
240
|
+
field.name,
|
|
241
|
+
updates.get(field.name, getattr(node, field.name)),
|
|
242
|
+
)
|
|
243
|
+
|
|
244
|
+
return new_node
|
|
213
245
|
elif is_jax_array_like(node):
|
|
214
246
|
sharding = getattr(node, "sharding", None)
|
|
215
247
|
# TODO: these are usually replicated. Is there a better way to tell?
|
|
@@ -232,7 +264,7 @@ def pspec_for(
|
|
|
232
264
|
else:
|
|
233
265
|
return None
|
|
234
266
|
|
|
235
|
-
return htu.tree_map(partition_spec, tree)
|
|
267
|
+
return htu.tree_map(partition_spec, tree, is_leaf=lambda x: isinstance(x, eqx.Module))
|
|
236
268
|
|
|
237
269
|
|
|
238
270
|
def infer_resource_partitions(
|
|
@@ -59,6 +59,34 @@ def test_pspec_for_named_axes():
|
|
|
59
59
|
assert specs.unnamed1 == PartitionSpec(None)
|
|
60
60
|
|
|
61
61
|
|
|
62
|
+
class ArrayModule(eqx.Module):
|
|
63
|
+
arr: Array = hax.field(axis_names=("dim2", "dim3"))
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def test_pspec_for_plain_array_axis_names():
|
|
67
|
+
mesh = Mesh(np.array(jax.devices()).reshape(-1, 1), (ResourceAxis.DATA, ResourceAxis.MODEL))
|
|
68
|
+
with axis_mapping(resource_map), mesh:
|
|
69
|
+
mod = ArrayModule(jnp.ones((Dim2.size, Dim3.size)))
|
|
70
|
+
|
|
71
|
+
specs: ArrayModule = pspec_for(mod, preserve_existing_shardings=False)
|
|
72
|
+
|
|
73
|
+
assert specs.arr == PartitionSpec(ResourceAxis.DATA, ResourceAxis.MODEL)
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
class NestedArrayModule(eqx.Module):
|
|
77
|
+
inner: ArrayModule
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
def test_pspec_for_plain_array_axis_names_nested_module():
|
|
81
|
+
mesh = Mesh(np.array(jax.devices()).reshape(-1, 1), (ResourceAxis.DATA, ResourceAxis.MODEL))
|
|
82
|
+
with axis_mapping(resource_map), mesh:
|
|
83
|
+
mod = NestedArrayModule(ArrayModule(jnp.ones((Dim2.size, Dim3.size))))
|
|
84
|
+
|
|
85
|
+
specs: NestedArrayModule = pspec_for(mod, preserve_existing_shardings=False)
|
|
86
|
+
|
|
87
|
+
assert specs.inner.arr == PartitionSpec(ResourceAxis.DATA, ResourceAxis.MODEL)
|
|
88
|
+
|
|
89
|
+
|
|
62
90
|
class MyModuleInit(eqx.Module):
|
|
63
91
|
named: NamedArray
|
|
64
92
|
unnamed1: Array
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev402"
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|