haliax 1.4.dev292__tar.gz → 1.4.dev294__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.dev292 → haliax-1.4.dev294}/PKG-INFO +1 -1
- haliax-1.4.dev294/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/partitioning.py +13 -1
- haliax-1.4.dev292/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev292 → haliax-1.4.dev294}/.coveragerc +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/.flake8 +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/.gitignore +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/LICENSE +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/README.md +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/api.md +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/css/material.css +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/faq.md +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/fp8.md +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/hof.md +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/index.md +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/indexing.md +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/matmul.md +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/nn.md +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/partitioning.md +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/rearrange.md +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/requirements.txt +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/docs/tutorial.md +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/mkdocs.yml +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/pyproject.toml +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/core.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/random.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/types.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/util.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/core_test.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/test_attention.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/test_axis.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/test_conv.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/test_debug.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/test_dot.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/test_hof.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/test_nn.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/test_ops.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/test_pool.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/test_random.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/test_scan.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev292 → haliax-1.4.dev294}/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.dev294
|
|
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.dev294"
|
|
@@ -3,6 +3,7 @@ import functools
|
|
|
3
3
|
import threading
|
|
4
4
|
import typing
|
|
5
5
|
import warnings
|
|
6
|
+
from itertools import chain
|
|
6
7
|
from math import prod
|
|
7
8
|
from typing import Callable, ContextManager, Mapping, Optional, ParamSpec, Sequence, TypeVar, Union
|
|
8
9
|
|
|
@@ -38,6 +39,7 @@ class ResourceAxis(StringHolderEnum):
|
|
|
38
39
|
|
|
39
40
|
MODEL = "model"
|
|
40
41
|
DATA = "data"
|
|
42
|
+
REPLICA = "replica"
|
|
41
43
|
|
|
42
44
|
|
|
43
45
|
class _ResourceMappingHolder:
|
|
@@ -584,7 +586,17 @@ def sharding_for_axis(
|
|
|
584
586
|
def pspec_for_axis(axis: AxisSelection, mapping: Optional[ResourceMapping] = None) -> PartitionSpec:
|
|
585
587
|
"""Get the PartitionSpec for a single axis"""
|
|
586
588
|
axis = ensure_tuple(axis)
|
|
587
|
-
|
|
589
|
+
phys_axes = []
|
|
590
|
+
for a in axis:
|
|
591
|
+
pa = physical_axis_name(a, mapping)
|
|
592
|
+
if pa is None or isinstance(pa, str):
|
|
593
|
+
phys_axes.append(pa)
|
|
594
|
+
else:
|
|
595
|
+
# I have no way to resolve the mypy check :)
|
|
596
|
+
for i in pa:
|
|
597
|
+
phys_axes.append(i)
|
|
598
|
+
|
|
599
|
+
return PartitionSpec(*phys_axes)
|
|
588
600
|
|
|
589
601
|
|
|
590
602
|
def round_axis_for_partitioning(axis: Axis, mapping: Optional[ResourceMapping] = None) -> Axis:
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev292"
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|