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