haliax 1.4.dev325__tar.gz → 1.4.dev327__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.dev325 → haliax-1.4.dev327}/PKG-INFO +3 -2
  2. haliax-1.4.dev327/src/haliax/__about__.py +1 -0
  3. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/__init__.py +4 -0
  4. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/_src/state_dict.py +10 -11
  5. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/axis.py +47 -0
  6. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/core.py +90 -1
  7. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/partitioning.py +6 -6
  8. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/quantization.py +1 -2
  9. haliax-1.4.dev325/src/haliax/__about__.py +0 -1
  10. {haliax-1.4.dev325 → haliax-1.4.dev327}/.coveragerc +0 -0
  11. {haliax-1.4.dev325 → haliax-1.4.dev327}/.flake8 +0 -0
  12. {haliax-1.4.dev325 → haliax-1.4.dev327}/.github/workflows/publish_dev.yaml +0 -0
  13. {haliax-1.4.dev325 → haliax-1.4.dev327}/.github/workflows/run_pre_commit.yaml +0 -0
  14. {haliax-1.4.dev325 → haliax-1.4.dev327}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  15. {haliax-1.4.dev325 → haliax-1.4.dev327}/.github/workflows/run_tests.yaml +0 -0
  16. {haliax-1.4.dev325 → haliax-1.4.dev327}/.gitignore +0 -0
  17. {haliax-1.4.dev325 → haliax-1.4.dev327}/.pre-commit-config.yaml +0 -0
  18. {haliax-1.4.dev325 → haliax-1.4.dev327}/.readthedocs.yaml +0 -0
  19. {haliax-1.4.dev325 → haliax-1.4.dev327}/CONTRIBUTING.md +0 -0
  20. {haliax-1.4.dev325 → haliax-1.4.dev327}/LICENSE +0 -0
  21. {haliax-1.4.dev325 → haliax-1.4.dev327}/README.md +0 -0
  22. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/api.md +0 -0
  23. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/broadcasting.md +0 -0
  24. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/cheatsheet.md +0 -0
  25. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/css/material.css +0 -0
  26. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/css/mkdocstrings.css +0 -0
  27. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/faq.md +0 -0
  28. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/data_parallel_mesh.png +0 -0
  29. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  30. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/device_mesh_1d.png +0 -0
  31. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/device_mesh_1d_zero.png +0 -0
  32. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/device_mesh_2d.png +0 -0
  33. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  34. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  35. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  36. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  37. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/figures/device_mesh_2d_zero.png +0 -0
  38. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/fp8.md +0 -0
  39. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/hof.md +0 -0
  40. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/index.md +0 -0
  41. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/indexing.md +0 -0
  42. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/matmul.md +0 -0
  43. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/nn.md +0 -0
  44. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/partitioning.md +0 -0
  45. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/rearrange.ipynb +0 -0
  46. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/rearrange.md +0 -0
  47. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/requirements.txt +0 -0
  48. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/state-dict.md +0 -0
  49. {haliax-1.4.dev325 → haliax-1.4.dev327}/docs/tutorial.md +0 -0
  50. {haliax-1.4.dev325 → haliax-1.4.dev327}/mkdocs.yml +0 -0
  51. {haliax-1.4.dev325 → haliax-1.4.dev327}/pyproject.toml +0 -0
  52. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/_src/__init__.py +0 -0
  53. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/_src/compile_utils.py +0 -0
  54. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/_src/dot.py +0 -0
  55. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/_src/einsum.py +0 -0
  56. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/_src/fp8.py +0 -0
  57. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/_src/parsing.py +0 -0
  58. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/_src/rearrange.py +0 -0
  59. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/_src/util.py +0 -0
  60. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/debug.py +0 -0
  61. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/hof.py +0 -0
  62. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/jax_utils.py +0 -0
  63. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/__init__.py +0 -0
  64. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/activations.py +0 -0
  65. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/attention.py +0 -0
  66. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/conv.py +0 -0
  67. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/dropout.py +0 -0
  68. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/embedding.py +0 -0
  69. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/linear.py +0 -0
  70. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/loss.py +0 -0
  71. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/mlp.py +0 -0
  72. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/normalization.py +0 -0
  73. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/pool.py +0 -0
  74. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/nn/scan.py +0 -0
  75. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/ops.py +0 -0
  76. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/random.py +0 -0
  77. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/specialized_fns.py +0 -0
  78. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/state_dict.py +0 -0
  79. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/tree_util.py +0 -0
  80. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/types.py +0 -0
  81. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/util.py +0 -0
  82. {haliax-1.4.dev325 → haliax-1.4.dev327}/src/haliax/wrap.py +0 -0
  83. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/core_test.py +0 -0
  84. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_attention.py +0 -0
  85. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_axis.py +0 -0
  86. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_conv.py +0 -0
  87. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_debug.py +0 -0
  88. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_dot.py +0 -0
  89. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_einsum.py +0 -0
  90. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_fp8.py +0 -0
  91. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_hof.py +0 -0
  92. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_nn.py +0 -0
  93. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_ops.py +0 -0
  94. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_parsing.py +0 -0
  95. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_partitioning.py +0 -0
  96. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_pool.py +0 -0
  97. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_random.py +0 -0
  98. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_rearrange.py +0 -0
  99. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_scan.py +0 -0
  100. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_specialized_fns.py +0 -0
  101. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_state_dict.py +0 -0
  102. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_tree_util.py +0 -0
  103. {haliax-1.4.dev325 → haliax-1.4.dev327}/tests/test_utils.py +0 -0
@@ -1,11 +1,12 @@
1
- Metadata-Version: 2.3
1
+ Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev325
3
+ Version: 1.4.dev327
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/
7
7
  Project-URL: Documentation, https://haliax.readthedocs.io/en/latest/
8
8
  Author-email: David Hall <dlwh@cs.stanford.edu>
9
+ License-File: LICENSE
9
10
  Classifier: Development Status :: 4 - Beta
10
11
  Classifier: Intended Audience :: Science/Research
11
12
  Classifier: License :: OSI Approved :: Apache Software License
@@ -0,0 +1 @@
1
+ __version__ = "1.4.dev327"
@@ -34,6 +34,8 @@ from .axis import (
34
34
  dslice,
35
35
  eliminate_axes,
36
36
  make_axes,
37
+ replace_axis,
38
+ resolve_axis,
37
39
  selects_axis,
38
40
  )
39
41
  from .core import (
@@ -1060,6 +1062,8 @@ __all__ = [
1060
1062
  "stack",
1061
1063
  "concatenate",
1062
1064
  "eliminate_axes",
1065
+ "resolve_axis",
1066
+ "replace_axis",
1063
1067
  "selects_axis",
1064
1068
  "concat_axes",
1065
1069
  "concat_axis_specs",
@@ -15,7 +15,8 @@ from jaxtyping import PyTree
15
15
 
16
16
  import haliax.partitioning as partitioning
17
17
  from haliax._src.util import index_where
18
- from haliax.core import NamedArray, named
18
+ from haliax.axis import Axis
19
+ from haliax.core import NamedArray, flatten_axes, named
19
20
  from haliax.jax_utils import is_jax_array_like, is_scalarish
20
21
 
21
22
 
@@ -390,24 +391,22 @@ def flatten_linear_layers(tree: T) -> T:
390
391
  weight = layer.weight
391
392
  bias = layer.bias
392
393
 
394
+ new_Out: Axis = flatten_axes(layer.Out, "__OUT__")
395
+ new_In: Axis = flatten_axes(layer.In, "__IN__")
396
+
393
397
  if weight.array is not None:
394
398
  out_first = layer.out_first
395
- weight = weight.flatten_axes(layer.Out, "__OUT__").flatten_axes(layer.In, "__IN__")
399
+ weight = weight.flatten_axes(layer.Out, new_Out).flatten_axes(layer.In, new_In)
396
400
 
397
401
  if out_first:
398
402
  weight = weight.rearrange((..., "__OUT__", "__IN__"))
399
403
  else:
400
404
  weight = weight.rearrange((..., "__IN__", "__OUT__"))
401
405
 
402
- if bias is not None:
403
- bias = bias.flatten_axes(layer.Out, "__OUT__")
404
-
405
- In = weight.resolve_axis("__IN__")
406
- Out = weight.resolve_axis("__OUT__")
406
+ if isinstance(bias, NamedArray):
407
+ bias = bias.flatten_axes(layer.Out, new_Out)
407
408
 
408
- return dataclasses.replace(layer, weight=weight, bias=bias, In=In, Out=Out) # type: ignore
409
- else:
410
- return layer
409
+ return dataclasses.replace(layer, weight=weight, bias=bias, In=new_In, Out=new_Out) # type: ignore
411
410
 
412
411
  return jax.tree.map(_flatten_linear, tree, is_leaf=lambda x: isinstance(x, Linear))
413
412
 
@@ -438,7 +437,7 @@ def unflatten_linear_layers(template: T, tree_with_flattened_linears: T) -> T:
438
437
  weight = weight.unflatten_axis("__OUT__", template.Out).unflatten_axis("__IN__", template.In)
439
438
  weight = weight.rearrange(template.weight.axes)
440
439
 
441
- if bias is not None:
440
+ if isinstance(bias, NamedArray):
442
441
  bias = bias.unflatten_axis("__OUT__", template.Out)
443
442
  assert template.bias is not None, "Flattened bias but template has no bias"
444
443
  bias = bias.rearrange(template.bias.axes)
@@ -394,6 +394,52 @@ def axis_size(ax: AxisSpec) -> int:
394
394
  return prod(axis.size for axis in ensure_tuple(ax)) # type: ignore
395
395
 
396
396
 
397
+ @typing.overload
398
+ def resolve_axis(axis_spec: AxisSpec, axis_selection: AxisSelector) -> Axis:
399
+ ...
400
+
401
+
402
+ @typing.overload
403
+ def resolve_axis(axis_spec: AxisSpec, axis_selection: AxisSelection) -> AxisSpec:
404
+ ...
405
+
406
+
407
+ def resolve_axis(axis_spec: AxisSpec, axis_selection: AxisSelection) -> AxisSpec:
408
+ """
409
+ Returns the axis or axes in axis_spec that match the name of axis_selection.
410
+
411
+ If axis_selection is a str or axis, returns a single Axis. If it is a sequence, returns a sequence of Axes.
412
+
413
+ If an axis is present with a different size, raises ValueError.
414
+ """
415
+ ax: Axis
416
+ if isinstance(axis_selection, str | Axis):
417
+ name = axis_name(axis_selection)
418
+ for ax in ensure_tuple(axis_spec):
419
+ if axis_name(ax) == name:
420
+ if isinstance(axis_selection, Axis) and axis_size(ax) != axis_size(axis_selection):
421
+ raise ValueError(f"Axis {name} has different sizes in {axis_spec} and {axis_selection}")
422
+ return ax
423
+ raise ValueError(f"Axis {name} not found in {axis_spec}")
424
+ else:
425
+ as_map = axis_spec_to_shape_dict(axis_spec)
426
+ out: list[Axis] = []
427
+
428
+ for ax in ensure_tuple(axis_selection): # type: ignore
429
+ name = axis_name(ax)
430
+ if name not in as_map:
431
+ raise ValueError(f"Axis {name} not found in {axis_spec}")
432
+ if isinstance(ax, Axis):
433
+ if as_map[name] != ax.size: # type: ignore
434
+ raise ValueError(f"Axis {name} has different sizes in {axis_spec} and {axis_selection}")
435
+ else:
436
+ out.append(ax)
437
+ else:
438
+ out.append(Axis(name, as_map[name]))
439
+
440
+ return tuple(out)
441
+
442
+
397
443
  class dslice(eqx.Module):
398
444
  """
399
445
  Dynamic slice, comprising a (start, length) pair. Also aliased as ds.
@@ -576,6 +622,7 @@ __all__ = [
576
622
  "is_axis_compatible",
577
623
  "overlapping_axes",
578
624
  "replace_axis",
625
+ "resolve_axis",
579
626
  "selects_axis",
580
627
  "union_axes",
581
628
  "without_axes",
@@ -1143,12 +1143,65 @@ def rename(array: NamedArray, renames: Mapping[AxisSelector, AxisSelector]) -> N
1143
1143
  return NamedArray(array.array, new_axes)
1144
1144
 
1145
1145
 
1146
+ @typing.overload
1147
+ def flatten_axes(axis: Axis, old_axes: Axis, new_axis: AxisSelector) -> Axis:
1148
+ pass
1149
+
1150
+
1151
+ @typing.overload
1152
+ def flatten_axes(axis: AxisSpec, new_axis: AxisSelector) -> Axis:
1153
+ pass
1154
+
1155
+
1156
+ @typing.overload
1157
+ def flatten_axes(axis: AxisSpec, old_axes: AxisSelection, new_axis: AxisSelector) -> AxisSpec:
1158
+ pass
1159
+
1160
+
1161
+ @typing.overload
1146
1162
  def flatten_axes(array: NamedArray, old_axes: AxisSelection, new_axis: AxisSelector) -> NamedArray:
1163
+ pass
1164
+
1165
+
1166
+ def flatten_axes( # type: ignore
1167
+ # array: NamedArray | AxisSpec, old_axes: AxisSelection, new_axis: AxisSelector
1168
+ *args,
1169
+ **kwargs,
1170
+ ) -> NamedArray | AxisSpec:
1147
1171
  """
1148
1172
  Merge a sequence of axes into a single axis. The new axis must have the same size as the product of the old axes.
1149
1173
 
1150
- The new axis is always inserted starting at the index of the first old axis in theunderlying array.
1174
+ The new axis is always inserted starting at the index of the first old axis in the underlying array.
1175
+
1176
+ This function can be used in two ways:
1177
+
1178
+ * `flatten_axes(array, old_axes, new_axis)`: merge the old axes of the array into a new axis
1179
+ * `flatten_axes(axes, old_axes, new_axis)`: merge the old axes into a new axis
1151
1180
  """
1181
+
1182
+ if len(args) + len(kwargs) == 2:
1183
+ return _simple_flatten(*args, **kwargs)
1184
+ else:
1185
+ return _full_flatten(*args, **kwargs)
1186
+
1187
+
1188
+ def _simple_flatten(axis: AxisSpec, new_axis: AxisSelector) -> Axis:
1189
+ size = haliax.axis_size(axis)
1190
+ if isinstance(new_axis, Axis):
1191
+ if new_axis.size != size:
1192
+ raise ValueError(f"Cannot merge {axis} into {new_axis}: size mismatch")
1193
+ return new_axis
1194
+
1195
+ assert isinstance(new_axis, str)
1196
+ return Axis(new_axis, size)
1197
+
1198
+
1199
+ def _full_flatten(
1200
+ array: NamedArray | AxisSpec, old_axes: AxisSelection, new_axis: AxisSelector
1201
+ ) -> NamedArray | AxisSpec:
1202
+ if isinstance(array, str | Axis | Sequence):
1203
+ return _flatten_axis_spec(array, old_axes, new_axis)
1204
+
1152
1205
  old_axes = ensure_tuple(old_axes)
1153
1206
  old_axes = array.resolve_axis(old_axes)
1154
1207
  total_axis_size = haliax.axis_size(old_axes)
@@ -1187,6 +1240,42 @@ def flatten_axes(array: NamedArray, old_axes: AxisSelection, new_axis: AxisSelec
1187
1240
  return NamedArray(raw_array, tuple(new_axes))
1188
1241
 
1189
1242
 
1243
+ def _flatten_axis_spec(axes: AxisSpec, old_axes: AxisSelection, new_axis: AxisSelector) -> AxisSpec:
1244
+ axes = ensure_tuple(axes)
1245
+ old_axes = ensure_tuple(old_axes)
1246
+ old_axes = haliax.axis.resolve_axis(axes, old_axes)
1247
+ total_axis_size = haliax.axis_size(old_axes)
1248
+
1249
+ if isinstance(new_axis, Axis):
1250
+ if new_axis.size != total_axis_size:
1251
+ raise ValueError(f"Cannot merge {old_axes} into {new_axis}: size mismatch")
1252
+ else:
1253
+ assert isinstance(new_axis, str)
1254
+ new_axis = Axis(new_axis, total_axis_size)
1255
+
1256
+ if len(old_axes) == 0: # type: ignore
1257
+ return (new_axis,) + axes
1258
+
1259
+ # ensure that the old_axes are contiguous
1260
+ # we basically ensure that the old_axes occur after the index of the first old_axis
1261
+ intermediate_axes: List[Axis] = []
1262
+ new_axes: List[Axis] = []
1263
+ index_of_first_old_axis = None
1264
+ for i, ax in enumerate(axes):
1265
+ if ax in old_axes: # type: ignore
1266
+ if index_of_first_old_axis is None:
1267
+ index_of_first_old_axis = i
1268
+ intermediate_axes.extend(old_axes) # type: ignore
1269
+ new_axes.append(new_axis)
1270
+ else:
1271
+ continue
1272
+ else:
1273
+ intermediate_axes.append(ax)
1274
+ new_axes.append(ax)
1275
+
1276
+ return tuple(new_axes)
1277
+
1278
+
1190
1279
  def unflatten_axis(array: NamedArray, axis: AxisSelector, new_axes: AxisSpec) -> NamedArray:
1191
1280
  """
1192
1281
  Split an axis into a sequence of axes. The old axis must have the same size as the product of the new axes.
@@ -8,7 +8,7 @@ from typing import Callable, ContextManager, Mapping, Optional, ParamSpec, Seque
8
8
 
9
9
  import equinox as eqx
10
10
  import jax
11
- from equinox import module_update_wrapper
11
+ from equinox import is_array, module_update_wrapper
12
12
  from jax.lax import with_sharding_constraint
13
13
  from jax.sharding import Mesh, NamedSharding, PartitionSpec, SingleDeviceSharding
14
14
  from jaxtyping import PyTree
@@ -20,7 +20,7 @@ from .axis import Axis, AxisSelection, AxisSelector
20
20
  from .core import NamedArray
21
21
  from .jax_utils import Static, is_in_jit, is_jax_array_like, is_on_mac_metal
22
22
  from .tree_util import hashable_combine, hashable_partition
23
- from .util import StringHolderEnum, ensure_tuple, is_named_array
23
+ from .util import StringHolderEnum, ensure_tuple
24
24
 
25
25
 
26
26
  PhysicalAxisSpec = Union[(str), Sequence[str]]
@@ -274,7 +274,7 @@ class _NamedJitWrapper(eqx.Module):
274
274
  if out_axis_resources is None:
275
275
  out_axis_resources = axis_resources
276
276
 
277
- dynamic_argspec, static_argspec = hashable_partition((args, kwargs), is_jax_array_like)
277
+ dynamic_argspec, static_argspec = hashable_partition((args, kwargs), is_array)
278
278
  dynamic = (self._dynamic_fun, dynamic_argspec)
279
279
 
280
280
  donate_args = self._donate_args
@@ -436,7 +436,7 @@ def named_jit(
436
436
  **pjit_args,
437
437
  )
438
438
 
439
- dynamic_fun, static_fun = hashable_partition(fn, is_jax_array_like)
439
+ dynamic_fun, static_fun = hashable_partition(fn, is_array)
440
440
 
441
441
  wrapper = _NamedJitWrapper(
442
442
  fn,
@@ -514,7 +514,7 @@ def _named_pjit_cache(fun_names, **jitkwargs) -> WrappedCallable:
514
514
  fun = hashable_combine(dynamic_fun, static_fun)
515
515
  args, kwargs = hashable_combine(dynamic_spec, static_spec)
516
516
  out = fun(*args, **kwargs)
517
- out_dynamic, out_static = hashable_partition(out, is_jax_array_like)
517
+ out_dynamic, out_static = hashable_partition(out, is_array)
518
518
  return out_dynamic, Static(out_static)
519
519
 
520
520
  fun_name, fun_qualname = fun_names
@@ -543,7 +543,7 @@ def _cached_filter_eval_shape(fun, *args, **kwargs):
543
543
  eval_shape is surprisingly expensive, so we cache it. We use this for named_pjit for evaluating resource partitions
544
544
  of the output.
545
545
  """
546
- dynamic, static = hashable_partition((fun, args, kwargs), is_jax_array_like)
546
+ dynamic, static = hashable_partition((fun, args, kwargs), is_array)
547
547
  if static not in _eval_shape_cache:
548
548
  _eval_shape_cache[static] = eqx.filter_eval_shape(fun, *args, **kwargs)
549
549
 
@@ -10,7 +10,6 @@ from typing import Optional, Protocol, TypeVar
10
10
  import equinox as eqx
11
11
  import jax
12
12
  from jax import numpy as jnp
13
- from jax._src.tree_util import BuiltInKeyEntry
14
13
  from jax.tree_util import DictKey, FlattenedIndexKey, GetAttrKey, SequenceKey
15
14
  from jax.typing import DTypeLike
16
15
 
@@ -253,7 +252,7 @@ def _matches_target_fp8(key_path, config: Fp8Config) -> bool:
253
252
  return re.match(config.targets, key_path_str) is not None
254
253
 
255
254
 
256
- def _key_path_to_str(key_path: tuple[BuiltInKeyEntry, ...]) -> str:
255
+ def _key_path_to_str(key_path: tuple) -> str:
257
256
  out = ""
258
257
  for k in key_path:
259
258
  match k:
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev325"
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