haliax 1.4.dev447__tar.gz → 1.4.dev452__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 (138) hide show
  1. {haliax-1.4.dev447 → haliax-1.4.dev452}/PKG-INFO +1 -1
  2. {haliax-1.4.dev447 → haliax-1.4.dev452}/pyproject.toml +1 -0
  3. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/__about__.py +1 -1
  4. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/_src/state_dict.py +8 -2
  5. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/nn/linear.py +7 -2
  6. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_state_dict.py +40 -1
  7. {haliax-1.4.dev447 → haliax-1.4.dev452}/uv.lock +2 -0
  8. {haliax-1.4.dev447 → haliax-1.4.dev452}/.agents/projects/api_parity.md +0 -0
  9. {haliax-1.4.dev447 → haliax-1.4.dev452}/.agents/refs.md +0 -0
  10. {haliax-1.4.dev447 → haliax-1.4.dev452}/.coveragerc +0 -0
  11. {haliax-1.4.dev447 → haliax-1.4.dev452}/.flake8 +0 -0
  12. {haliax-1.4.dev447 → haliax-1.4.dev452}/.github/workflows/publish_dev.yaml +0 -0
  13. {haliax-1.4.dev447 → haliax-1.4.dev452}/.github/workflows/run_pre_commit.yaml +0 -0
  14. {haliax-1.4.dev447 → haliax-1.4.dev452}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  15. {haliax-1.4.dev447 → haliax-1.4.dev452}/.github/workflows/run_tests.yaml +0 -0
  16. {haliax-1.4.dev447 → haliax-1.4.dev452}/.gitignore +0 -0
  17. {haliax-1.4.dev447 → haliax-1.4.dev452}/.playbooks/add-types.md +0 -0
  18. {haliax-1.4.dev447 → haliax-1.4.dev452}/.playbooks/wrap-non-named.md +0 -0
  19. {haliax-1.4.dev447 → haliax-1.4.dev452}/.pre-commit-config.yaml +0 -0
  20. {haliax-1.4.dev447 → haliax-1.4.dev452}/.readthedocs.yaml +0 -0
  21. {haliax-1.4.dev447 → haliax-1.4.dev452}/AGENTS.md +0 -0
  22. {haliax-1.4.dev447 → haliax-1.4.dev452}/AUTHORS.md +0 -0
  23. {haliax-1.4.dev447 → haliax-1.4.dev452}/CONTRIBUTING.md +0 -0
  24. {haliax-1.4.dev447 → haliax-1.4.dev452}/CONTRIBUTORS.md +0 -0
  25. {haliax-1.4.dev447 → haliax-1.4.dev452}/LICENSE +0 -0
  26. {haliax-1.4.dev447 → haliax-1.4.dev452}/README.md +0 -0
  27. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/api.md +0 -0
  28. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/broadcasting.md +0 -0
  29. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/cheatsheet.md +0 -0
  30. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/css/material.css +0 -0
  31. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/css/mkdocstrings.css +0 -0
  32. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/faq.md +0 -0
  33. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/figures/data_parallel_mesh.png +0 -0
  34. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  35. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/figures/device_mesh_1d.png +0 -0
  36. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/figures/device_mesh_1d_zero.png +0 -0
  37. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/figures/device_mesh_2d.png +0 -0
  38. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  39. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  40. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  41. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  42. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/figures/device_mesh_2d_zero.png +0 -0
  43. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/fp8.md +0 -0
  44. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/index.md +0 -0
  45. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/indexing.md +0 -0
  46. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/matmul.md +0 -0
  47. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/mutable-refs.md +0 -0
  48. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/nn.md +0 -0
  49. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/partitioning.md +0 -0
  50. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/primer.md +0 -0
  51. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/rearrange.ipynb +0 -0
  52. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/rearrange.md +0 -0
  53. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/requirements.txt +0 -0
  54. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/scan.md +0 -0
  55. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/state-dict.md +0 -0
  56. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/tutorial.md +0 -0
  57. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/typing.md +0 -0
  58. {haliax-1.4.dev447 → haliax-1.4.dev452}/docs/vmap.md +0 -0
  59. {haliax-1.4.dev447 → haliax-1.4.dev452}/etc/license_header.txt +0 -0
  60. {haliax-1.4.dev447 → haliax-1.4.dev452}/mkdocs.yml +0 -0
  61. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/__init__.py +0 -0
  62. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/_src/__init__.py +0 -0
  63. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/_src/compile_utils.py +0 -0
  64. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/_src/dot.py +0 -0
  65. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/_src/einsum.py +0 -0
  66. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/_src/fp8.py +0 -0
  67. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/_src/parsing.py +0 -0
  68. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/_src/rearrange.py +0 -0
  69. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/_src/scan.py +0 -0
  70. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/_src/util.py +0 -0
  71. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/axis.py +0 -0
  72. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/core.py +0 -0
  73. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/debug.py +0 -0
  74. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/fft.py +0 -0
  75. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/field.py +0 -0
  76. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/haxtyping.py +0 -0
  77. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/hof.py +0 -0
  78. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/jax_utils.py +0 -0
  79. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/nn/__init__.py +0 -0
  80. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/nn/activations.py +0 -0
  81. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/nn/attention.py +0 -0
  82. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/nn/conv.py +0 -0
  83. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/nn/dropout.py +0 -0
  84. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/nn/embedding.py +0 -0
  85. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/nn/loss.py +0 -0
  86. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/nn/mlp.py +0 -0
  87. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/nn/mup.py +0 -0
  88. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/nn/normalization.py +0 -0
  89. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/nn/pool.py +0 -0
  90. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/nn/scan.py +0 -0
  91. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/ops.py +0 -0
  92. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/partitioning.py +0 -0
  93. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/poly.py +0 -0
  94. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/quantization.py +0 -0
  95. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/random.py +0 -0
  96. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/ref.py +0 -0
  97. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/specialized_fns.py +0 -0
  98. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/state_dict.py +0 -0
  99. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/tree.py +0 -0
  100. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/tree_util.py +0 -0
  101. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/types.py +0 -0
  102. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/util.py +0 -0
  103. {haliax-1.4.dev447 → haliax-1.4.dev452}/src/haliax/wrap.py +0 -0
  104. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/core_test.py +0 -0
  105. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_attention.py +0 -0
  106. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_axis.py +0 -0
  107. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_bitwise_ops.py +0 -0
  108. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_conv.py +0 -0
  109. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_debug.py +0 -0
  110. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_dot.py +0 -0
  111. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_dtype_typing.py +0 -0
  112. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_einsum.py +0 -0
  113. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_fft.py +0 -0
  114. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_field.py +0 -0
  115. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_fp8.py +0 -0
  116. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_hof.py +0 -0
  117. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_int8.py +0 -0
  118. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_moe_linear.py +0 -0
  119. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_mup_coordinate_check.py +0 -0
  120. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_mup_embedding.py +0 -0
  121. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_mup_linear.py +0 -0
  122. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_named_ref.py +0 -0
  123. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_namedarray_typing.py +0 -0
  124. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_nan_reductions.py +0 -0
  125. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_nn.py +0 -0
  126. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_ops.py +0 -0
  127. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_parsing.py +0 -0
  128. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_partitioning.py +0 -0
  129. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_poly_ops.py +0 -0
  130. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_pool.py +0 -0
  131. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_random.py +0 -0
  132. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_rearrange.py +0 -0
  133. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_scan.py +0 -0
  134. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_scatter_gather.py +0 -0
  135. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_specialized_fns.py +0 -0
  136. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_tree_util.py +0 -0
  137. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_utils.py +0 -0
  138. {haliax-1.4.dev447 → haliax-1.4.dev452}/tests/test_visualize_sharding.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev447
