haliax 1.4.dev402__tar.gz → 1.4.dev403__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 (120) hide show
  1. {haliax-1.4.dev402 → haliax-1.4.dev403}/PKG-INFO +1 -1
  2. haliax-1.4.dev403/src/haliax/__about__.py +1 -0
  3. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/partitioning.py +33 -1
  4. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_partitioning.py +28 -0
  5. haliax-1.4.dev402/src/haliax/__about__.py +0 -1
  6. {haliax-1.4.dev402 → haliax-1.4.dev403}/.coveragerc +0 -0
  7. {haliax-1.4.dev402 → haliax-1.4.dev403}/.flake8 +0 -0
  8. {haliax-1.4.dev402 → haliax-1.4.dev403}/.github/workflows/publish_dev.yaml +0 -0
  9. {haliax-1.4.dev402 → haliax-1.4.dev403}/.github/workflows/run_pre_commit.yaml +0 -0
  10. {haliax-1.4.dev402 → haliax-1.4.dev403}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  11. {haliax-1.4.dev402 → haliax-1.4.dev403}/.github/workflows/run_tests.yaml +0 -0
  12. {haliax-1.4.dev402 → haliax-1.4.dev403}/.gitignore +0 -0
  13. {haliax-1.4.dev402 → haliax-1.4.dev403}/.playbooks/add-types.md +0 -0
  14. {haliax-1.4.dev402 → haliax-1.4.dev403}/.playbooks/wrap-non-named.md +0 -0
  15. {haliax-1.4.dev402 → haliax-1.4.dev403}/.pre-commit-config.yaml +0 -0
  16. {haliax-1.4.dev402 → haliax-1.4.dev403}/.readthedocs.yaml +0 -0
  17. {haliax-1.4.dev402 → haliax-1.4.dev403}/AGENTS.md +0 -0
  18. {haliax-1.4.dev402 → haliax-1.4.dev403}/CONTRIBUTING.md +0 -0
  19. {haliax-1.4.dev402 → haliax-1.4.dev403}/LICENSE +0 -0
  20. {haliax-1.4.dev402 → haliax-1.4.dev403}/README.md +0 -0
  21. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/api.md +0 -0
  22. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/broadcasting.md +0 -0
  23. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/cheatsheet.md +0 -0
  24. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/css/material.css +0 -0
  25. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/css/mkdocstrings.css +0 -0
  26. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/faq.md +0 -0
  27. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/data_parallel_mesh.png +0 -0
  28. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  29. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/device_mesh_1d.png +0 -0
  30. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/device_mesh_1d_zero.png +0 -0
  31. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/device_mesh_2d.png +0 -0
  32. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  33. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  34. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  35. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  36. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/figures/device_mesh_2d_zero.png +0 -0
  37. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/fp8.md +0 -0
  38. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/index.md +0 -0
  39. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/indexing.md +0 -0
  40. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/matmul.md +0 -0
  41. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/nn.md +0 -0
  42. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/partitioning.md +0 -0
  43. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/primer.md +0 -0
  44. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/rearrange.ipynb +0 -0
  45. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/rearrange.md +0 -0
  46. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/requirements.txt +0 -0
  47. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/scan.md +0 -0
  48. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/state-dict.md +0 -0
  49. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/tutorial.md +0 -0
  50. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/typing.md +0 -0
  51. {haliax-1.4.dev402 → haliax-1.4.dev403}/docs/vmap.md +0 -0
  52. {haliax-1.4.dev402 → haliax-1.4.dev403}/mkdocs.yml +0 -0
  53. {haliax-1.4.dev402 → haliax-1.4.dev403}/pyproject.toml +0 -0
  54. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/__init__.py +0 -0
  55. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/__init__.py +0 -0
  56. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/compile_utils.py +0 -0
  57. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/dot.py +0 -0
  58. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/einsum.py +0 -0
  59. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/fp8.py +0 -0
  60. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/parsing.py +0 -0
  61. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/rearrange.py +0 -0
  62. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/scan.py +0 -0
  63. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/state_dict.py +0 -0
  64. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/_src/util.py +0 -0
  65. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/axis.py +0 -0
  66. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/core.py +0 -0
  67. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/debug.py +0 -0
  68. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/field.py +0 -0
  69. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/haxtyping.py +0 -0
  70. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/hof.py +0 -0
  71. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/jax_utils.py +0 -0
  72. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/__init__.py +0 -0
  73. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/activations.py +0 -0
  74. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/attention.py +0 -0
  75. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/conv.py +0 -0
  76. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/dropout.py +0 -0
  77. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/embedding.py +0 -0
  78. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/linear.py +0 -0
  79. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/loss.py +0 -0
  80. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/mlp.py +0 -0
  81. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/normalization.py +0 -0
  82. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/pool.py +0 -0
  83. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/nn/scan.py +0 -0
  84. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/ops.py +0 -0
  85. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/quantization.py +0 -0
  86. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/random.py +0 -0
  87. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/specialized_fns.py +0 -0
  88. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/state_dict.py +0 -0
  89. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/tree_util.py +0 -0
  90. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/types.py +0 -0
  91. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/util.py +0 -0
  92. {haliax-1.4.dev402 → haliax-1.4.dev403}/src/haliax/wrap.py +0 -0
  93. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/core_test.py +0 -0
  94. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_attention.py +0 -0
  95. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_axis.py +0 -0
  96. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_conv.py +0 -0
  97. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_debug.py +0 -0
  98. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_dot.py +0 -0
  99. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_dtype_typing.py +0 -0
  100. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_einsum.py +0 -0
  101. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_field.py +0 -0
  102. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_fp8.py +0 -0
  103. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_hof.py +0 -0
  104. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_int8.py +0 -0
  105. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_moe_linear.py +0 -0
  106. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_namedarray_typing.py +0 -0
  107. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_nn.py +0 -0
  108. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_ops.py +0 -0
  109. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_parsing.py +0 -0
  110. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_pool.py +0 -0
  111. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_random.py +0 -0
  112. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_rearrange.py +0 -0
  113. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_scan.py +0 -0
  114. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_scatter_gather.py +0 -0
  115. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_specialized_fns.py +0 -0
  116. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_state_dict.py +0 -0
  117. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_tree_util.py +0 -0
  118. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_utils.py +0 -0
  119. {haliax-1.4.dev402 → haliax-1.4.dev403}/tests/test_visualize_sharding.py +0 -0
  120. {haliax-1.4.dev402 → haliax-1.4.dev403}/uv.lock +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev402
