haliax 1.4.dev324__tar.gz → 1.4.dev325__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.dev324 → haliax-1.4.dev325}/PKG-INFO +1 -1
- haliax-1.4.dev325/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/_src/state_dict.py +2 -1
- haliax-1.4.dev324/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev324 → haliax-1.4.dev325}/.coveragerc +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/.flake8 +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/.gitignore +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/LICENSE +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/README.md +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/api.md +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/css/material.css +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/faq.md +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/fp8.md +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/hof.md +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/index.md +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/indexing.md +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/matmul.md +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/nn.md +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/partitioning.md +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/rearrange.md +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/requirements.txt +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/state-dict.md +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/docs/tutorial.md +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/mkdocs.yml +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/pyproject.toml +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/core.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/random.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/types.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/util.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/core_test.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_attention.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_axis.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_conv.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_debug.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_dot.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_hof.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_nn.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_ops.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_pool.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_random.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_scan.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev324 → haliax-1.4.dev325}/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.dev325
|
|
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.dev325"
|
|
@@ -320,6 +320,7 @@ def to_numpy_state_dict(model, prefix: Optional[str] = None) -> StateDict:
|
|
|
320
320
|
process_mesh = Mesh(
|
|
321
321
|
np.array(jax.devices()).reshape((jax.process_count(), -1)), ("process", "device")
|
|
322
322
|
)
|
|
323
|
+
|
|
323
324
|
# now we need to find an axis along which we can shard the array.
|
|
324
325
|
# for this, we need to find an axis s.t. size(axis) % local_devices == 0
|
|
325
326
|
|
|
@@ -332,7 +333,7 @@ def to_numpy_state_dict(model, prefix: Optional[str] = None) -> StateDict:
|
|
|
332
333
|
|
|
333
334
|
shardings = [None if i != axis_to_shard else "device" for i in range(len(arr.shape))]
|
|
334
335
|
sharding = NamedSharding(process_mesh, PartitionSpec(*shardings))
|
|
335
|
-
out = jax.
|
|
336
|
+
out = jax.device_put(arr, sharding)
|
|
336
337
|
return np.array(out)
|
|
337
338
|
elif is_scalarish(arr):
|
|
338
339
|
return np.asarray(arr)
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev324"
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|