haliax 1.4.dev348__tar.gz → 1.4.dev351__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 (106) hide show
  1. {haliax-1.4.dev348 → haliax-1.4.dev351}/PKG-INFO +1 -1
  2. haliax-1.4.dev351/src/haliax/__about__.py +1 -0
  3. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/_src/state_dict.py +2 -0
  4. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_state_dict.py +24 -0
  5. haliax-1.4.dev348/src/haliax/__about__.py +0 -1
  6. {haliax-1.4.dev348 → haliax-1.4.dev351}/.coveragerc +0 -0
  7. {haliax-1.4.dev348 → haliax-1.4.dev351}/.flake8 +0 -0
  8. {haliax-1.4.dev348 → haliax-1.4.dev351}/.github/workflows/publish_dev.yaml +0 -0
  9. {haliax-1.4.dev348 → haliax-1.4.dev351}/.github/workflows/run_pre_commit.yaml +0 -0
  10. {haliax-1.4.dev348 → haliax-1.4.dev351}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  11. {haliax-1.4.dev348 → haliax-1.4.dev351}/.github/workflows/run_tests.yaml +0 -0
  12. {haliax-1.4.dev348 → haliax-1.4.dev351}/.gitignore +0 -0
  13. {haliax-1.4.dev348 → haliax-1.4.dev351}/.pre-commit-config.yaml +0 -0
  14. {haliax-1.4.dev348 → haliax-1.4.dev351}/.readthedocs.yaml +0 -0
  15. {haliax-1.4.dev348 → haliax-1.4.dev351}/CONTRIBUTING.md +0 -0
  16. {haliax-1.4.dev348 → haliax-1.4.dev351}/LICENSE +0 -0
  17. {haliax-1.4.dev348 → haliax-1.4.dev351}/README.md +0 -0
  18. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/api.md +0 -0
  19. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/broadcasting.md +0 -0
  20. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/cheatsheet.md +0 -0
  21. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/css/material.css +0 -0
  22. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/css/mkdocstrings.css +0 -0
  23. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/faq.md +0 -0
  24. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/figures/data_parallel_mesh.png +0 -0
  25. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  26. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/figures/device_mesh_1d.png +0 -0
  27. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/figures/device_mesh_1d_zero.png +0 -0
  28. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/figures/device_mesh_2d.png +0 -0
  29. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  30. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  31. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  32. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  33. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/figures/device_mesh_2d_zero.png +0 -0
  34. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/fp8.md +0 -0
  35. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/index.md +0 -0
  36. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/indexing.md +0 -0
  37. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/matmul.md +0 -0
  38. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/nn.md +0 -0
  39. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/partitioning.md +0 -0
  40. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/rearrange.ipynb +0 -0
  41. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/rearrange.md +0 -0
  42. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/requirements.txt +0 -0
  43. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/scan.md +0 -0
  44. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/state-dict.md +0 -0
  45. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/tutorial.md +0 -0
  46. {haliax-1.4.dev348 → haliax-1.4.dev351}/docs/vmap.md +0 -0
  47. {haliax-1.4.dev348 → haliax-1.4.dev351}/mkdocs.yml +0 -0
  48. {haliax-1.4.dev348 → haliax-1.4.dev351}/pyproject.toml +0 -0
  49. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/__init__.py +0 -0
  50. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/_src/__init__.py +0 -0
  51. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/_src/compile_utils.py +0 -0
  52. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/_src/dot.py +0 -0
  53. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/_src/einsum.py +0 -0
  54. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/_src/fp8.py +0 -0
  55. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/_src/parsing.py +0 -0
  56. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/_src/rearrange.py +0 -0
  57. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/_src/scan.py +0 -0
  58. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/_src/util.py +0 -0
  59. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/axis.py +0 -0
  60. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/core.py +0 -0
  61. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/debug.py +0 -0
  62. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/hof.py +0 -0
  63. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/jax_utils.py +0 -0
  64. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/nn/__init__.py +0 -0
  65. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/nn/activations.py +0 -0
  66. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/nn/attention.py +0 -0
  67. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/nn/conv.py +0 -0
  68. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/nn/dropout.py +0 -0
  69. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/nn/embedding.py +0 -0
  70. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/nn/linear.py +0 -0
  71. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/nn/loss.py +0 -0
  72. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/nn/mlp.py +0 -0
  73. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/nn/normalization.py +0 -0
  74. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/nn/pool.py +0 -0
  75. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/nn/scan.py +0 -0
  76. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/ops.py +0 -0
  77. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/partitioning.py +0 -0
  78. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/quantization.py +0 -0
  79. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/random.py +0 -0
  80. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/specialized_fns.py +0 -0
  81. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/state_dict.py +0 -0
  82. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/tree_util.py +0 -0
  83. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/types.py +0 -0
  84. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/util.py +0 -0
  85. {haliax-1.4.dev348 → haliax-1.4.dev351}/src/haliax/wrap.py +0 -0
  86. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/core_test.py +0 -0
  87. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_attention.py +0 -0
  88. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_axis.py +0 -0
  89. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_conv.py +0 -0
  90. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_debug.py +0 -0
  91. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_dot.py +0 -0
  92. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_einsum.py +0 -0
  93. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_fp8.py +0 -0
  94. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_hof.py +0 -0
  95. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_int8.py +0 -0
  96. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_nn.py +0 -0
  97. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_ops.py +0 -0
  98. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_parsing.py +0 -0
  99. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_partitioning.py +0 -0
  100. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_pool.py +0 -0
  101. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_random.py +0 -0
  102. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_rearrange.py +0 -0
  103. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_scan.py +0 -0
  104. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_specialized_fns.py +0 -0
  105. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_tree_util.py +0 -0
  106. {haliax-1.4.dev348 → haliax-1.4.dev351}/tests/test_utils.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev348