3
+ Version: 1.4.dev403
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.dev403"
@@ -1,4 +1,5 @@
1
1
  import contextlib
2
+ import dataclasses
2
3
  import functools
3
4
  import threading
4
5
  import typing
@@ -210,6 +211,37 @@ def pspec_for(
210
211
  return None
211
212
  else:
212
213
  return pspec_for_axis(node.axes, resource_mapping)
214
+ elif isinstance(node, eqx.Module):
215
+ # handle eqx.Module explicitly so that we can look at axis_names metadata
216
+ updates: dict[str, typing.Any] = {}
217
+ for field in dataclasses.fields(node):
218
+ if field.metadata.get("static", False):
219
+ continue
220
+
221
+ value = getattr(node, field.name)
222
+ axis_names = field.metadata.get("axis_names") if field.metadata is not None else None
223
+ if axis_names is not None and is_jax_array_like(value):
224
+ current_sharding = (
225
+ getattr(value, "sharding", None) if preserve_existing_shardings else None
226
+ )
227
+ if current_sharding is not None:
228
+ updates[field.name] = None
229
+ else:
230
+ updates[field.name] = pspec_for_axis(axis_names, resource_mapping)
231
+ else:
232
+ updates[field.name] = htu.tree_map(
233
+ partition_spec, value, is_leaf=lambda x: isinstance(x, eqx.Module)
234
+ )
235
+
236
+ new_node = object.__new__(type(node))
237
+ for field in dataclasses.fields(node):
238
+ object.__setattr__(
239
+ new_node,
240
+ field.name,
241
+ updates.get(field.name, getattr(node, field.name)),
242
+ )
243
+
244
+ return new_node
213
245
  elif is_jax_array_like(node):
214
246
  sharding = getattr(node, "sharding", None)
215
247
  # TODO: these are usually replicated. Is there a better way to tell?
@@ -232,7 +264,7 @@ def pspec_for(
232
264
  else:
233
265
  return None
234
266
 
235
- return htu.tree_map(partition_spec, tree)
267
+ return htu.tree_map(partition_spec, tree, is_leaf=lambda x: isinstance(x, eqx.Module))
236
268
 
237
269
 
238
270
  def infer_resource_partitions(
@@ -59,6 +59,34 @@ def test_pspec_for_named_axes():
59
59
  assert specs.unnamed1 == PartitionSpec(None)
60
60
 
61
61
 
62
+ class ArrayModule(eqx.Module):
63
+ arr: Array = hax.field(axis_names=("dim2", "dim3"))
64
+
65
+
66
+ def test_pspec_for_plain_array_axis_names():
67
+ mesh = Mesh(np.array(jax.devices()).reshape(-1, 1), (ResourceAxis.DATA, ResourceAxis.MODEL))
68
+ with axis_mapping(resource_map), mesh:
69
+ mod = ArrayModule(jnp.ones((Dim2.size, Dim3.size)))
70
+
71
+ specs: ArrayModule = pspec_for(mod, preserve_existing_shardings=False)
72
+
73
+ assert specs.arr == PartitionSpec(ResourceAxis.DATA, ResourceAxis.MODEL)
74
+
75
+
76
+ class NestedArrayModule(eqx.Module):
77
+ inner: ArrayModule
78
+
79
+
80
+ def test_pspec_for_plain_array_axis_names_nested_module():
81
+ mesh = Mesh(np.array(jax.devices()).reshape(-1, 1), (ResourceAxis.DATA, ResourceAxis.MODEL))
82
+ with axis_mapping(resource_map), mesh:
83
+ mod = NestedArrayModule(ArrayModule(jnp.ones((Dim2.size, Dim3.size))))
84
+
85
+ specs: NestedArrayModule = pspec_for(mod, preserve_existing_shardings=False)
86
+
87
+ assert specs.inner.arr == PartitionSpec(ResourceAxis.DATA, ResourceAxis.MODEL)
88
+
89
+
62
90
  class MyModuleInit(eqx.Module):
63
91
  named: NamedArray
64
92
  unnamed1: Array
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev402"
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