haliax 1.4.dev371__tar.gz → 1.4.dev373__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.dev371 → haliax-1.4.dev373}/PKG-INFO +1 -1
- haliax-1.4.dev373/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/_src/state_dict.py +1 -1
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/partitioning.py +35 -3
- haliax-1.4.dev373/uv.lock +1711 -0
- haliax-1.4.dev371/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev371 → haliax-1.4.dev373}/.coveragerc +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/.flake8 +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/.gitignore +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/AGENTS.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/LICENSE +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/README.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/api.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/css/material.css +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/faq.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/fp8.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/index.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/indexing.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/matmul.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/nn.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/partitioning.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/rearrange.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/requirements.txt +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/scan.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/state-dict.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/tutorial.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/typing.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/docs/vmap.md +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/mkdocs.yml +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/pyproject.toml +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/core.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/haxtyping.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/random.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/types.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/util.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/core_test.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_attention.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_axis.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_conv.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_debug.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_dot.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_hof.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_int8.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_nn.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_ops.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_pool.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_random.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_scan.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev371 → haliax-1.4.dev373}/tests/test_utils.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev373
|
|
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.dev373"
|
|
@@ -197,7 +197,7 @@ def from_state_dict(tree: T, state_dict: StateDict, prefix: Optional[str] = None
|
|
|
197
197
|
if isinstance(array, np.ndarray):
|
|
198
198
|
mesh = partitioning._get_mesh()
|
|
199
199
|
# TODO: modernize this
|
|
200
|
-
if
|
|
200
|
+
if jax.device_count() > 1: # this happens with the default mesh
|
|
201
201
|
pspec = partitioning.pspec_for_axis(tree.axes)
|
|
202
202
|
sharding = jax.sharding.NamedSharding(mesh, pspec)
|
|
203
203
|
array = jax.make_array_from_callback(tree.array.shape, sharding, lambda indices: array[indices])
|
|
@@ -10,7 +10,24 @@ import equinox as eqx
|
|
|
10
10
|
import jax
|
|
11
11
|
from equinox import is_array, module_update_wrapper
|
|
12
12
|
from jax.lax import with_sharding_constraint
|
|
13
|
-
from jax.sharding import
|
|
13
|
+
from jax.sharding import (
|
|
14
|
+
Mesh,
|
|
15
|
+
NamedSharding,
|
|
16
|
+
PartitionSpec,
|
|
17
|
+
SingleDeviceSharding,
|
|
18
|
+
)
|
|
19
|
+
|
|
20
|
+
try: # jax>=0.4.26
|
|
21
|
+
from jax.sharding import AbstractMesh, get_abstract_mesh
|
|
22
|
+
except Exception: # pragma: no cover - older JAX versions
|
|
23
|
+
AbstractMesh = Mesh # type: ignore[misc,assignment]
|
|
24
|
+
def get_abstract_mesh(): # type: ignore[dead-code]
|
|
25
|
+
try:
|
|
26
|
+
from jax.interpreters.pxla import thread_resources
|
|
27
|
+
except Exception:
|
|
28
|
+
from jax.experimental.maps import thread_resources
|
|
29
|
+
|
|
30
|
+
return thread_resources.env.physical_mesh
|
|
14
31
|
from jaxtyping import PyTree
|
|
15
32
|
|
|
16
33
|
import haliax.tree_util as htu
|
|
@@ -604,10 +621,25 @@ def round_axis_for_partitioning(axis: Axis, mapping: Optional[ResourceMapping] =
|
|
|
604
621
|
return Axis(axis.name, new_size)
|
|
605
622
|
|
|
606
623
|
|
|
607
|
-
def _get_mesh() -> Mesh:
|
|
624
|
+
def _get_mesh() -> Mesh | AbstractMesh:
|
|
625
|
+
"""Return the current mesh.
|
|
626
|
+
|
|
627
|
+
On newer versions of JAX this prefers ``get_abstract_mesh`` which does not
|
|
628
|
+
capture concrete devices. If no abstract mesh is currently active we fall
|
|
629
|
+
back to the concrete mesh used by ``Mesh``'s context manager so existing
|
|
630
|
+
code continues to work.
|
|
631
|
+
"""
|
|
632
|
+
|
|
633
|
+
try: # jax>=0.4.26
|
|
634
|
+
mesh = get_abstract_mesh()
|
|
635
|
+
if not getattr(mesh, "empty", False):
|
|
636
|
+
return mesh
|
|
637
|
+
except Exception: # pragma: no cover - older JAX versions
|
|
638
|
+
pass
|
|
639
|
+
|
|
608
640
|
try:
|
|
609
641
|
from jax.interpreters.pxla import thread_resources
|
|
610
|
-
except
|
|
642
|
+
except Exception: # pragma: no cover - jax<0.4
|
|
611
643
|
from jax.experimental.maps import thread_resources
|
|
612
644
|
|
|
613
645
|
return thread_resources.env.physical_mesh
|