haliax 1.4.dev341__tar.gz → 1.4.dev342__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.dev341 → haliax-1.4.dev342}/PKG-INFO +1 -1
- haliax-1.4.dev342/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/_src/state_dict.py +2 -2
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_state_dict.py +25 -1
- haliax-1.4.dev341/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev341 → haliax-1.4.dev342}/.coveragerc +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/.flake8 +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/.gitignore +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/LICENSE +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/README.md +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/api.md +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/css/material.css +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/faq.md +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/fp8.md +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/hof.md +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/index.md +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/indexing.md +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/matmul.md +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/nn.md +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/partitioning.md +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/rearrange.md +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/requirements.txt +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/stacked.md +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/state-dict.md +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/tutorial.md +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/mkdocs.yml +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/pyproject.toml +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/core.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/random.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/types.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/util.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/core_test.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_attention.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_axis.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_conv.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_debug.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_dot.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_hof.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_int8.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_nn.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_ops.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_pool.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_random.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_scan.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_utils.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev342
|
|
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.dev342"
|
|
@@ -58,7 +58,7 @@ def flatten_modules_for_export(t: T) -> T:
|
|
|
58
58
|
def _flatten_module(module):
|
|
59
59
|
if isinstance(module, ModuleWithStateDictSerialization):
|
|
60
60
|
module = module.flatten_for_export()
|
|
61
|
-
module =
|
|
61
|
+
module = scan_aware_tree_map(
|
|
62
62
|
_flatten_module,
|
|
63
63
|
module,
|
|
64
64
|
is_leaf=lambda x: x is not module and isinstance(x, ModuleWithStateDictSerialization),
|
|
@@ -76,7 +76,7 @@ def unflatten_modules_from_export(t: T, template: T) -> T:
|
|
|
76
76
|
def _unflatten_module(module, template):
|
|
77
77
|
if isinstance(module, ModuleWithStateDictSerialization):
|
|
78
78
|
module = module.unflatten_from_export(template)
|
|
79
|
-
module =
|
|
79
|
+
module = scan_aware_tree_map(
|
|
80
80
|
_unflatten_module,
|
|
81
81
|
module,
|
|
82
82
|
template,
|
|
@@ -7,8 +7,9 @@ import jax.numpy as jnp
|
|
|
7
7
|
import pytest
|
|
8
8
|
|
|
9
9
|
import haliax as hax
|
|
10
|
+
from haliax._src.state_dict import flatten_modules_for_export, unflatten_modules_from_export
|
|
10
11
|
from haliax.nn import Linear
|
|
11
|
-
from haliax.nn.scan import _stack_state_dict, _unstack_state_dict
|
|
12
|
+
from haliax.nn.scan import Stacked, _stack_state_dict, _unstack_state_dict
|
|
12
13
|
from haliax.state_dict import from_state_dict, to_state_dict
|
|
13
14
|
|
|
14
15
|
|
|
@@ -151,3 +152,26 @@ def test_export_layer_norm():
|
|
|
151
152
|
new_layer_norm = flat_layer_norm.unflatten_from_export(layer_norm2)
|
|
152
153
|
|
|
153
154
|
assert layer_norm == new_layer_norm
|
|
155
|
+
|
|
156
|
+
|
|
157
|
+
def test_stacked_layer_norm():
|
|
158
|
+
L = hax.Axis("L", 4)
|
|
159
|
+
D = hax.Axis("D", 10)
|
|
160
|
+
E = hax.Axis("E", 20)
|
|
161
|
+
|
|
162
|
+
norms = Stacked.init(L, hax.nn.LayerNorm)((D, E), eps=1e-5, use_weight=True, use_bias=True)
|
|
163
|
+
|
|
164
|
+
norms_flat = flatten_modules_for_export(norms)
|
|
165
|
+
|
|
166
|
+
flat_state_dict = to_state_dict(norms_flat)
|
|
167
|
+
|
|
168
|
+
assert flat_state_dict["0.weight"].shape == (D.size * E.size,)
|
|
169
|
+
assert flat_state_dict["0.bias"].shape == (D.size * E.size,)
|
|
170
|
+
assert flat_state_dict["1.weight"].shape == (D.size * E.size,)
|
|
171
|
+
|
|
172
|
+
# now unflatten it
|
|
173
|
+
norms2 = Stacked.init(L, hax.nn.LayerNorm)((D, E), eps=1e-5, use_weight=True, use_bias=True)
|
|
174
|
+
|
|
175
|
+
new_norms = unflatten_modules_from_export(norms_flat, norms2)
|
|
176
|
+
|
|
177
|
+
assert norms == new_norms
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev341"
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|