3
+ Version: 1.4.dev351
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.dev351"
@@ -215,6 +215,8 @@ def from_state_dict(tree: T, state_dict: StateDict, prefix: Optional[str] = None
215
215
  raise ValueError("Cannot extract a leaf value from a state dict without a prefix")
216
216
  # TODO: add "strict" flag so we can return None in cases where it's just missing
217
217
  return jnp.array(state_dict[prefix])
218
+ elif tree is None:
219
+ return None
218
220
  else:
219
221
  if prefix is None:
220
222
  return tree
@@ -175,3 +175,27 @@ def test_stacked_layer_norm():
175
175
  new_norms = unflatten_modules_from_export(norms_flat, norms2)
176
176
 
177
177
  assert norms == new_norms
178
+
179
+
180
+ def test_linear_doesnt_read_bias_if_it_didnt_have_bias():
181
+ H = hax.Axis("H", 10)
182
+ W = hax.Axis("W", 20)
183
+ D = hax.Axis("D", 30)
184
+ B = hax.Axis("B", 40)
185
+
186
+ linear = hax.nn.Linear.init((H, W), (D, B), key=jax.random.PRNGKey(0), use_bias=False, out_first=True)
187
+
188
+ flat_linear = linear.flatten_for_export()
189
+
190
+ flat_state_dict = to_state_dict(flat_linear)
191
+
192
+ assert "bias" not in flat_state_dict
193
+ flat_state_dict["bias"] = jnp.zeros((D.size * B.size,)) # add a dummy bias
194
+
195
+ # now unflatten it
196
+ linear2 = Linear.init((H, W), (D, B), key=jax.random.PRNGKey(1), use_bias=False, out_first=True)
197
+ flinear2 = linear2.flatten_for_export()
198
+ flinear2 = from_state_dict(flinear2, flat_state_dict)
199
+ new_linear = flinear2.unflatten_from_export(linear2)
200
+
201
+ assert linear == new_linear
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev348"
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes