haliax 1.4.dev372__tar.gz → 1.4.dev373__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 (114) hide show
  1. {haliax-1.4.dev372 → haliax-1.4.dev373}/PKG-INFO +1 -1
  2. haliax-1.4.dev373/src/haliax/__about__.py +1 -0
  3. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/_src/state_dict.py +1 -1
  4. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/partitioning.py +35 -3
  5. haliax-1.4.dev372/src/haliax/__about__.py +0 -1
  6. {haliax-1.4.dev372 → haliax-1.4.dev373}/.coveragerc +0 -0
  7. {haliax-1.4.dev372 → haliax-1.4.dev373}/.flake8 +0 -0
  8. {haliax-1.4.dev372 → haliax-1.4.dev373}/.github/workflows/publish_dev.yaml +0 -0
  9. {haliax-1.4.dev372 → haliax-1.4.dev373}/.github/workflows/run_pre_commit.yaml +0 -0
  10. {haliax-1.4.dev372 → haliax-1.4.dev373}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  11. {haliax-1.4.dev372 → haliax-1.4.dev373}/.github/workflows/run_tests.yaml +0 -0
  12. {haliax-1.4.dev372 → haliax-1.4.dev373}/.gitignore +0 -0
  13. {haliax-1.4.dev372 → haliax-1.4.dev373}/.playbooks/add-types.md +0 -0
  14. {haliax-1.4.dev372 → haliax-1.4.dev373}/.pre-commit-config.yaml +0 -0
  15. {haliax-1.4.dev372 → haliax-1.4.dev373}/.readthedocs.yaml +0 -0
  16. {haliax-1.4.dev372 → haliax-1.4.dev373}/AGENTS.md +0 -0
  17. {haliax-1.4.dev372 → haliax-1.4.dev373}/CONTRIBUTING.md +0 -0
  18. {haliax-1.4.dev372 → haliax-1.4.dev373}/LICENSE +0 -0
  19. {haliax-1.4.dev372 → haliax-1.4.dev373}/README.md +0 -0
  20. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/api.md +0 -0
  21. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/broadcasting.md +0 -0
  22. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/cheatsheet.md +0 -0
  23. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/css/material.css +0 -0
  24. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/css/mkdocstrings.css +0 -0
  25. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/faq.md +0 -0
  26. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/figures/data_parallel_mesh.png +0 -0
  27. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  28. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/figures/device_mesh_1d.png +0 -0
  29. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/figures/device_mesh_1d_zero.png +0 -0
  30. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/figures/device_mesh_2d.png +0 -0
  31. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  32. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  33. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  34. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  35. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/figures/device_mesh_2d_zero.png +0 -0
  36. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/fp8.md +0 -0
  37. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/index.md +0 -0
  38. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/indexing.md +0 -0
  39. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/matmul.md +0 -0
  40. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/nn.md +0 -0
  41. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/partitioning.md +0 -0
  42. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/rearrange.ipynb +0 -0
  43. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/rearrange.md +0 -0
  44. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/requirements.txt +0 -0
  45. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/scan.md +0 -0
  46. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/state-dict.md +0 -0
  47. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/tutorial.md +0 -0
  48. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/typing.md +0 -0
  49. {haliax-1.4.dev372 → haliax-1.4.dev373}/docs/vmap.md +0 -0
  50. {haliax-1.4.dev372 → haliax-1.4.dev373}/mkdocs.yml +0 -0
  51. {haliax-1.4.dev372 → haliax-1.4.dev373}/pyproject.toml +0 -0
  52. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/__init__.py +0 -0
  53. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/_src/__init__.py +0 -0
  54. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/_src/compile_utils.py +0 -0
  55. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/_src/dot.py +0 -0
  56. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/_src/einsum.py +0 -0
  57. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/_src/fp8.py +0 -0
  58. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/_src/parsing.py +0 -0
  59. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/_src/rearrange.py +0 -0
  60. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/_src/scan.py +0 -0
  61. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/_src/util.py +0 -0
  62. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/axis.py +0 -0
  63. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/core.py +0 -0
  64. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/debug.py +0 -0
  65. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/haxtyping.py +0 -0
  66. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/hof.py +0 -0
  67. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/jax_utils.py +0 -0
  68. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/nn/__init__.py +0 -0
  69. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/nn/activations.py +0 -0
  70. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/nn/attention.py +0 -0
  71. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/nn/conv.py +0 -0
  72. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/nn/dropout.py +0 -0
  73. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/nn/embedding.py +0 -0
  74. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/nn/linear.py +0 -0
  75. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/nn/loss.py +0 -0
  76. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/nn/mlp.py +0 -0
  77. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/nn/normalization.py +0 -0
  78. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/nn/pool.py +0 -0
  79. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/nn/scan.py +0 -0
  80. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/ops.py +0 -0
  81. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/quantization.py +0 -0
  82. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/random.py +0 -0
  83. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/specialized_fns.py +0 -0
  84. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/state_dict.py +0 -0
  85. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/tree_util.py +0 -0
  86. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/types.py +0 -0
  87. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/util.py +0 -0
  88. {haliax-1.4.dev372 → haliax-1.4.dev373}/src/haliax/wrap.py +0 -0
  89. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/core_test.py +0 -0
  90. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_attention.py +0 -0
  91. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_axis.py +0 -0
  92. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_conv.py +0 -0
  93. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_debug.py +0 -0
  94. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_dot.py +0 -0
  95. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_dtype_typing.py +0 -0
  96. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_einsum.py +0 -0
  97. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_fp8.py +0 -0
  98. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_hof.py +0 -0
  99. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_int8.py +0 -0
  100. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_namedarray_typing.py +0 -0
  101. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_nn.py +0 -0
  102. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_ops.py +0 -0
  103. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_parsing.py +0 -0
  104. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_partitioning.py +0 -0
  105. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_pool.py +0 -0
  106. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_random.py +0 -0
  107. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_rearrange.py +0 -0
  108. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_scan.py +0 -0
  109. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_scatter_gather.py +0 -0
  110. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_specialized_fns.py +0 -0
  111. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_state_dict.py +0 -0
  112. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_tree_util.py +0 -0
  113. {haliax-1.4.dev372 → haliax-1.4.dev373}/tests/test_utils.py +0 -0
  114. {haliax-1.4.dev372 → haliax-1.4.dev373}/uv.lock +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev372
