haliax 1.4.dev379__tar.gz → 1.4.dev381__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.dev379 → haliax-1.4.dev381}/PKG-INFO +1 -1
  2. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/api.md +1 -0
  3. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/indexing.md +5 -0
  4. haliax-1.4.dev381/src/haliax/__about__.py +1 -0
  5. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/__init__.py +2 -1
  6. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/axis.py +11 -6
  7. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/core.py +23 -62
  8. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/ops.py +41 -3
  9. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/core_test.py +29 -4
  10. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_ops.py +13 -0
  11. haliax-1.4.dev379/src/haliax/__about__.py +0 -1
  12. {haliax-1.4.dev379 → haliax-1.4.dev381}/.coveragerc +0 -0
  13. {haliax-1.4.dev379 → haliax-1.4.dev381}/.flake8 +0 -0
  14. {haliax-1.4.dev379 → haliax-1.4.dev381}/.github/workflows/publish_dev.yaml +0 -0
  15. {haliax-1.4.dev379 → haliax-1.4.dev381}/.github/workflows/run_pre_commit.yaml +0 -0
  16. {haliax-1.4.dev379 → haliax-1.4.dev381}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  17. {haliax-1.4.dev379 → haliax-1.4.dev381}/.github/workflows/run_tests.yaml +0 -0
  18. {haliax-1.4.dev379 → haliax-1.4.dev381}/.gitignore +0 -0
  19. {haliax-1.4.dev379 → haliax-1.4.dev381}/.playbooks/add-types.md +0 -0
  20. {haliax-1.4.dev379 → haliax-1.4.dev381}/.pre-commit-config.yaml +0 -0
  21. {haliax-1.4.dev379 → haliax-1.4.dev381}/.readthedocs.yaml +0 -0
  22. {haliax-1.4.dev379 → haliax-1.4.dev381}/AGENTS.md +0 -0
  23. {haliax-1.4.dev379 → haliax-1.4.dev381}/CONTRIBUTING.md +0 -0
  24. {haliax-1.4.dev379 → haliax-1.4.dev381}/LICENSE +0 -0
  25. {haliax-1.4.dev379 → haliax-1.4.dev381}/README.md +0 -0
  26. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/broadcasting.md +0 -0
  27. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/cheatsheet.md +0 -0
  28. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/css/material.css +0 -0
  29. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/css/mkdocstrings.css +0 -0
  30. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/faq.md +0 -0
  31. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/data_parallel_mesh.png +0 -0
  32. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  33. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/device_mesh_1d.png +0 -0
  34. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/device_mesh_1d_zero.png +0 -0
  35. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/device_mesh_2d.png +0 -0
  36. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  37. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  38. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  39. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  40. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/device_mesh_2d_zero.png +0 -0
  41. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/fp8.md +0 -0
  42. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/index.md +0 -0
  43. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/matmul.md +0 -0
  44. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/nn.md +0 -0
  45. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/partitioning.md +0 -0
  46. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/rearrange.ipynb +0 -0
  47. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/rearrange.md +0 -0
  48. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/requirements.txt +0 -0
  49. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/scan.md +0 -0
  50. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/state-dict.md +0 -0
  51. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/tutorial.md +0 -0
  52. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/typing.md +0 -0
  53. {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/vmap.md +0 -0
  54. {haliax-1.4.dev379 → haliax-1.4.dev381}/mkdocs.yml +0 -0
  55. {haliax-1.4.dev379 → haliax-1.4.dev381}/pyproject.toml +0 -0
  56. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/__init__.py +0 -0
  57. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/compile_utils.py +0 -0
  58. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/dot.py +0 -0
  59. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/einsum.py +0 -0
  60. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/fp8.py +0 -0
  61. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/parsing.py +0 -0
  62. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/rearrange.py +0 -0
  63. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/scan.py +0 -0
  64. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/state_dict.py +0 -0
  65. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/util.py +0 -0
  66. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/debug.py +0 -0
  67. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/haxtyping.py +0 -0
  68. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/hof.py +0 -0
  69. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/jax_utils.py +0 -0
  70. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/__init__.py +0 -0
  71. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/activations.py +0 -0
  72. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/attention.py +0 -0
  73. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/conv.py +0 -0
  74. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/dropout.py +0 -0
  75. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/embedding.py +0 -0
  76. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/linear.py +0 -0
  77. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/loss.py +0 -0
  78. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/mlp.py +0 -0
  79. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/normalization.py +0 -0
  80. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/pool.py +0 -0
  81. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/scan.py +0 -0
  82. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/partitioning.py +0 -0
  83. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/quantization.py +0 -0
  84. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/random.py +0 -0
  85. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/specialized_fns.py +0 -0
  86. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/state_dict.py +0 -0
  87. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/tree_util.py +0 -0
  88. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/types.py +0 -0
  89. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/util.py +0 -0
  90. {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/wrap.py +0 -0
  91. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_attention.py +0 -0
  92. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_axis.py +0 -0
  93. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_conv.py +0 -0
  94. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_debug.py +0 -0
  95. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_dot.py +0 -0
  96. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_dtype_typing.py +0 -0
  97. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_einsum.py +0 -0
  98. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_fp8.py +0 -0
  99. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_hof.py +0 -0
  100. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_int8.py +0 -0
  101. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_namedarray_typing.py +0 -0
  102. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_nn.py +0 -0
  103. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_parsing.py +0 -0
  104. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_partitioning.py +0 -0
  105. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_pool.py +0 -0
  106. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_random.py +0 -0
  107. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_rearrange.py +0 -0
  108. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_scan.py +0 -0
  109. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_scatter_gather.py +0 -0
  110. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_specialized_fns.py +0 -0
  111. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_state_dict.py +0 -0
  112. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_tree_util.py +0 -0
  113. {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_utils.py +0 -0
  114. {haliax-1.4.dev379 → haliax-1.4.dev381}/uv.lock +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev379
3
+ Version: 1.4.dev381
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/
@@ -258,6 +258,7 @@ These are all more or less directly from JAX's NumPy API.
258
258
 
259
259
  ::: haliax.clip
260
260
  ::: haliax.isclose
261
+ ::: haliax.pad
261
262
  ::: haliax.top_k
262
263
  ::: haliax.trace
263
264
  ::: haliax.tril
@@ -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:
@@ -0,0 +1 @@
1
+ __version__ = "1.4.dev381"
@@ -66,7 +66,7 @@ from .core import (
66
66
  from .haxtyping import Named
67
67
  from .hof import fold, map, scan, vmap
68
68
  from .jax_utils import tree_checkpoint_name
69
- from .ops import clip, isclose, pad_left, trace, tril, triu, where
69
+ from .ops import clip, isclose, pad_left, pad, trace, tril, triu, where
70
70
  from .partitioning import auto_sharded, axis_mapping, fsdp, named_jit, shard, shard_with_axis_mapping
71
71
  from .specialized_fns import top_k
72
72
  from .types import Scalar
@@ -1082,6 +1082,7 @@ __all__ = [
1082
1082
  "are_shape_checks_enabled",
1083
1083
  "isclose",
1084
1084
  "pad_left",
1085
+ "pad",
1085
1086
  "stack",
1086
1087
  "concatenate",
1087
1088
  "eliminate_axes",
@@ -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
 
@@ -1,10 +1,10 @@
1
1
  import typing
2
- from typing import Optional, Union
2
+ from typing import Mapping, Optional, Union
3
3
 
4
4
  import jax
5
5
  import jax.numpy as jnp
6
6
 
7
- from .axis import Axis, AxisSelector
7
+ from .axis import Axis, AxisSelector, axis_name
8
8
  from .core import NamedArray, NamedOrNumeric, broadcast_arrays, broadcast_arrays_and_return_axes, named
9
9
  from .jax_utils import is_scalarish
10
10
 
@@ -152,10 +152,48 @@ def pad_left(array: NamedArray, axis: Axis, new_axis: Axis, value=0) -> NamedArr
152
152
  return NamedArray(padded, array.axes[:idx] + (new_axis,) + array.axes[idx + 1 :])
153
153
 
154
154
 
155
+ def pad(
156
+ array: NamedArray,
157
+ pad_width: Mapping[AxisSelector, tuple[int, int]],
158
+ *,
159
+ mode: str = "constant",
160
+ constant_values: NamedOrNumeric = 0,
161
+ **kwargs,
162
+ ) -> NamedArray:
163
+ """Version of ``jax.numpy.pad`` that works with ``NamedArray``.
164
+
165
+ ``pad_width`` should be a mapping from axis (or axis name) to a ``(before, after)``
166
+ tuple specifying how much padding to add on each side of that axis. Any axis
167
+ not present in ``pad_width`` will not be padded.
168
+ """
169
+
170
+ padding = []
171
+ new_axes = []
172
+ for ax in array.axes:
173
+ left_right = pad_width.get(ax)
174
+ if left_right is None:
175
+ left_right = pad_width.get(axis_name(ax)) # type: ignore[arg-type]
176
+ if left_right is None:
177
+ left_right = (0, 0)
178
+ left, right = left_right
179
+ padding.append((left, right))
180
+ new_axes.append(ax.resize(ax.size + left + right))
181
+
182
+ result = jnp.pad(
183
+ array.array,
184
+ padding,
185
+ mode=mode,
186
+ constant_values=raw_array_or_scalar(constant_values),
187
+ **kwargs,
188
+ )
189
+
190
+ return NamedArray(result, tuple(new_axes))
191
+
192
+
155
193
  def raw_array_or_scalar(x: NamedOrNumeric):
156
194
  if isinstance(x, NamedArray):
157
195
  return x.array
158
196
  return x
159
197
 
160
198
 
161
- __all__ = ["trace", "where", "tril", "triu", "isclose", "pad_left", "clip"]
199
+ __all__ = ["trace", "where", "tril", "triu", "isclose", "pad_left", "pad", "clip"]
@@ -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():
@@ -237,3 +237,16 @@ def test_reductions_produce_scalar_named_arrays_when_None_axis():
237
237
  # But if we specify axes, we always get a NamedArray, even if it's a scalar
238
238
  assert isinstance(hax.mean(named1, axis=("Height", "Width")), NamedArray)
239
239
  assert hax.mean(named1, axis=("Height", "Width")).axes == ()
240
+
241
+
242
+ def test_pad():
243
+ Height = Axis("Height", 3)
244
+ Width = Axis("Width", 2)
245
+
246
+ arr = hax.arange((Height, Width))
247
+ padded = hax.pad(arr, {Height: (1, 2), Width: (0, 1)}, mode="constant", constant_values=0)
248
+
249
+ expected = jnp.pad(arr.array, [(1, 2), (0, 1)], mode="constant", constant_values=0)
250
+ assert padded.axes[0].size == Height.size + 3
251
+ assert padded.axes[1].size == Width.size + 1
252
+ assert jnp.all(expected == padded.array)
@@ -1 +0,0 @@
1
- __version__ = "1.4.dev379"
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