haliax 1.4.dev388__tar.gz → 1.4.dev390__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 (116) hide show
  1. {haliax-1.4.dev388 → haliax-1.4.dev390}/PKG-INFO +1 -1
  2. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/faq.md +6 -0
  3. haliax-1.4.dev390/src/haliax/__about__.py +1 -0
  4. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/_src/scan.py +1 -1
  5. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/debug.py +71 -1
  6. haliax-1.4.dev390/tests/test_visualize_sharding.py +67 -0
  7. haliax-1.4.dev388/src/haliax/__about__.py +0 -1
  8. {haliax-1.4.dev388 → haliax-1.4.dev390}/.coveragerc +0 -0
  9. {haliax-1.4.dev388 → haliax-1.4.dev390}/.flake8 +0 -0
  10. {haliax-1.4.dev388 → haliax-1.4.dev390}/.github/workflows/publish_dev.yaml +0 -0
  11. {haliax-1.4.dev388 → haliax-1.4.dev390}/.github/workflows/run_pre_commit.yaml +0 -0
  12. {haliax-1.4.dev388 → haliax-1.4.dev390}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  13. {haliax-1.4.dev388 → haliax-1.4.dev390}/.github/workflows/run_tests.yaml +0 -0
  14. {haliax-1.4.dev388 → haliax-1.4.dev390}/.gitignore +0 -0
  15. {haliax-1.4.dev388 → haliax-1.4.dev390}/.playbooks/add-types.md +0 -0
  16. {haliax-1.4.dev388 → haliax-1.4.dev390}/.playbooks/wrap-non-named.md +0 -0
  17. {haliax-1.4.dev388 → haliax-1.4.dev390}/.pre-commit-config.yaml +0 -0
  18. {haliax-1.4.dev388 → haliax-1.4.dev390}/.readthedocs.yaml +0 -0
  19. {haliax-1.4.dev388 → haliax-1.4.dev390}/AGENTS.md +0 -0
  20. {haliax-1.4.dev388 → haliax-1.4.dev390}/CONTRIBUTING.md +0 -0
  21. {haliax-1.4.dev388 → haliax-1.4.dev390}/LICENSE +0 -0
  22. {haliax-1.4.dev388 → haliax-1.4.dev390}/README.md +0 -0
  23. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/api.md +0 -0
  24. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/broadcasting.md +0 -0
  25. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/cheatsheet.md +0 -0
  26. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/css/material.css +0 -0
  27. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/css/mkdocstrings.css +0 -0
  28. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/figures/data_parallel_mesh.png +0 -0
  29. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  30. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/figures/device_mesh_1d.png +0 -0
  31. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/figures/device_mesh_1d_zero.png +0 -0
  32. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/figures/device_mesh_2d.png +0 -0
  33. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  34. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  35. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  36. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  37. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/figures/device_mesh_2d_zero.png +0 -0
  38. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/fp8.md +0 -0
  39. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/index.md +0 -0
  40. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/indexing.md +0 -0
  41. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/matmul.md +0 -0
  42. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/nn.md +0 -0
  43. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/partitioning.md +0 -0
  44. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/rearrange.ipynb +0 -0
  45. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/rearrange.md +0 -0
  46. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/requirements.txt +0 -0
  47. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/scan.md +0 -0
  48. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/state-dict.md +0 -0
  49. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/tutorial.md +0 -0
  50. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/typing.md +0 -0
  51. {haliax-1.4.dev388 → haliax-1.4.dev390}/docs/vmap.md +0 -0
  52. {haliax-1.4.dev388 → haliax-1.4.dev390}/mkdocs.yml +0 -0
  53. {haliax-1.4.dev388 → haliax-1.4.dev390}/pyproject.toml +0 -0
  54. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/__init__.py +0 -0
  55. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/_src/__init__.py +0 -0
  56. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/_src/compile_utils.py +0 -0
  57. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/_src/dot.py +0 -0
  58. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/_src/einsum.py +0 -0
  59. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/_src/fp8.py +0 -0
  60. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/_src/parsing.py +0 -0
  61. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/_src/rearrange.py +0 -0
  62. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/_src/state_dict.py +0 -0
  63. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/_src/util.py +0 -0
  64. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/axis.py +0 -0
  65. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/core.py +0 -0
  66. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/haxtyping.py +0 -0
  67. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/hof.py +0 -0
  68. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/jax_utils.py +0 -0
  69. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/nn/__init__.py +0 -0
  70. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/nn/activations.py +0 -0
  71. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/nn/attention.py +0 -0
  72. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/nn/conv.py +0 -0
  73. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/nn/dropout.py +0 -0
  74. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/nn/embedding.py +0 -0
  75. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/nn/linear.py +0 -0
  76. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/nn/loss.py +0 -0
  77. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/nn/mlp.py +0 -0
  78. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/nn/normalization.py +0 -0
  79. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/nn/pool.py +0 -0
  80. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/nn/scan.py +0 -0
  81. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/ops.py +0 -0
  82. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/partitioning.py +0 -0
  83. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/quantization.py +0 -0
  84. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/random.py +0 -0
  85. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/specialized_fns.py +0 -0
  86. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/state_dict.py +0 -0
  87. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/tree_util.py +0 -0
  88. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/types.py +0 -0
  89. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/util.py +0 -0
  90. {haliax-1.4.dev388 → haliax-1.4.dev390}/src/haliax/wrap.py +0 -0
  91. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/core_test.py +0 -0
  92. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_attention.py +0 -0
  93. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_axis.py +0 -0
  94. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_conv.py +0 -0
  95. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_debug.py +0 -0
  96. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_dot.py +0 -0
  97. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_dtype_typing.py +0 -0
  98. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_einsum.py +0 -0
  99. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_fp8.py +0 -0
  100. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_hof.py +0 -0
  101. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_int8.py +0 -0
  102. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_namedarray_typing.py +0 -0
  103. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_nn.py +0 -0
  104. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_ops.py +0 -0
  105. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_parsing.py +0 -0
  106. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_partitioning.py +0 -0
  107. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_pool.py +0 -0
  108. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_random.py +0 -0
  109. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_rearrange.py +0 -0
  110. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_scan.py +0 -0
  111. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_scatter_gather.py +0 -0
  112. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_specialized_fns.py +0 -0
  113. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_state_dict.py +0 -0
  114. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_tree_util.py +0 -0
  115. {haliax-1.4.dev388 → haliax-1.4.dev390}/tests/test_utils.py +0 -0
  116. {haliax-1.4.dev388 → haliax-1.4.dev390}/uv.lock +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev388