3
+ Version: 1.4.dev373
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.dev373"
@@ -197,7 +197,7 @@ def from_state_dict(tree: T, state_dict: StateDict, prefix: Optional[str] = None
197
197
  if isinstance(array, np.ndarray):
198
198
  mesh = partitioning._get_mesh()
199
199
  # TODO: modernize this
200
- if mesh.devices.size > 1: # this happens with the default mesh
200
+ if jax.device_count() > 1: # this happens with the default mesh
201
201
  pspec = partitioning.pspec_for_axis(tree.axes)
202
202
  sharding = jax.sharding.NamedSharding(mesh, pspec)
203
203
  array = jax.make_array_from_callback(tree.array.shape, sharding, lambda indices: array[indices])
@@ -10,7 +10,24 @@ import equinox as eqx
10
10
  import jax
11
11
  from equinox import is_array, module_update_wrapper
12
12
  from jax.lax import with_sharding_constraint
13
- from jax.sharding import Mesh, NamedSharding, PartitionSpec, SingleDeviceSharding
13
+ from jax.sharding import (
14
+ Mesh,
15
+ NamedSharding,
16
+ PartitionSpec,
17
+ SingleDeviceSharding,
18
+ )
19
+
20
+ try: # jax>=0.4.26
21
+ from jax.sharding import AbstractMesh, get_abstract_mesh
22
+ except Exception: # pragma: no cover - older JAX versions
23
+ AbstractMesh = Mesh # type: ignore[misc,assignment]
24
+ def get_abstract_mesh(): # type: ignore[dead-code]
25
+ try:
26
+ from jax.interpreters.pxla import thread_resources
27
+ except Exception:
28
+ from jax.experimental.maps import thread_resources
29
+
30
+ return thread_resources.env.physical_mesh
14
31
  from jaxtyping import PyTree
15
32
 
16
33
  import haliax.tree_util as htu
@@ -604,10 +621,25 @@ def round_axis_for_partitioning(axis: Axis, mapping: Optional[ResourceMapping] =
604
621
  return Axis(axis.name, new_size)
605
622
 
606
623
 
607
- def _get_mesh() -> Mesh:
624
+ def _get_mesh() -> Mesh | AbstractMesh:
625
+ """Return the current mesh.
626
+
627
+ On newer versions of JAX this prefers ``get_abstract_mesh`` which does not
628
+ capture concrete devices. If no abstract mesh is currently active we fall
629
+ back to the concrete mesh used by ``Mesh``'s context manager so existing
630
+ code continues to work.
631
+ """
632
+
633
+ try: # jax>=0.4.26
634
+ mesh = get_abstract_mesh()
635
+ if not getattr(mesh, "empty", False):
636
+ return mesh
637
+ except Exception: # pragma: no cover - older JAX versions
638
+ pass
639
+
608
640
  try:
609
641
  from jax.interpreters.pxla import thread_resources
610
- except ImportError:
642
+ except Exception: # pragma: no cover - jax<0.4
611
643
  from jax.experimental.maps import thread_resources
612
644
 
613
645
  return thread_resources.env.physical_mesh
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev372"
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