haliax 1.4.dev303__tar.gz → 1.4.dev305__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.dev303 → haliax-1.4.dev305}/PKG-INFO +1 -1
- haliax-1.4.dev305/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/partitioning.py +6 -2
- haliax-1.4.dev303/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev303 → haliax-1.4.dev305}/.coveragerc +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/.flake8 +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/.gitignore +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/LICENSE +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/README.md +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/api.md +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/css/material.css +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/faq.md +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/fp8.md +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/hof.md +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/index.md +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/indexing.md +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/matmul.md +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/nn.md +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/partitioning.md +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/rearrange.md +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/requirements.txt +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/docs/tutorial.md +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/mkdocs.yml +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/pyproject.toml +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/core.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/random.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/types.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/util.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/core_test.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/test_attention.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/test_axis.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/test_conv.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/test_debug.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/test_dot.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/test_hof.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/test_nn.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/test_ops.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/test_pool.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/test_random.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/test_scan.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev303 → haliax-1.4.dev305}/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.dev305
|
|
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.dev305"
|
|
@@ -141,9 +141,13 @@ def shard(x: T, mapping: Optional[ResourceMapping] = None, mesh: Optional[Mesh]
|
|
|
141
141
|
assert isinstance(sharding, NamedSharding)
|
|
142
142
|
if is_in_jit():
|
|
143
143
|
return with_sharding_constraint(x, sharding)
|
|
144
|
-
|
|
144
|
+
elif sharding.is_fully_addressable:
|
|
145
145
|
sharded_array = jax.device_put(x.array, sharding)
|
|
146
146
|
return NamedArray(sharded_array, x.axes)
|
|
147
|
+
else:
|
|
148
|
+
# sharded_array = jax.device_put(x.array, sharding)
|
|
149
|
+
ret = eqx.filter_jit(lambda x: with_sharding_constraint(x, sharding))(x)
|
|
150
|
+
return ret
|
|
147
151
|
|
|
148
152
|
return htu.tree_map(_do_device_put, x)
|
|
149
153
|
|
|
@@ -309,7 +313,7 @@ class _NamedJitWrapper(eqx.Module):
|
|
|
309
313
|
output_shape = _cached_filter_eval_shape(self._fn, *args, **kwargs)
|
|
310
314
|
my_pjit_args = dict(**self._pjit_args)
|
|
311
315
|
|
|
312
|
-
if in_axis_resources is not None
|
|
316
|
+
if in_axis_resources is not None:
|
|
313
317
|
in_resources = infer_resource_partitions(
|
|
314
318
|
(dynamic_donated, dynamic_reserved),
|
|
315
319
|
in_axis_resources,
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev303"
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|