3
+ Version: 1.4.dev452
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/
@@ -45,6 +45,7 @@ dev = [
45
45
  "pymdown-extensions",
46
46
  "pygments",
47
47
  "chex>=0.1.86",
48
+ "safetensors>=0.4.3",
48
49
  "pre-commit",
49
50
  ]
50
51
 
@@ -3,4 +3,4 @@
3
3
  # SPDX-License-Identifier: Apache-2.0
4
4
 
5
5
 
6
- __version__ = "1.4.dev447"
6
+ __version__ = "1.4.dev452"
@@ -26,8 +26,10 @@ from haliax.tree_util import scan_aware_tree_map
26
26
 
27
27
  try:
28
28
  import safetensors
29
+ import safetensors.numpy as safetensors_numpy
29
30
  except ImportError:
30
31
  safetensors = None
32
+ safetensors_numpy = None
31
33
 
32
34
 
33
35
  StateDict = dict[str, Any]
@@ -436,7 +438,9 @@ def save_state_dict(state_dict: StateDict, path):
436
438
  state_dict = {k: v for k, v in state_dict.items() if v is not None}
437
439
  if jax.process_index() == 0:
438
440
  # the "pt" is a lie but it doesn't seem to actually matter and HF demands it
439
- safetensors.numpy.save_file(state_dict, path, metadata={"format": "pt"})
441
+ if safetensors_numpy is None:
442
+ raise ImportError("safetensors_numpy is not installed")
443
+ safetensors_numpy.save_file(state_dict, path, metadata={"format": "pt"})
440
444
  global _GLOBAL_SAVE_COUNT
441
445
  sync_global_devices(f"save_state_dict {_GLOBAL_SAVE_COUNT}")
442
446
  _GLOBAL_SAVE_COUNT += 1
@@ -447,5 +451,7 @@ def load_state_dict(path):
447
451
  Load a model's state dict from a file, bringing all tensors to the CPU first and then converting to numpy.
448
452
  This will load using safetensors format
