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.
Files changed (105) hide show
  1. {haliax-1.4.dev341 → haliax-1.4.dev342}/PKG-INFO +1 -1
  2. haliax-1.4.dev342/src/haliax/__about__.py +1 -0
  3. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/_src/state_dict.py +2 -2
  4. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_state_dict.py +25 -1
  5. haliax-1.4.dev341/src/haliax/__about__.py +0 -1
  6. {haliax-1.4.dev341 → haliax-1.4.dev342}/.coveragerc +0 -0
  7. {haliax-1.4.dev341 → haliax-1.4.dev342}/.flake8 +0 -0
  8. {haliax-1.4.dev341 → haliax-1.4.dev342}/.github/workflows/publish_dev.yaml +0 -0
  9. {haliax-1.4.dev341 → haliax-1.4.dev342}/.github/workflows/run_pre_commit.yaml +0 -0
  10. {haliax-1.4.dev341 → haliax-1.4.dev342}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  11. {haliax-1.4.dev341 → haliax-1.4.dev342}/.github/workflows/run_tests.yaml +0 -0
  12. {haliax-1.4.dev341 → haliax-1.4.dev342}/.gitignore +0 -0
  13. {haliax-1.4.dev341 → haliax-1.4.dev342}/.pre-commit-config.yaml +0 -0
  14. {haliax-1.4.dev341 → haliax-1.4.dev342}/.readthedocs.yaml +0 -0
  15. {haliax-1.4.dev341 → haliax-1.4.dev342}/CONTRIBUTING.md +0 -0
  16. {haliax-1.4.dev341 → haliax-1.4.dev342}/LICENSE +0 -0
  17. {haliax-1.4.dev341 → haliax-1.4.dev342}/README.md +0 -0
  18. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/api.md +0 -0
  19. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/broadcasting.md +0 -0
  20. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/cheatsheet.md +0 -0
  21. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/css/material.css +0 -0
  22. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/css/mkdocstrings.css +0 -0
  23. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/faq.md +0 -0
  24. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/data_parallel_mesh.png +0 -0
  25. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  26. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/device_mesh_1d.png +0 -0
  27. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/device_mesh_1d_zero.png +0 -0
  28. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/device_mesh_2d.png +0 -0
  29. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  30. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  31. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  32. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  33. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/figures/device_mesh_2d_zero.png +0 -0
  34. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/fp8.md +0 -0
  35. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/hof.md +0 -0
  36. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/index.md +0 -0
  37. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/indexing.md +0 -0
  38. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/matmul.md +0 -0
  39. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/nn.md +0 -0
  40. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/partitioning.md +0 -0
  41. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/rearrange.ipynb +0 -0
  42. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/rearrange.md +0 -0
  43. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/requirements.txt +0 -0
  44. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/stacked.md +0 -0
  45. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/state-dict.md +0 -0
  46. {haliax-1.4.dev341 → haliax-1.4.dev342}/docs/tutorial.md +0 -0
  47. {haliax-1.4.dev341 → haliax-1.4.dev342}/mkdocs.yml +0 -0
  48. {haliax-1.4.dev341 → haliax-1.4.dev342}/pyproject.toml +0 -0
  49. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/__init__.py +0 -0
  50. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/_src/__init__.py +0 -0
  51. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/_src/compile_utils.py +0 -0
  52. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/_src/dot.py +0 -0
  53. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/_src/einsum.py +0 -0
  54. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/_src/fp8.py +0 -0
  55. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/_src/parsing.py +0 -0
  56. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/_src/rearrange.py +0 -0
  57. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/_src/util.py +0 -0
  58. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/axis.py +0 -0
  59. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/core.py +0 -0
  60. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/debug.py +0 -0
  61. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/hof.py +0 -0
  62. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/jax_utils.py +0 -0
  63. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/__init__.py +0 -0
  64. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/activations.py +0 -0
  65. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/attention.py +0 -0
  66. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/conv.py +0 -0
  67. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/dropout.py +0 -0
  68. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/embedding.py +0 -0
  69. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/linear.py +0 -0
  70. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/loss.py +0 -0
  71. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/mlp.py +0 -0
  72. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/normalization.py +0 -0
  73. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/pool.py +0 -0
  74. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/nn/scan.py +0 -0
  75. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/ops.py +0 -0
  76. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/partitioning.py +0 -0
  77. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/quantization.py +0 -0
  78. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/random.py +0 -0
  79. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/specialized_fns.py +0 -0
  80. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/state_dict.py +0 -0
  81. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/tree_util.py +0 -0
  82. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/types.py +0 -0
  83. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/util.py +0 -0
  84. {haliax-1.4.dev341 → haliax-1.4.dev342}/src/haliax/wrap.py +0 -0
  85. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/core_test.py +0 -0
  86. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_attention.py +0 -0
  87. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_axis.py +0 -0
  88. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_conv.py +0 -0
  89. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_debug.py +0 -0
  90. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_dot.py +0 -0
  91. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_einsum.py +0 -0
  92. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_fp8.py +0 -0
  93. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_hof.py +0 -0
  94. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_int8.py +0 -0
  95. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_nn.py +0 -0
  96. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_ops.py +0 -0
  97. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_parsing.py +0 -0
  98. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_partitioning.py +0 -0
  99. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_pool.py +0 -0
  100. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_random.py +0 -0
  101. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_rearrange.py +0 -0
  102. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_scan.py +0 -0
  103. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_specialized_fns.py +0 -0
  104. {haliax-1.4.dev341 → haliax-1.4.dev342}/tests/test_tree_util.py +0 -0
  105. {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.dev341
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 = jax.tree.map(
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 = jax.tree.map(
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