haliax 1.4.dev306__tar.gz → 1.4.dev307__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.dev306 → haliax-1.4.dev307}/PKG-INFO +1 -1
- haliax-1.4.dev307/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/partitioning.py +2 -1
- haliax-1.4.dev306/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev306 → haliax-1.4.dev307}/.coveragerc +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/.flake8 +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/.gitignore +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/LICENSE +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/README.md +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/api.md +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/css/material.css +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/faq.md +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/fp8.md +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/hof.md +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/index.md +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/indexing.md +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/matmul.md +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/nn.md +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/partitioning.md +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/rearrange.md +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/requirements.txt +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/docs/tutorial.md +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/mkdocs.yml +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/pyproject.toml +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/core.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/random.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/types.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/util.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/core_test.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/test_attention.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/test_axis.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/test_conv.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/test_debug.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/test_dot.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/test_hof.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/test_nn.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/test_ops.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/test_pool.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/test_random.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/test_scan.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev306 → haliax-1.4.dev307}/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.dev307
|
|
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.dev307"
|
|
@@ -139,10 +139,11 @@ def shard(x: T, mapping: Optional[ResourceMapping] = None, mesh: Optional[Mesh]
|
|
|
139
139
|
|
|
140
140
|
sharding = infer_resource_partitions(named, mapping, mesh=mesh, preserve_existing_shardings=False)
|
|
141
141
|
assert isinstance(sharding, NamedSharding)
|
|
142
|
+
in_sharding = getattr(named.array, "sharding", None)
|
|
142
143
|
if is_in_jit():
|
|
143
144
|
return with_sharding_constraint(named, sharding)
|
|
144
145
|
# as a special case, SingleDeviceShardings are routed through jit
|
|
145
|
-
elif isinstance(
|
|
146
|
+
elif isinstance(in_sharding, SingleDeviceSharding) and in_sharding._device in sharding.device_set:
|
|
146
147
|
# TODO(dlwh): this should be unnecessary in JAX soon. Check after 2024-08-01
|
|
147
148
|
sharded_array = jax.jit(lambda x: x, out_shardings=sharding)(named)
|
|
148
149
|
return sharded_array
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev306"
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|