3
+ Version: 1.4.dev390
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/
@@ -9,3 +9,9 @@ Currently, we diagnose:
9
9
 
10
10
  * Reuse of arrays or NamedArrays in a field. [Equinox modules must be trees.](https://docs.kidger.site/equinox/faq/#a-module-saved-in-two-places-has-become-two-independent-copies)
11
11
  * Use of arrays or NamedArrays in a static field. Static data in JAX/Equinox must be hashable, and arrays are not hashable.
12
+
13
+ ## Tip 2: `hax.debug.visualize_shardings`
14
+
15
+ Use `hax.debug.visualize_shardings` to quickly inspect how a PyTree is sharded.
16
+ It prints the sharding of each array leaf, including the mapping from named axes
17
+ to physical axes for :class:`haliax.NamedArray` leaves.
@@ -0,0 +1 @@
1
+ __version__ = "1.4.dev390"
@@ -33,7 +33,7 @@ class ScanFn(Protocol[Carry, Args, Y]):
33
33
  ...
34
34
 
35
35
 
36
- @dataclasses.dataclass
36
+ @dataclasses.dataclass(frozen=True)
37
37
  class ScanCheckpointPolicy:
38
38
  """
39
39
  A class that represents a gradient checkpoint policy for blocks in a Stacked module. This is used to control
@@ -1,11 +1,14 @@
1
1
  import dataclasses
2
- from typing import List, Tuple, Union
2
+ from typing import List, Tuple, Union, Sequence
3
3
 
4
4
  import equinox as eqx
5
+ import jax
5
6
  import jax.numpy as jnp
6
7
  import jax.tree_util as jtu
7
8
 
9
+
8
10
  from haliax.core import NamedArray
11
+ from haliax.axis import Axis
9
12
  from haliax.util import is_jax_or_hax_array_like
10
13
 
11
14
  from ._src.util import IdentityMap
@@ -117,3 +120,70 @@ def _check_for_static_arrays(problems, module):
117
120
 
118
121
  if static_arrays:
119
122
  problems.static_arrays.extend(static_arrays)
123
+
124
+
125
+ def _pspec_parts(spec_part) -> str:
126
+ if spec_part is None:
127
+ return "unsharded"
128
+ elif isinstance(spec_part, (tuple, list)):
129
+ return "+".join(str(p) for p in spec_part)
130
+ else:
131
+ return str(spec_part)
132
+
133
+
134
+ def visualize_named_sharding(axes: Sequence[Axis], sharding: jax.sharding.Sharding) -> None:
135
+ """Visualize the sharding for a set of named axes.
136
+
137
+ This extends :func:`jax.debug.visualize_sharding` to handle arrays with more
138
+ than two dimensions by falling back to a textual description when necessary.
139
+ """
140
+
141
+ try:
142
+ pspec = sharding.spec # type: ignore[attr-defined]
143
+ except Exception:
144
+ pspec = (None,) * len(axes)
145
+
146
+ parts = [_pspec_parts(p) for p in pspec]
147
+ num_sharded = sum(p != "unsharded" for p in parts)
148
+
149
+ if num_sharded <= 2:
150
+ try:
151
+ jax.debug.visualize_sharding([ax.size for ax in axes], sharding)
152
+ except Exception:
153
+ pass
154
+
155
+ mapping = ", ".join(f"{ax.name}->{part}" for ax, part in zip(axes, parts))
156
+ print(mapping)
157
+
158
+
159
+ def visualize_shardings(tree) -> None:
160
+ """Print the sharding for each array-like leaf in ``tree``.
161
+
162
+ Both :class:`NamedArray` and regular JAX arrays are supported. NamedArrays
163
+ will show the mapping from logical axis names to physical axes. Plain arrays
164
+ will fall back to :func:`jax.debug.visualize_sharding`.
165
+ """
166
+
167
+ import haliax.tree_util as htu
168
+
169
+ def _show(x):
170
+ if isinstance(x, NamedArray):
171
+ arr = x.array
172
+ axes = x.axes
173
+ else:
174
+ arr = x
175
+ axes = None
176
+
177
+ def cb(sh):
178
+ if axes is not None:
179
+ visualize_named_sharding(axes, sh)
180
+ else:
181
+ try:
182
+ jax.debug.visualize_sharding(arr.shape, sh)
183
+ except Exception:
184
+ pass
185
+
186
+ jax.debug.inspect_array_sharding(arr, callback=cb)
187
+ return x
188
+
189
+ htu.tree_map(_show, tree, is_leaf=is_jax_or_hax_array_like)
@@ -0,0 +1,67 @@
1
+ import numpy as np
2
+ import jax
3
+ import jax.numpy as jnp
4
+
5
+ import haliax as hax
6
+ from haliax import Axis
7
+ from haliax.partitioning import ResourceAxis, axis_mapping, named_jit
8
+ from test_utils import skip_if_not_enough_devices
9
+ from haliax.debug import visualize_shardings
10
+
11
+ Dim1 = Axis("dim1", 8)
12
+ Dim2 = Axis("dim2", 8)
13
+ Dim3 = Axis("dim3", 2)
14
+
15
+ resource_map = {
16
+ "dim1": ResourceAxis.DATA,
17
+ "dim2": ResourceAxis.MODEL,
18
+ "dim3": ResourceAxis.REPLICA,
19
+ }
20
+
21
+
22
+ def test_visualize_shardings_runs(capsys):
23
+ mesh = jax.sharding.Mesh(
24
+ np.array(jax.devices()).reshape(-1, 1, 1),
25
+ (ResourceAxis.DATA, ResourceAxis.MODEL, ResourceAxis.REPLICA),
26
+ )
27
+ with axis_mapping(resource_map), mesh:
28
+ arr = hax.ones((Dim1, Dim2, Dim3))
29
+ visualize_shardings(arr)
30
+
31
+ out = capsys.readouterr().out
32
+ assert "dim1" in out and "dim2" in out and "dim3" in out
33
+
34
+
35
+ def test_visualize_shardings_inside_jit(capsys):
36
+ mesh = jax.sharding.Mesh(np.array(jax.devices()).reshape(-1, 1), (ResourceAxis.DATA, ResourceAxis.MODEL))
37
+
38
+ @named_jit(out_axis_resources={"dim1": ResourceAxis.DATA})
39
+ def fn(x):
40
+ visualize_shardings(x)
41
+ return x
42
+
43
+ with axis_mapping({"dim1": ResourceAxis.DATA}), mesh:
44
+ x = hax.ones(Dim1)
45
+ fn(x)
46
+
47
+ out = capsys.readouterr().out
48
+ assert "dim1" in out
49
+
50
+
51
+ def test_visualize_shardings_plain_array(capsys):
52
+ x = jnp.ones((4, 4))
53
+ visualize_shardings(x)
54
+ out = capsys.readouterr().out
55
+ assert out.strip() != ""
56
+
57
+
58
+ @skip_if_not_enough_devices(2)
59
+ def test_visualize_shardings_model_axis(capsys):
60
+ devices = jax.devices()
61
+ mesh = jax.sharding.Mesh(np.array(devices).reshape(-1, 2), (ResourceAxis.DATA, ResourceAxis.MODEL))
62
+ with axis_mapping({"dim1": ResourceAxis.DATA, "dim2": ResourceAxis.MODEL}), mesh:
63
+ arr = hax.ones((Dim1, Dim2))
64
+ visualize_shardings(arr)
65
+
66
+ out = capsys.readouterr().out
67
+ assert "dim2" in out
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev388"
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