449
453
  """
450
- state_dict = safetensors.numpy.load_file(path)
454
+ if safetensors_numpy is None:
455
+ raise ImportError("safetensors_numpy is not installed")
456
+ state_dict = safetensors_numpy.load_file(path)
451
457
  return state_dict
@@ -158,12 +158,17 @@ class Linear(ModuleWithStateDictSerialization, ReparamEnabled):
158
158
  return self.weight.axes[-len(self.Out) :] != self.Out
159
159
 
160
160
  def to_state_dict(self, prefix: Optional[str] = None) -> StateDict:
161
- scaled = dataclasses.replace(self, weight=self.weight * self.reparam.active_scale)
161
+ # weight can be None for certain filtering things like LoRA
162
+ scaled = dataclasses.replace(
163
+ self, weight=self.weight * self.reparam.active_scale if self.weight is not None else None
164
+ )
162
165
  return default_eqx_module_to_state_dict(scaled, prefix)
163
166
 
164
167
  def from_state_dict(self: Mod, state_dict: StateDict, prefix: Optional[str] = None) -> Mod:
165
168
  unscaled = default_eqx_module_from_state_dict(self, state_dict, prefix)
166
- return dataclasses.replace(unscaled, weight=unscaled.weight / self.reparam.active_scale)
169
+ if unscaled.weight is not None:
170
+ unscaled = dataclasses.replace(unscaled, weight=unscaled.weight / self.reparam.active_scale)
171
+ return unscaled
167
172
 
168
173
  @staticmethod
169
174
  def input_reparam(use_mup: bool = True) -> type[AbstractLinearReparam]:
@@ -9,13 +9,14 @@ from typing import Any
9
9
  import equinox as eqx
10
10
  import jax
11
11
  import jax.numpy as jnp
12
+ import numpy as np
12
13
  import pytest
13
14
 
14
15
  import haliax as hax
15
16
  from haliax._src.state_dict import flatten_modules_for_export, unflatten_modules_from_export
16
17
  from haliax.nn import Linear
17
18
  from haliax.nn.scan import Stacked, _stack_state_dict, _unstack_state_dict
18
- from haliax.state_dict import from_state_dict, to_state_dict
19
+ from haliax.state_dict import from_state_dict, load_state_dict, save_state_dict, to_state_dict
19
20
 
20
21
 
21
22
  @pytest.mark.parametrize("out_dims_first", [True, False])
@@ -159,6 +160,44 @@ def test_export_layer_norm():
159
160
  assert layer_norm == new_layer_norm
160
161
 
161
162
 
163
+ def test_save_state_dict_roundtrip(tmp_path):
164
+ pytest.importorskip("safetensors.numpy")
165
+ state_dict = {
166
+ "weight": np.arange(6, dtype=np.float32).reshape(2, 3),
167
+ "bias": np.linspace(0.0, 1.0, 3, dtype=np.float32),
168
+ }
169
+ path = tmp_path / "state.safetensors"
170
+
171
+ save_state_dict(state_dict, path)
172
+
173
+ assert path.exists()
174
+ loaded = load_state_dict(path)
175
+
176
+ assert set(loaded.keys()) == {"weight", "bias"}
177
+ np.testing.assert_array_equal(loaded["weight"], state_dict["weight"])
178
+ np.testing.assert_array_equal(loaded["bias"], state_dict["bias"])
179
+
180
+
181
+ def test_save_state_dict_filters_none_and_sets_metadata(tmp_path):
182
+ safetensors = pytest.importorskip("safetensors")
183
+
184
+ state_dict = {
185
+ "weight": np.ones((4,), dtype=np.float32),
186
+ "optional": None,
187
+ }
188
+ path = tmp_path / "state.safetensors"
189
+
190
+ save_state_dict(state_dict, path)
191
+
192
+ with safetensors.safe_open(str(path), framework="numpy") as f:
193
+ keys = list(f.keys())
194
+
195
+ assert "weight" in keys
196
+ assert "optional" not in keys
197
+ np.testing.assert_array_equal(f.get_tensor("weight"), state_dict["weight"])
198
+ assert f.metadata().get("format") == "pt"
199
+
200
+
162
201
  def test_stacked_layer_norm():
163
202
  L = hax.Axis("L", 4)
164
203
  D = hax.Axis("D", 10)
@@ -315,6 +315,7 @@ dev = [
315
315
  { name = "pygments" },
316
316
  { name = "pymdown-extensions" },
317
317
  { name = "pytest" },
318
+ { name = "safetensors" },
318
319
  ]
319
320
 
320
321
  [package.metadata]
@@ -342,6 +343,7 @@ dev = [
342
343
  { name = "pygments" },
343
344
  { name = "pymdown-extensions" },
344
345
  { name = "pytest", specifier = ">=7.4.0" },
346
+ { name = "safetensors", specifier = ">=0.4.3" },
345
347
  ]
346
348
 
347
349
  [[package]]
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes