haliax 1.4.dev380__tar.gz → 1.4.dev382__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.dev380 → haliax-1.4.dev382}/PKG-INFO +3 -3
  2. {haliax-1.4.dev380 → haliax-1.4.dev382}/README.md +2 -2
  3. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/indexing.md +6 -1
  4. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/nn.md +1 -1
  5. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/state-dict.md +1 -1
  6. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/typing.md +1 -1
  7. haliax-1.4.dev382/src/haliax/__about__.py +1 -0
  8. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/axis.py +11 -6
  9. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/core.py +23 -62
  10. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/nn/conv.py +1 -1
  11. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/nn/linear.py +1 -1
  12. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/core_test.py +29 -4
  13. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_rearrange.py +1 -1
  14. haliax-1.4.dev380/src/haliax/__about__.py +0 -1
  15. {haliax-1.4.dev380 → haliax-1.4.dev382}/.coveragerc +0 -0
  16. {haliax-1.4.dev380 → haliax-1.4.dev382}/.flake8 +0 -0
  17. {haliax-1.4.dev380 → haliax-1.4.dev382}/.github/workflows/publish_dev.yaml +0 -0
  18. {haliax-1.4.dev380 → haliax-1.4.dev382}/.github/workflows/run_pre_commit.yaml +0 -0
  19. {haliax-1.4.dev380 → haliax-1.4.dev382}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  20. {haliax-1.4.dev380 → haliax-1.4.dev382}/.github/workflows/run_tests.yaml +0 -0
  21. {haliax-1.4.dev380 → haliax-1.4.dev382}/.gitignore +0 -0
  22. {haliax-1.4.dev380 → haliax-1.4.dev382}/.playbooks/add-types.md +0 -0
  23. {haliax-1.4.dev380 → haliax-1.4.dev382}/.pre-commit-config.yaml +0 -0
  24. {haliax-1.4.dev380 → haliax-1.4.dev382}/.readthedocs.yaml +0 -0
  25. {haliax-1.4.dev380 → haliax-1.4.dev382}/AGENTS.md +0 -0
  26. {haliax-1.4.dev380 → haliax-1.4.dev382}/CONTRIBUTING.md +0 -0
  27. {haliax-1.4.dev380 → haliax-1.4.dev382}/LICENSE +0 -0
  28. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/api.md +0 -0
  29. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/broadcasting.md +0 -0
  30. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/cheatsheet.md +0 -0
  31. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/css/material.css +0 -0
  32. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/css/mkdocstrings.css +0 -0
  33. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/faq.md +0 -0
  34. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/figures/data_parallel_mesh.png +0 -0
  35. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  36. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/figures/device_mesh_1d.png +0 -0
  37. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/figures/device_mesh_1d_zero.png +0 -0
  38. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/figures/device_mesh_2d.png +0 -0
  39. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  40. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  41. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  42. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  43. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/figures/device_mesh_2d_zero.png +0 -0
  44. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/fp8.md +0 -0
  45. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/index.md +0 -0
  46. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/matmul.md +0 -0
  47. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/partitioning.md +0 -0
  48. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/rearrange.ipynb +0 -0
  49. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/rearrange.md +0 -0
  50. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/requirements.txt +0 -0
  51. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/scan.md +0 -0
  52. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/tutorial.md +0 -0
  53. {haliax-1.4.dev380 → haliax-1.4.dev382}/docs/vmap.md +0 -0
  54. {haliax-1.4.dev380 → haliax-1.4.dev382}/mkdocs.yml +0 -0
  55. {haliax-1.4.dev380 → haliax-1.4.dev382}/pyproject.toml +0 -0
  56. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/__init__.py +0 -0
  57. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/_src/__init__.py +0 -0
  58. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/_src/compile_utils.py +0 -0
  59. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/_src/dot.py +0 -0
  60. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/_src/einsum.py +0 -0
  61. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/_src/fp8.py +0 -0
  62. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/_src/parsing.py +0 -0
  63. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/_src/rearrange.py +0 -0
  64. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/_src/scan.py +0 -0
  65. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/_src/state_dict.py +0 -0
  66. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/_src/util.py +0 -0
  67. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/debug.py +0 -0
  68. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/haxtyping.py +0 -0
  69. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/hof.py +0 -0
  70. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/jax_utils.py +0 -0
  71. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/nn/__init__.py +0 -0
  72. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/nn/activations.py +0 -0
  73. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/nn/attention.py +0 -0
  74. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/nn/dropout.py +0 -0
  75. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/nn/embedding.py +0 -0
  76. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/nn/loss.py +0 -0
  77. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/nn/mlp.py +0 -0
  78. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/nn/normalization.py +0 -0
  79. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/nn/pool.py +0 -0
  80. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/nn/scan.py +0 -0
  81. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/ops.py +0 -0
  82. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/partitioning.py +0 -0
  83. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/quantization.py +0 -0
  84. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/random.py +0 -0
  85. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/specialized_fns.py +0 -0
  86. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/state_dict.py +0 -0
  87. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/tree_util.py +0 -0
  88. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/types.py +0 -0
  89. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/util.py +0 -0
  90. {haliax-1.4.dev380 → haliax-1.4.dev382}/src/haliax/wrap.py +0 -0
  91. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_attention.py +0 -0
  92. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_axis.py +0 -0
  93. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_conv.py +0 -0
  94. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_debug.py +0 -0
  95. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_dot.py +0 -0
  96. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_dtype_typing.py +0 -0
  97. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_einsum.py +0 -0
  98. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_fp8.py +0 -0
  99. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_hof.py +0 -0
  100. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_int8.py +0 -0
  101. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_namedarray_typing.py +0 -0
  102. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_nn.py +0 -0
  103. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_ops.py +0 -0
  104. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_parsing.py +0 -0
  105. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_partitioning.py +0 -0
  106. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_pool.py +0 -0
  107. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_random.py +0 -0
  108. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_scan.py +0 -0
  109. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_scatter_gather.py +0 -0
  110. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_specialized_fns.py +0 -0
  111. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_state_dict.py +0 -0
  112. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_tree_util.py +0 -0
  113. {haliax-1.4.dev380 → haliax-1.4.dev382}/tests/test_utils.py +0 -0
  114. {haliax-1.4.dev380 → haliax-1.4.dev382}/uv.lock +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev380
3
+ Version: 1.4.dev382
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/
@@ -33,14 +33,14 @@ Description-Content-Type: text/markdown
33
33
  <a href="">
34
34
  <img alt="License" src="https://img.shields.io/github/license/stanford-crfm/haliax?color=blue" />
35
35
  </a>
36
- <a href="https://https://pypi.org/project/haliax/">
36
+ <a href="https://pypi.org/project/haliax/">
37
37
  <img alt="PyPI" src="https://img.shields.io/pypi/v/haliax?color=blue" />
38
38
  </a>
39
39
 
40
40
  > *Though you don’t seem to be much for listening, it’s best to be careful. If you managed to catch hold of even just a piece of my name, you’d have all manner of power over me.*<br/>
41
41
  > — Patrick Rothfuss, *The Name of the Wind*
42
42
 
43
- Haliax is a [JAX](https:://github.com/google/jax) library for building neural networks with named tensors, in the tradition of Alexander Rush's [Tensor Considered Harmful](https://nlp.seas.harvard.edu/NamedTensor).
43
+ Haliax is a [JAX](https://github.com/google/jax) library for building neural networks with named tensors, in the tradition of Alexander Rush's [Tensor Considered Harmful](https://nlp.seas.harvard.edu/NamedTensor).
44
44
  Named tensors improve the **legibility** and **compositionality** of tensor programs by using named axes instead of positional indices
45
45
  as typically used in NumPy, PyTorch, etc.
46
46
 
@@ -10,14 +10,14 @@
10
10
  <a href="">
11
11
  <img alt="License" src="https://img.shields.io/github/license/stanford-crfm/haliax?color=blue" />
12
12
  </a>
13
- <a href="https://https://pypi.org/project/haliax/">
13
+ <a href="https://pypi.org/project/haliax/">
14
14
  <img alt="PyPI" src="https://img.shields.io/pypi/v/haliax?color=blue" />
15
15
  </a>
16
16
 
17
17
  > *Though you don’t seem to be much for listening, it’s best to be careful. If you managed to catch hold of even just a piece of my name, you’d have all manner of power over me.*<br/>
18
18
  > — Patrick Rothfuss, *The Name of the Wind*
19
19
 
20
- Haliax is a [JAX](https:://github.com/google/jax) library for building neural networks with named tensors, in the tradition of Alexander Rush's [Tensor Considered Harmful](https://nlp.seas.harvard.edu/NamedTensor).
20
+ Haliax is a [JAX](https://github.com/google/jax) library for building neural networks with named tensors, in the tradition of Alexander Rush's [Tensor Considered Harmful](https://nlp.seas.harvard.edu/NamedTensor).
21
21
  Named tensors improve the **legibility** and **compositionality** of tensor programs by using named axes instead of positional indices
22
22
  as typically used in NumPy, PyTorch, etc.
23
23
 
@@ -1,7 +1,7 @@
1
1
  # Indexing and Slicing
2
2
 
3
3
  Haliax supports Numpy-style indexing, including so-called [Advanced Indexing](https://numpy.org/doc/stable/user/basics.indexing.html#advanced-indexing),
4
- though the syntax is necessarily different. Most forms of indexing are supporting, except we don't support indexing with
4
+ though the syntax is necessarily different. Most forms of indexing are supported, except we don't support indexing with
5
5
  booleans right now. (JAX doesn't support indexing with non-constant bool arrays anyway,
6
6
  so I don't think it's worth the effort to implement it in Haliax.)
7
7
 
@@ -147,6 +147,11 @@ def f(x, slice_size: int):
147
147
  f(q, 2)
148
148
  ```
149
149
 
150
+ When indexing with ``dslice`` the slice is gathered starting at ``start`` for
151
+ ``size`` elements. Reads beyond the end of the array return the ``fill_value``
152
+ (0 by default). When used with ``at`` updates, any writes outside the bounds of
153
+ the array are dropped. These semantics match JAX's scatter/gather behavior.
154
+
150
155
  For convenience/brevity, `dslice` is aliased as `ds`. In addition, we also expose `dblock`, which is a convenience
151
156
  function for computing the start and size of a slice given a block index and the size of the slice. Thus, the above
152
157
  example can be written as follows:
@@ -6,7 +6,7 @@
6
6
  Haliax provides a small number of neural network modules that are compatible with Equinox, though
7
7
  they naturally all use [haliax.NamedArray][]. (We welcome PRs for more modules! Nothing too exotic though.)
8
8
 
9
- The most interesting of these modules is [haliax.nn.Stacked][], which allows you to create homogenous "stacks"
9
+ The most interesting of these modules is [haliax.nn.Stacked][], which allows you to create homogeneous "stacks"
10
10
  of the same module (e.g. transformer blocks), which is a common pattern in deep learning.
11
11
 
12
12
  ### Linear
@@ -226,7 +226,7 @@ any Axis members to match the new shape.
226
226
  ::: haliax.state_dict.save_state_dict
227
227
  ::: haliax.state_dict.load_state_dict
228
228
 
229
- ### Converting betweewn State Dicts and Modules
229
+ ### Converting between State Dicts and Modules
230
230
 
231
231
  ::: haliax.state_dict.from_state_dict
232
232
  ::: haliax.state_dict.to_state_dict
@@ -1,4 +1,4 @@
1
- from haliax import NamedArrayfrom haliax import NamedArray
1
+ from haliax import NamedArray
2
2
 
3
3
  # NamedArray Type Annotations
4
4
 
@@ -0,0 +1 @@
1
+ __version__ = "1.4.dev382"
@@ -598,14 +598,19 @@ def resolve_axis(axis_spec: AxisSpec, axis_selection: AxisSelection) -> AxisSpec
598
598
 
599
599
 
600
600
  class dslice(eqx.Module):
601
- """
602
- Dynamic slice, comprising a (start, length) pair. Also aliased as ds.
601
+ """Dynamic slice, comprising a (start, length) pair. Also aliased as ``ds``.
602
+
603
+ NumPy-style slices like ``a[i:i+16]`` don't work inside :func:`jax.jit`, because
604
+ JAX requires slice bounds to be static. ``dslice`` works around this by
605
+ separating the dynamic ``start`` from the static ``size`` so that you can
606
+ write ``a[dslice(i, 16)]`` or simply ``a[ds(i, 16)]``.
603
607
 
604
- Normal numpy-isms like a[i:i+16] don't work in Jax inside jit, because slice doesn't like tracers and JAX
605
- can't see that the slice is constant. This is a workaround that lets you do a[dslice(i, 16)] or even a[ds(i, 16)]
606
- instead.
608
+ When used in indexing or ``at`` updates, ``dslice`` behaves like a gather of
609
+ ``size`` elements starting at ``start``. Reads beyond the end of the array are
610
+ filled with a value (0 by default) and writes outside the array bounds are
611
+ dropped, matching JAX's default scatter/gather semantics.
607
612
 
608
- This class's name is taken from [jax.experimental.pallas.dslice][].
613
+ This class's name is taken from :mod:`jax.experimental.pallas`.
609
614
  """
610
615
 
611
616
  start: int
@@ -1163,6 +1163,11 @@ def index(array: NamedArray, slices: Mapping[AxisSelector, NamedIndex]) -> Named
1163
1163
  you might use `array[{"batch": slice(0, 10)}]` or `array["batch", 0:10]` to select the first 10 elements
1164
1164
  of the 'batch' axis.
1165
1165
 
1166
+ When indexing with a ``dslice`` the slice is gathered starting at the given
1167
+ ``start`` for ``size`` elements. Values read past the end of the array are
1168
+ filled with the ``fill_value`` (defaults to ``0``), and writes outside the
1169
+ bounds are dropped.
1170
+
1166
1171
  See Also:
1167
1172
  * [haliax.NamedArray.at][] for a functional equivalent of in-place array modifications.
1168
1173
 
@@ -1172,8 +1177,12 @@ def index(array: NamedArray, slices: Mapping[AxisSelector, NamedIndex]) -> Named
1172
1177
  """
1173
1178
  # indices where we have array args
1174
1179
  new_axes, ordered_slices = _compute_new_axes_and_slices_for_index(array, slices)
1175
- sliced, ordered_slices = _handle_dynamic_slices(array.array, ordered_slices)
1176
- sliced = sliced[tuple(ordered_slices)]
1180
+ needs_fill = any(isinstance(s, dslice) or is_pallas_dslice(s) for s in slices.values())
1181
+ ordered_slices = [s.array if isinstance(s, NamedArray) else s for s in ordered_slices]
1182
+ if needs_fill:
1183
+ sliced = array.array.at[tuple(ordered_slices)].get(mode="fill", fill_value=0)
1184
+ else:
1185
+ sliced = array.array[tuple(ordered_slices)]
1177
1186
 
1178
1187
  return haliax.named(sliced, new_axes)
1179
1188
 
@@ -1190,9 +1199,19 @@ def _compute_new_axes_and_slices_for_index(
1190
1199
  axis_index = array.axis_indices(axis)
1191
1200
  if axis_index is None:
1192
1201
  raise ValueError(f"axis {axis} not found in {array}")
1193
- if isinstance(slice_, py_slice) or isinstance(slice_, dslice) or is_pallas_dslice(slice_):
1202
+ if isinstance(slice_, py_slice):
1194
1203
  ordered_slices[axis_index] = slice_
1195
1204
  kept_axes[axis_index] = True
1205
+ elif is_pallas_dslice(slice_) or isinstance(slice_, dslice):
1206
+ # we want to treat this like a slice, but the closest approximation is an arange, but those have to broadcast.
1207
+ # So instead we make a named array arange with the same name as the axis.
1208
+ start = slice_.start
1209
+ size = slice_.size
1210
+ orig_axis = array.axes[axis_index]
1211
+ ordered_slices[axis_index] = haliax.arange(orig_axis.resize(size), start=start)
1212
+ kept_axes[axis_index] = False
1213
+ array_slice_indices.append(axis_index)
1214
+ index_axis_names.add(orig_axis.name)
1196
1215
  elif isinstance(slice_, int):
1197
1216
  ordered_slices[axis_index] = slice_
1198
1217
  kept_axes[axis_index] = False
@@ -1292,34 +1311,6 @@ def _compute_new_axes_and_slices_for_index(
1292
1311
  return new_axes, ordered_slices
1293
1312
 
1294
1313
 
1295
- def _handle_dynamic_slices(array: jnp.ndarray, slices):
1296
- """
1297
- Helper function to handle dynamic slices in the array. These have to be handled with jax.lax.dynamic_slice,
1298
- which is for when the start index is not known at compile time. (Sizes must always be known at compile time.)
1299
-
1300
- Notes:
1301
- **MUTATES `slices` IN PLACE**
1302
-
1303
- Returns:
1304
- array.array: the sliced array
1305
-
1306
- """
1307
- indices_for_dslice = [0] * array.ndim
1308
- lengths_for_dslice = list(array.shape)
1309
- dslice_indices = []
1310
- need_to_slice = False
1311
- for axis_index, slice_ in enumerate(slices):
1312
- if isinstance(slice_, dslice) or is_pallas_dslice(slice_):
1313
- dslice_indices.append(axis_index)
1314
- indices_for_dslice[axis_index] = slice_.start
1315
- lengths_for_dslice[axis_index] = slice_.size
1316
- need_to_slice = True
1317
- if need_to_slice:
1318
- array = jax.lax.dynamic_slice(array, indices_for_dslice, lengths_for_dslice)
1319
- for i in dslice_indices:
1320
- slices[i] = py_slice(None, None, None)
1321
- return array, slices
1322
-
1323
1314
 
1324
1315
  def split(a: NamedArray, axis: AxisSelector, new_axes: Sequence[Axis]) -> Sequence[NamedArray]:
1325
1316
  """
@@ -2056,38 +2047,8 @@ class _NamedIndexUpdateRef:
2056
2047
 
2057
2048
  def _raw_indices_for_at(array, indexes):
2058
2049
  sliced_axes, ordered_slices = _compute_new_axes_and_slices_for_index(array, indexes)
2059
- del sliced_axes
2060
- # this isn't the fastest (it does the _compute_new_axes_and_slices_for_index twice)
2061
- # but it's easy
2062
2050
  _sliced = index(array, indexes)
2063
- # we have to handle dslices differently than for normal indexing, because we can't use
2064
- # extra dynamic_slices...
2065
- # we'd like to just replace these with iota, but we have account for broadcasting semantics
2066
- # for the other arrays
2067
- dslice_sizes = tuple(x.size for x in ordered_slices if isinstance(x, dslice) or is_pallas_dslice(x)) # type: ignore
2068
- current_array_slice_shape = next((x.shape for x in ordered_slices if is_jax_array_like(x)), None) # type: ignore
2069
- dims_to_expand = list(range(len(dslice_sizes)))
2070
- if current_array_slice_shape is not None:
2071
- iota_shape = dslice_sizes + current_array_slice_shape
2072
- else:
2073
- iota_shape = dslice_sizes
2074
-
2075
- def iota_for_dslice(dslice, cur_dynamic_slice):
2076
- return jax.lax.broadcasted_iota(int, iota_shape, cur_dynamic_slice) + dslice.start
2077
-
2078
- if len(dslice_sizes) > 0:
2079
- cur_dynamic_slice = 0
2080
- for i in range(len(ordered_slices)):
2081
- if isinstance(ordered_slices[i], dslice) or is_pallas_dslice(ordered_slices[i]):
2082
- ordered_slices[i] = iota_for_dslice(ordered_slices[i], cur_dynamic_slice)
2083
- cur_dynamic_slice += 1
2084
- elif is_jax_array_like(ordered_slices[i]):
2085
- # prepend array slices with one 1 for each dynamic slice
2086
- ordered_slices[i] = jnp.expand_dims(ordered_slices[i], axis=dims_to_expand)
2087
-
2088
- assert cur_dynamic_slice == len(dslice_sizes)
2089
-
2090
- # ok the ordered slices are now correct
2051
+ ordered_slices = [s.array if isinstance(s, NamedArray) else s for s in ordered_slices]
2091
2052
  return ordered_slices, _sliced.axes
2092
2053
 
2093
2054
 
@@ -213,7 +213,7 @@ class Conv(_ConvBase):
213
213
  return x
214
214
 
215
215
  def _do_conv(self, inputs):
216
- # _do_conv expects there ot be a single __batch__ dimension
216
+ # _do_conv expects there to be a single __batch__ dimension
217
217
  output_axes = _compute_output_axes(inputs, "__batch__", self.In, self.Out)
218
218
 
219
219
  batch_index = _index_of_name(inputs.axes, "__batch__")
@@ -143,7 +143,7 @@ class MoELinear(eqx.Module):
143
143
  Experts: AxisSpec = eqx.field(static=True)
144
144
  In: Axis = eqx.field(static=True)
145
145
  Out: Axis = eqx.field(static=True)
146
- # TODO: support quanitization for ragged_dot?
146
+ # TODO: support quantization for ragged_dot?
147
147
  # dot_general: DotGeneralOp = eqx.field(default_factory=DotGeneralOp.default)
148
148
 
149
149
  use_gmm: bool = eqx.field(static=True)
@@ -586,11 +586,36 @@ def test_slice_nd_dslice():
586
586
  named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
587
587
  from haliax import ds
588
588
 
589
- assert jnp.all(jnp.equal(named1["H", ds(0, 5), "D", ds(3, 7)].array, named1.array[0:5, :, 3:10]))
589
+ ref = jnp.take(named1.array, jnp.arange(0, 5), axis=0, mode="fill", fill_value=0)
590
+ ref = jnp.take(ref, jnp.arange(3, 10), axis=2, mode="fill", fill_value=0)
591
+ ref = jnp.transpose(ref, (0, 2, 1))
592
+ assert jnp.all(jnp.equal(named1["H", ds(0, 5), "D", ds(3, 7)].array, ref))
590
593
  # test mixed normal and dslice
591
- assert jnp.all(jnp.equal(named1["H", ds(1, 5), "D", 3:7].array, named1.array[1:6, :, 3:7]))
592
- assert jnp.all(jnp.equal(named1["H", ds(2, 5), "D", 3].array, named1.array[2:7, :, 3]))
593
- assert jnp.all(jnp.equal(named1["H", ds(3, 5), "D", 3:10:2].array, named1.array[3:8, :, 3:10:2]))
594
+ ref = jnp.take(named1.array, jnp.arange(1, 6), axis=0, mode="fill", fill_value=0)
595
+ ref = jnp.take(ref, jnp.arange(3, 7), axis=2, mode="fill", fill_value=0)
596
+ assert jnp.all(jnp.equal(named1["H", ds(1, 5), "D", 3:7].array, ref))
597
+ ref = jnp.take(named1.array, jnp.arange(2, 7), axis=0, mode="fill", fill_value=0)
598
+ ref = jnp.take(ref, jnp.arange(3, 4), axis=2, mode="fill", fill_value=0)
599
+ ref = ref.squeeze(2)
600
+ assert jnp.all(jnp.equal(named1["H", ds(2, 5), "D", 3].array, ref))
601
+ ref = jnp.take(named1.array, jnp.arange(3, 8), axis=0, mode="fill", fill_value=0)
602
+ ref = jnp.take(ref, jnp.arange(3, 10, 2), axis=2, mode="fill", fill_value=0)
603
+ assert jnp.all(jnp.equal(named1["H", ds(3, 5), "D", 3:10:2].array, ref))
604
+
605
+
606
+ def test_dslice_oob_read_and_write():
607
+ Seq = hax.Axis("seq", 5)
608
+ from haliax import ds
609
+
610
+ arr = hax.arange((Seq,), dtype=int)
611
+ out = arr[{"seq": ds(3, 4)}]
612
+ ref = jnp.take(arr.array, jnp.arange(3, 7), mode="fill", fill_value=0)
613
+ assert jnp.array_equal(out.array, ref)
614
+
615
+ upd = hax.arange((Seq.resize(4),), dtype=int)
616
+ updated = arr.at[{"seq": ds(3, 4)}].set(upd)
617
+ ref_upd = arr.array.at[jnp.arange(3, 7)].set(upd.array, mode="drop")
618
+ assert jnp.array_equal(updated.array, ref_upd)
594
619
 
595
620
 
596
621
  def test_slice_nd_array_present_dims():
@@ -293,7 +293,7 @@ def test_examples():
293
293
  r = einops_rearrange(z, "{B (H: h1 h) (W: w1 w) C} -> (B: B h1 w1) ... (C: C h w) ", h1=2, w1=2)
294
294
  assert r.axes == (Axis("B", B.size * 2 * 2), D, Axis("C", C.size * sH.size * sW.size))
295
295
  # unet attention reordering:
296
- # postional: (qkv heads c) h w -> qkv heads c (h w)
296
+ # positional: (qkv heads c) h w -> qkv heads c (h w)
297
297
  # named: { (embed: qkv heads c) h w } -> qkv heads c (pos: h w)
298
298
  Embed = Axis("embed", 3 * 4 * C.size)
299
299
  attn = hax.random.randint(PRNGKey(0), (Embed, H, W), 0, 255)
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev380"
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