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