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.
- {haliax-1.4.dev379 → haliax-1.4.dev381}/PKG-INFO +1 -1
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/api.md +1 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/indexing.md +5 -0
- haliax-1.4.dev381/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/__init__.py +2 -1
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/axis.py +11 -6
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/core.py +23 -62
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/ops.py +41 -3
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/core_test.py +29 -4
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_ops.py +13 -0
- haliax-1.4.dev379/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev379 → haliax-1.4.dev381}/.coveragerc +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/.flake8 +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/.gitignore +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/AGENTS.md +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/LICENSE +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/README.md +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/css/material.css +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/faq.md +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/fp8.md +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/index.md +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/matmul.md +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/nn.md +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/partitioning.md +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/rearrange.md +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/requirements.txt +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/scan.md +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/state-dict.md +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/tutorial.md +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/typing.md +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/docs/vmap.md +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/mkdocs.yml +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/pyproject.toml +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/haxtyping.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/random.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/types.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/util.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_attention.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_axis.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_conv.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_debug.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_dot.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_hof.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_int8.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_nn.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_pool.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_random.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_scan.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev379 → haliax-1.4.dev381}/tests/test_utils.py +0 -0
- {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.
|
|
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/
|
|
@@ -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
|
-
|
|
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
|
-
|
|
605
|
-
|
|
606
|
-
|
|
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
|
|
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
|
-
|
|
1176
|
-
|
|
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)
|
|
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
|
-
|
|
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
|
-
|
|
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
|
-
|
|
592
|
-
|
|
593
|
-
assert jnp.all(jnp.equal(named1["H", ds(
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|