haliax 1.4.dev305__tar.gz → 1.4.dev306__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.dev305 → haliax-1.4.dev306}/.github/workflows/run_tests.yaml +1 -1
- {haliax-1.4.dev305 → haliax-1.4.dev306}/PKG-INFO +1 -1
- haliax-1.4.dev306/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/partitioning.py +13 -12
- haliax-1.4.dev305/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev305 → haliax-1.4.dev306}/.coveragerc +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/.flake8 +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/.gitignore +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/LICENSE +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/README.md +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/api.md +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/css/material.css +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/faq.md +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/fp8.md +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/hof.md +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/index.md +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/indexing.md +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/matmul.md +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/nn.md +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/partitioning.md +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/rearrange.md +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/requirements.txt +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/docs/tutorial.md +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/mkdocs.yml +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/pyproject.toml +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/core.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/random.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/types.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/util.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/core_test.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_attention.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_axis.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_conv.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_debug.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_dot.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_hof.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_nn.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_ops.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_pool.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_random.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_scan.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev305 → haliax-1.4.dev306}/tests/test_utils.py +0 -0
|
@@ -17,7 +17,7 @@ jobs:
|
|
|
17
17
|
run: |
|
|
18
18
|
python -m pip install --upgrade pip
|
|
19
19
|
pip install flake8 pytest
|
|
20
|
-
pip install --upgrade "jax[cpu]==0.4.
|
|
20
|
+
pip install --upgrade "jax[cpu]==0.4.30" "jaxlib[cpu]==0.4.30"
|
|
21
21
|
pip install .[dev]
|
|
22
22
|
- name: Test with pytest
|
|
23
23
|
run: |
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.3
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev306
|
|
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.dev306"
|
|
@@ -128,25 +128,26 @@ def shard(x: T, mapping: Optional[ResourceMapping] = None, mesh: Optional[Mesh]
|
|
|
128
128
|
warnings.warn("Sharding constraints are not supported in jit on metal", RuntimeWarning)
|
|
129
129
|
return x
|
|
130
130
|
|
|
131
|
-
def _do_device_put(
|
|
132
|
-
if not isinstance(
|
|
133
|
-
return
|
|
131
|
+
def _do_device_put(named):
|
|
132
|
+
if not isinstance(named, NamedArray):
|
|
133
|
+
return named
|
|
134
134
|
|
|
135
|
-
if not is_jax_array_like(
|
|
135
|
+
if not is_jax_array_like(named.array):
|
|
136
136
|
# this happens when we filter out params for things like lora.
|
|
137
137
|
# could use eqx.partition to avoid this, but eh
|
|
138
|
-
return
|
|
138
|
+
return named
|
|
139
139
|
|
|
140
|
-
sharding = infer_resource_partitions(
|
|
140
|
+
sharding = infer_resource_partitions(named, mapping, mesh=mesh, preserve_existing_shardings=False)
|
|
141
141
|
assert isinstance(sharding, NamedSharding)
|
|
142
142
|
if is_in_jit():
|
|
143
|
-
return with_sharding_constraint(
|
|
144
|
-
|
|
145
|
-
|
|
146
|
-
|
|
143
|
+
return with_sharding_constraint(named, sharding)
|
|
144
|
+
# as a special case, SingleDeviceShardings are routed through jit
|
|
145
|
+
elif isinstance(named.array.sharding, SingleDeviceSharding):
|
|
146
|
+
# TODO(dlwh): this should be unnecessary in JAX soon. Check after 2024-08-01
|
|
147
|
+
sharded_array = jax.jit(lambda x: x, out_shardings=sharding)(named)
|
|
148
|
+
return sharded_array
|
|
147
149
|
else:
|
|
148
|
-
|
|
149
|
-
ret = eqx.filter_jit(lambda x: with_sharding_constraint(x, sharding))(x)
|
|
150
|
+
ret = jax.device_put(named, sharding)
|
|
150
151
|
return ret
|
|
151
152
|
|
|
152
153
|
return htu.tree_map(_do_device_put, x)
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev305"
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|