haliax 1.4.dev315__tar.gz → 1.4.dev317__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.dev315 → haliax-1.4.dev317}/PKG-INFO +1 -1
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/api.md +1 -0
- haliax-1.4.dev317/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/__init__.py +4 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/axis.py +11 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/core.py +2 -3
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/mlp.py +20 -21
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/scan.py +1 -1
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/core_test.py +26 -394
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_attention.py +2 -6
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_axis.py +2 -4
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_einsum.py +1 -6
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_nn.py +5 -15
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_rearrange.py +1 -5
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_specialized_fns.py +3 -4
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_tree_util.py +1 -4
- haliax-1.4.dev315/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev315 → haliax-1.4.dev317}/.coveragerc +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/.flake8 +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/.gitignore +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/LICENSE +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/README.md +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/css/material.css +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/faq.md +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/fp8.md +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/hof.md +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/index.md +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/indexing.md +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/matmul.md +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/nn.md +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/partitioning.md +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/rearrange.md +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/requirements.txt +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/docs/tutorial.md +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/mkdocs.yml +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/pyproject.toml +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/embedding.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/random.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/types.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/util.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_conv.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_debug.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_dot.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_hof.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_ops.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_pool.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_random.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_scan.py +0 -0
- {haliax-1.4.dev315 → haliax-1.4.dev317}/tests/test_utils.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.3
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev317
|
|
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/
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "1.4.dev317"
|
|
@@ -32,6 +32,7 @@ from .axis import (
|
|
|
32
32
|
ds,
|
|
33
33
|
dslice,
|
|
34
34
|
eliminate_axes,
|
|
35
|
+
make_axes,
|
|
35
36
|
selects_axis,
|
|
36
37
|
)
|
|
37
38
|
from .core import (
|
|
@@ -893,6 +894,9 @@ __all__ = [
|
|
|
893
894
|
"AxisSpec",
|
|
894
895
|
"AxisSelection",
|
|
895
896
|
"AxisSelector",
|
|
897
|
+
"make_axes",
|
|
898
|
+
"axis_name",
|
|
899
|
+
"axis_size",
|
|
896
900
|
"NamedArray",
|
|
897
901
|
"broadcast_to",
|
|
898
902
|
"broadcast_axis",
|
|
@@ -28,6 +28,17 @@ class Axis:
|
|
|
28
28
|
return f"{self.name}({self.size})"
|
|
29
29
|
|
|
30
30
|
|
|
31
|
+
def make_axes(**kwargs: int) -> Tuple[Axis, ...]:
|
|
32
|
+
"""
|
|
33
|
+
Convenience function for creating a tuple of Axis objects.
|
|
34
|
+
|
|
35
|
+
Example:
|
|
36
|
+
```
|
|
37
|
+
X, Y = axes(X=10, Y=20)
|
|
38
|
+
"""
|
|
39
|
+
return tuple(Axis(name, size) for name, size in kwargs.items())
|
|
40
|
+
|
|
41
|
+
|
|
31
42
|
AxisSelector = Union[Axis, str]
|
|
32
43
|
"""AxisSelector is a type that can be used to select a single axis from an array. str or Axis"""
|
|
33
44
|
AxisSelection = Union[AxisSelector, Sequence[AxisSelector]]
|
|
@@ -373,11 +373,10 @@ class NamedArray:
|
|
|
373
373
|
|
|
374
374
|
Supports indexing like:
|
|
375
375
|
|
|
376
|
-
>>> X =
|
|
377
|
-
>>> Y = Axis("y", 20)
|
|
376
|
+
>>> X, Y = haliax.make_axes(X=10, Y=20)
|
|
378
377
|
>>> arr = haliax.random.randint(jax.random.PRNGKey(0), (X, Y), 0, X.size)
|
|
379
378
|
# slice with ints or slices
|
|
380
|
-
>>> arr[{"x": 1, "y": slice(0,10,
|
|
379
|
+
>>> arr[{"x": 1, "y": slice(0,10,2)}]
|
|
381
380
|
>>> Z = Axis("z", 3)
|
|
382
381
|
# so-called "advanced indexing" with NamedArrays.
|
|
383
382
|
>>> index_arr = NamedArray(np.array([1, 2, 3]), Z)
|
|
@@ -44,6 +44,7 @@ class MLP(eqx.Module):
|
|
|
44
44
|
depth: int,
|
|
45
45
|
activation: Callable = relu,
|
|
46
46
|
*,
|
|
47
|
+
out_first: bool = True,
|
|
47
48
|
use_bias: bool = True,
|
|
48
49
|
use_final_bias: bool = True,
|
|
49
50
|
key: PRNGKeyArray,
|
|
@@ -57,36 +58,34 @@ class MLP(eqx.Module):
|
|
|
57
58
|
|
|
58
59
|
layers = []
|
|
59
60
|
|
|
61
|
+
kwargs: dict = {
|
|
62
|
+
"use_bias": use_bias,
|
|
63
|
+
"dot_general": dot_general,
|
|
64
|
+
"init_scale": init_scale,
|
|
65
|
+
"out_first": out_first,
|
|
66
|
+
}
|
|
67
|
+
|
|
68
|
+
last_kwargs: dict = {
|
|
69
|
+
"use_bias": use_final_bias,
|
|
70
|
+
"dot_general": dot_general,
|
|
71
|
+
"init_scale": init_scale,
|
|
72
|
+
"out_first": out_first,
|
|
73
|
+
}
|
|
74
|
+
|
|
60
75
|
if depth == 0:
|
|
61
76
|
# special case: no hidden layers
|
|
62
|
-
layers.append(
|
|
63
|
-
Linear.init(
|
|
64
|
-
Input, Output, use_bias=use_final_bias, key=keys[0], dot_general=dot_general, init_scale=init_scale
|
|
65
|
-
)
|
|
66
|
-
)
|
|
77
|
+
layers.append(Linear.init(Input, Output, key=keys[0], **last_kwargs))
|
|
67
78
|
else:
|
|
68
79
|
# first hidden layer
|
|
69
|
-
layers.append(
|
|
70
|
-
Linear.init(
|
|
71
|
-
Input, Width, use_bias=use_bias, key=keys[0], dot_general=dot_general, init_scale=init_scale
|
|
72
|
-
)
|
|
73
|
-
)
|
|
80
|
+
layers.append(Linear.init(Input, Width, key=keys[0], **kwargs))
|
|
74
81
|
# middle hidden layers
|
|
75
82
|
cur = Width
|
|
76
83
|
next = Width2
|
|
77
84
|
for i in range(1, depth):
|
|
78
|
-
layers.append(
|
|
79
|
-
Linear.init(
|
|
80
|
-
cur, next, use_bias=use_bias, key=keys[i], dot_general=dot_general, init_scale=init_scale
|
|
81
|
-
)
|
|
82
|
-
)
|
|
85
|
+
layers.append(Linear.init(cur, next, key=keys[i], **kwargs))
|
|
83
86
|
cur, next = next, cur
|
|
84
|
-
# final
|
|
85
|
-
layers.append(
|
|
86
|
-
Linear.init(
|
|
87
|
-
cur, Output, use_bias=use_final_bias, key=keys[-1], dot_general=dot_general, init_scale=init_scale
|
|
88
|
-
)
|
|
89
|
-
)
|
|
87
|
+
# final layer
|
|
88
|
+
layers.append(Linear.init(cur, Output, key=keys[-1], **last_kwargs))
|
|
90
89
|
|
|
91
90
|
return MLP(
|
|
92
91
|
layers=tuple(layers),
|
|
@@ -300,7 +300,7 @@ class Stacked(eqx.Module, Generic[M]):
|
|
|
300
300
|
"""
|
|
301
301
|
|
|
302
302
|
def unbatch_leaf(x):
|
|
303
|
-
if haliax.
|
|
303
|
+
if isinstance(x, haliax.core.NamedArray):
|
|
304
304
|
if haliax.selects_axis(x.axes, self.Block):
|
|
305
305
|
return haliax.unbind(x, self.Block)
|
|
306
306
|
else:
|
|
@@ -9,9 +9,7 @@ from haliax import Axis, NamedArray
|
|
|
9
9
|
|
|
10
10
|
|
|
11
11
|
def test_unary_np_functions():
|
|
12
|
-
Height =
|
|
13
|
-
Width = Axis("Width", 3)
|
|
14
|
-
Depth = Axis("Depth", 4)
|
|
12
|
+
Height, Width, Depth = hax.make_axes(Height=2, Width=3, Depth=4)
|
|
15
13
|
|
|
16
14
|
m1 = NamedArray(jnp.ones((Height.size, Width.size, Depth.size)), (Height, Width, Depth))
|
|
17
15
|
|
|
@@ -29,9 +27,7 @@ def test_unary_np_functions():
|
|
|
29
27
|
|
|
30
28
|
|
|
31
29
|
def test_reduction_functions():
|
|
32
|
-
Height =
|
|
33
|
-
Width = Axis("Width", 3)
|
|
34
|
-
Depth = Axis("Depth", 4)
|
|
30
|
+
Height, Width, Depth = hax.make_axes(Height=2, Width=3, Depth=4)
|
|
35
31
|
|
|
36
32
|
rand_m = jax.random.uniform(PRNGKey(0), (Height.size, Width.size, Depth.size))
|
|
37
33
|
|
|
@@ -69,9 +65,7 @@ def test_reduction_functions():
|
|
|
69
65
|
|
|
70
66
|
|
|
71
67
|
def test_reduction_functions_with_where():
|
|
72
|
-
H =
|
|
73
|
-
W = Axis("W", 3)
|
|
74
|
-
D = Axis("D", 4)
|
|
68
|
+
H, W, D = hax.make_axes(H=2, W=3, D=4)
|
|
75
69
|
|
|
76
70
|
rand_m = jax.random.uniform(PRNGKey(0), (H.size, W.size, D.size))
|
|
77
71
|
|
|
@@ -118,9 +112,7 @@ def test_reduction_functions_with_where():
|
|
|
118
112
|
|
|
119
113
|
|
|
120
114
|
def test_split():
|
|
121
|
-
Height =
|
|
122
|
-
Width = Axis("Width", 3)
|
|
123
|
-
Depth = Axis("Depth", 4)
|
|
115
|
+
Height, Width, Depth = hax.make_axes(Height=2, Width=3, Depth=4)
|
|
124
116
|
|
|
125
117
|
D10 = Axis("Depth", Depth.size * 10)
|
|
126
118
|
|
|
@@ -144,9 +136,7 @@ def test_split():
|
|
|
144
136
|
|
|
145
137
|
|
|
146
138
|
def test_take():
|
|
147
|
-
Height =
|
|
148
|
-
Width = Axis("Width", 3)
|
|
149
|
-
Depth = Axis("Depth", 4)
|
|
139
|
+
Height, Width, Depth = hax.make_axes(Height=2, Width=3, Depth=4)
|
|
150
140
|
named1 = hax.random.uniform(PRNGKey(0), (Height, Width, Depth))
|
|
151
141
|
|
|
152
142
|
assert jnp.all(jnp.equal(hax.take(named1, Height, 0).array, named1.array[0]))
|
|
@@ -178,9 +168,7 @@ def test_take():
|
|
|
178
168
|
|
|
179
169
|
|
|
180
170
|
def test_take_overlapping_names():
|
|
181
|
-
Height =
|
|
182
|
-
Width = Axis("Width", 30)
|
|
183
|
-
Depth = Axis("Depth", 40)
|
|
171
|
+
Height, Width, Depth = hax.make_axes(Height=20, Width=30, Depth=40)
|
|
184
172
|
named1 = hax.random.uniform(PRNGKey(0), (Height, Width, Depth))
|
|
185
173
|
|
|
186
174
|
Height2 = Axis("Height", 10)
|
|
@@ -198,9 +186,7 @@ def test_take_overlapping_2():
|
|
|
198
186
|
def cross_entropy(logits: hax.NamedArray, labels: hax.NamedArray) -> hax.NamedArray:
|
|
199
187
|
return hax.take(logits, Embed, labels) # extract log probability of the correct token
|
|
200
188
|
|
|
201
|
-
Embed =
|
|
202
|
-
Block = Axis("Block", 20)
|
|
203
|
-
Batch = Axis("Batch", 30)
|
|
189
|
+
Embed, Block, Batch = hax.make_axes(Embed=10, Block=20, Batch=30)
|
|
204
190
|
logits = hax.random.uniform(PRNGKey(0), (Batch, Block, Embed))
|
|
205
191
|
labels = hax.random.randint(PRNGKey(0), (Batch, Block), 0, Embed.size)
|
|
206
192
|
|
|
@@ -223,9 +209,7 @@ def test_take_overlapping_2():
|
|
|
223
209
|
|
|
224
210
|
|
|
225
211
|
def test_cumsum_etc():
|
|
226
|
-
Height =
|
|
227
|
-
Width = Axis("Width", 3)
|
|
228
|
-
Depth = Axis("Depth", 4)
|
|
212
|
+
Height, Width, Depth = hax.make_axes(Height=2, Width=3, Depth=4)
|
|
229
213
|
|
|
230
214
|
named1 = hax.random.uniform(PRNGKey(0), (Height, Width, Depth))
|
|
231
215
|
|
|
@@ -255,10 +239,7 @@ def test_cumsum_etc():
|
|
|
255
239
|
|
|
256
240
|
|
|
257
241
|
def test_rearrange():
|
|
258
|
-
H =
|
|
259
|
-
W = Axis("W", 3)
|
|
260
|
-
D = Axis("D", 4)
|
|
261
|
-
C = Axis("C", 5)
|
|
242
|
+
H, W, D, C = hax.make_axes(H=2, W=3, D=4, C=5)
|
|
262
243
|
|
|
263
244
|
named1 = hax.random.uniform(PRNGKey(0), (H, W, D, C))
|
|
264
245
|
|
|
@@ -304,9 +285,7 @@ def test_rearrange():
|
|
|
304
285
|
|
|
305
286
|
def test_rearrange_unused_ellipsis():
|
|
306
287
|
# Make sure we just ignore the ellipsis if all axes are specified in addition
|
|
307
|
-
H =
|
|
308
|
-
W = Axis("Width", 3)
|
|
309
|
-
D = Axis("Depth", 4)
|
|
288
|
+
H, W, D = hax.make_axes(Height=2, Width=3, Depth=4)
|
|
310
289
|
|
|
311
290
|
named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
|
|
312
291
|
|
|
@@ -321,7 +300,7 @@ def test_rearrange_unused_ellipsis():
|
|
|
321
300
|
|
|
322
301
|
|
|
323
302
|
def test_arange():
|
|
324
|
-
H =
|
|
303
|
+
(H,) = hax.make_axes(H=10)
|
|
325
304
|
|
|
326
305
|
assert jnp.all(jnp.equal(hax.arange(H).array, jnp.arange(10)))
|
|
327
306
|
assert hax.arange(H).axes == (H,)
|
|
@@ -334,8 +313,7 @@ def test_arange():
|
|
|
334
313
|
|
|
335
314
|
|
|
336
315
|
def test_stack():
|
|
337
|
-
H =
|
|
338
|
-
W = Axis("W", 3)
|
|
316
|
+
H, W = hax.make_axes(H=4, W=3)
|
|
339
317
|
|
|
340
318
|
named1 = hax.random.uniform(PRNGKey(0), (H, W))
|
|
341
319
|
named2 = hax.random.uniform(PRNGKey(1), (H, W))
|
|
@@ -350,9 +328,8 @@ def test_stack():
|
|
|
350
328
|
|
|
351
329
|
|
|
352
330
|
def test_concatenate():
|
|
353
|
-
H1 =
|
|
354
|
-
H2 =
|
|
355
|
-
W = Axis("W", 3)
|
|
331
|
+
H1, W = hax.make_axes(H=4, W=3)
|
|
332
|
+
H2 = H1.resize(3)
|
|
356
333
|
|
|
357
334
|
named1 = hax.random.uniform(PRNGKey(0), (H1, W))
|
|
358
335
|
named2 = hax.random.uniform(PRNGKey(1), (H2, W))
|
|
@@ -392,8 +369,7 @@ def test_repeat():
|
|
|
392
369
|
# [3, 4],
|
|
393
370
|
# [3, 4]])
|
|
394
371
|
|
|
395
|
-
H =
|
|
396
|
-
W = Axis("W", 2)
|
|
372
|
+
H, W = hax.make_axes(H=2, W=2)
|
|
397
373
|
|
|
398
374
|
named1 = hax.named([[1, 2], [3, 4]], (H, W))
|
|
399
375
|
|
|
@@ -459,9 +435,7 @@ def test_tile():
|
|
|
459
435
|
|
|
460
436
|
|
|
461
437
|
def test_unflatten_axis():
|
|
462
|
-
H =
|
|
463
|
-
W = Axis("Width", 3)
|
|
464
|
-
D = Axis("Depth", 4)
|
|
438
|
+
H, W, D = hax.make_axes(Height=2, Width=3, Depth=4)
|
|
465
439
|
|
|
466
440
|
named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
|
|
467
441
|
flattened_HW = named1.flatten_axes((H, W), "Z")
|
|
@@ -483,9 +457,7 @@ def test_unflatten_axis():
|
|
|
483
457
|
|
|
484
458
|
|
|
485
459
|
def test_ravel():
|
|
486
|
-
H =
|
|
487
|
-
W = Axis("Width", 3)
|
|
488
|
-
D = Axis("Depth", 4)
|
|
460
|
+
H, W, D = hax.make_axes(Height=2, Width=3, Depth=4)
|
|
489
461
|
|
|
490
462
|
named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
|
|
491
463
|
raveled = named1.ravel("Z")
|
|
@@ -496,13 +468,8 @@ def test_ravel():
|
|
|
496
468
|
|
|
497
469
|
|
|
498
470
|
def test_rename():
|
|
499
|
-
H =
|
|
500
|
-
|
|
501
|
-
D = Axis("D", 4)
|
|
502
|
-
|
|
503
|
-
H2 = Axis("H2", 2)
|
|
504
|
-
W2 = Axis("W2", 3)
|
|
505
|
-
D2 = Axis("D2", 4)
|
|
471
|
+
H, W, D = hax.make_axes(H=2, W=3, D=4)
|
|
472
|
+
H2, W2, D2 = hax.make_axes(H2=2, W2=3, D2=4)
|
|
506
473
|
|
|
507
474
|
named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
|
|
508
475
|
|
|
@@ -517,9 +484,7 @@ def test_rename():
|
|
|
517
484
|
|
|
518
485
|
|
|
519
486
|
def test_index():
|
|
520
|
-
H =
|
|
521
|
-
W = Axis("W", 30)
|
|
522
|
-
D = Axis("D", 40)
|
|
487
|
+
H, W, D = hax.make_axes(H=20, W=30, D=40)
|
|
523
488
|
|
|
524
489
|
named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
|
|
525
490
|
|
|
@@ -544,9 +509,7 @@ def test_index():
|
|
|
544
509
|
|
|
545
510
|
|
|
546
511
|
def test_index_with_tracer():
|
|
547
|
-
H =
|
|
548
|
-
W = Axis("W", 30)
|
|
549
|
-
D = Axis("D", 40)
|
|
512
|
+
H, W, D = hax.make_axes(H=20, W=30, D=40)
|
|
550
513
|
named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
|
|
551
514
|
|
|
552
515
|
@jax.jit
|
|
@@ -562,12 +525,7 @@ def test_index_with_tracer():
|
|
|
562
525
|
|
|
563
526
|
def test_index_array_slices():
|
|
564
527
|
# fancier tests with array slices with named array args
|
|
565
|
-
H =
|
|
566
|
-
W = Axis("W", 20)
|
|
567
|
-
D = Axis("D", 30)
|
|
568
|
-
C = Axis("C", 40)
|
|
569
|
-
Q = Axis("Q", 50)
|
|
570
|
-
I0 = Axis("I0", 10)
|
|
528
|
+
H, W, D, C, Q, I0 = hax.make_axes(H=10, W=20, D=30, C=40, Q=50, I0=10)
|
|
571
529
|
|
|
572
530
|
named1 = hax.random.uniform(PRNGKey(0), (H, W, D, C, Q))
|
|
573
531
|
index_1 = hax.random.randint(PRNGKey(0), (I0,), 0, H.size)
|
|
@@ -586,9 +544,7 @@ def test_index_array_slices():
|
|
|
586
544
|
# https://numpy.org/doc/stable/user/basics.indexing.html#combining-advanced-and-basic-indexing
|
|
587
545
|
# Example
|
|
588
546
|
# Let x.shape be (10, 20, 30, 40, 50) and suppose ind_1 and ind_2 can be broadcast to the shape (2, 3, 4).
|
|
589
|
-
I1 =
|
|
590
|
-
I2 = Axis("I2", 3)
|
|
591
|
-
I3 = Axis("I3", 4)
|
|
547
|
+
I1, I2, I3 = hax.make_axes(I1=2, I2=3, I3=4)
|
|
592
548
|
|
|
593
549
|
ind_1 = hax.random.randint(PRNGKey(0), (I2, I3), 0, W.size)
|
|
594
550
|
ind_2 = hax.random.randint(PRNGKey(0), (I1, I3), 0, D.size)
|
|
@@ -615,9 +571,7 @@ def test_index_array_slices():
|
|
|
615
571
|
def test_slice_nd_shorthand_syntax():
|
|
616
572
|
# syntax like arr["X", 0:10, "Y", 0:10] is supported
|
|
617
573
|
|
|
618
|
-
H =
|
|
619
|
-
W = Axis("W", 20)
|
|
620
|
-
D = Axis("D", 30)
|
|
574
|
+
H, W, D = hax.make_axes(H=10, W=20, D=30)
|
|
621
575
|
|
|
622
576
|
named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
|
|
623
577
|
|
|
@@ -625,9 +579,7 @@ def test_slice_nd_shorthand_syntax():
|
|
|
625
579
|
|
|
626
580
|
|
|
627
581
|
def test_slice_nd_dslice():
|
|
628
|
-
H =
|
|
629
|
-
W = Axis("W", 20)
|
|
630
|
-
D = Axis("D", 30)
|
|
582
|
+
H, W, D = hax.make_axes(H=10, W=20, D=30)
|
|
631
583
|
|
|
632
584
|
named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
|
|
633
585
|
from haliax import ds
|
|
@@ -641,9 +593,7 @@ def test_slice_nd_dslice():
|
|
|
641
593
|
|
|
642
594
|
def test_slice_nd_array_present_dims():
|
|
643
595
|
# tests slicing with arrays that are already present in the named array, which is sometimes ok
|
|
644
|
-
H =
|
|
645
|
-
W = Axis("W", 20)
|
|
646
|
-
D = Axis("D", 30)
|
|
596
|
+
H, W, D = hax.make_axes(H=10, W=20, D=30)
|
|
647
597
|
|
|
648
598
|
named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
|
|
649
599
|
|
|
@@ -653,321 +603,3 @@ def test_slice_nd_array_present_dims():
|
|
|
653
603
|
assert jnp.all(jnp.equal(named1[{"H": index1}].array, named1.array[index1.array, :, :]))
|
|
654
604
|
|
|
655
605
|
# this is not ok, since the H would not be eliminated
|
|
656
|
-
with pytest.raises(ValueError):
|
|
657
|
-
named1[{W: index1}]
|
|
658
|
-
|
|
659
|
-
# this is not ok, but is trickier because the H has a different size
|
|
660
|
-
H2 = H.resize(5)
|
|
661
|
-
index2 = hax.random.randint(PRNGKey(0), (H2,), 0, H.size)
|
|
662
|
-
with pytest.raises(ValueError):
|
|
663
|
-
named1[{W: index2}]
|
|
664
|
-
|
|
665
|
-
# this is ok, since the H would be eliminated anyway
|
|
666
|
-
assert jnp.all(jnp.equal(named1[{"H": index2}].array, named1.array[index2.array, :, :]))
|
|
667
|
-
|
|
668
|
-
|
|
669
|
-
def test_slice_nd_array_unnamed_slice():
|
|
670
|
-
# tests slicing with arrays that are already present in the named array, which is sometimes ok
|
|
671
|
-
H = Axis("H", 10)
|
|
672
|
-
W = Axis("W", 20)
|
|
673
|
-
D = Axis("D", 30)
|
|
674
|
-
|
|
675
|
-
named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
|
|
676
|
-
|
|
677
|
-
index1 = jax.random.randint(PRNGKey(1), (4,), 0, H.size)
|
|
678
|
-
assert jnp.all(jnp.equal(named1[{"H": index1}].array, named1.array[index1, :, :]))
|
|
679
|
-
|
|
680
|
-
# hidden behavior: if we also pass in an H index to e.g. D, it is zipped together
|
|
681
|
-
index2 = hax.random.randint(PRNGKey(2), Axis("H", 4), 0, D.size)
|
|
682
|
-
assert jnp.all(jnp.equal(named1[{"H": index1, "D": index2}].array, named1.array[index1, :, index2.array]))
|
|
683
|
-
|
|
684
|
-
# this is different though:
|
|
685
|
-
index2r = index2.array
|
|
686
|
-
assert jnp.all(
|
|
687
|
-
jnp.equal(
|
|
688
|
-
named1[{"H": index1, "D": index2r}].array, named1.array[index1.reshape(1, -1), :, index2r.reshape(-1, 1)]
|
|
689
|
-
)
|
|
690
|
-
)
|
|
691
|
-
assert named1[{"H": index1, "D": index2r}].shape != named1[{"H": index1, "D": index2}].shape
|
|
692
|
-
|
|
693
|
-
index1 = list(index1)
|
|
694
|
-
assert jnp.all(jnp.equal(named1[{"H": index1}].array, named1.array[index1, :, :]))
|
|
695
|
-
|
|
696
|
-
|
|
697
|
-
def test_full_indexing_returns_named_array():
|
|
698
|
-
H = Axis("H", 10)
|
|
699
|
-
W = Axis("W", 20)
|
|
700
|
-
D = Axis("D", 30)
|
|
701
|
-
|
|
702
|
-
named1 = hax.random.uniform(PRNGKey(0), (H, W, D))
|
|
703
|
-
sliced = named1[{"H": 0, "W": 0, "D": 0}]
|
|
704
|
-
|
|
705
|
-
assert isinstance(sliced, NamedArray)
|
|
706
|
-
assert sliced.shape == {}
|
|
707
|
-
|
|
708
|
-
|
|
709
|
-
def test_indexing_bug_from_docs():
|
|
710
|
-
X = hax.Axis("X", 10)
|
|
711
|
-
Y = hax.Axis("Y", 20)
|
|
712
|
-
Z = hax.Axis("Z", 30)
|
|
713
|
-
|
|
714
|
-
a = hax.random.uniform(jax.random.PRNGKey(0), (X, Y, Z))
|
|
715
|
-
|
|
716
|
-
I1 = hax.Axis("I1", 5)
|
|
717
|
-
I2 = hax.Axis("I2", 5)
|
|
718
|
-
I3 = hax.Axis("I3", 5)
|
|
719
|
-
ind1 = hax.random.randint(jax.random.PRNGKey(0), (I1,), 0, 10)
|
|
720
|
-
ind2 = hax.random.randint(jax.random.PRNGKey(0), (I2, I3), 0, 20)
|
|
721
|
-
|
|
722
|
-
# assert a[{"X": ind1, "Y": ind2}].axes == (I1, I2, I3, Z)
|
|
723
|
-
assert a[{"X": ind1, "Y": ind2, "Z": 3}].axes == (I1, I2, I3)
|
|
724
|
-
|
|
725
|
-
|
|
726
|
-
def test_duplicate_axis_names_in_slicing():
|
|
727
|
-
X = hax.Axis("X", 10)
|
|
728
|
-
Y = hax.Axis("Y", 20)
|
|
729
|
-
Z = hax.Axis("Z", 30)
|
|
730
|
-
|
|
731
|
-
X2 = hax.Axis("X", 5)
|
|
732
|
-
Y2 = hax.Axis("Y", 5)
|
|
733
|
-
|
|
734
|
-
a = hax.random.uniform(jax.random.PRNGKey(0), (X, Y, Z))
|
|
735
|
-
ind1 = hax.random.randint(jax.random.PRNGKey(0), (X2,), 0, 10)
|
|
736
|
-
ind2 = hax.random.randint(jax.random.PRNGKey(0), (Y2,), 0, 10)
|
|
737
|
-
|
|
738
|
-
a[{"X": ind1, "Y": ind2}] # returns a NamedArray with axes = Axis("X", 5), Axis("Y", 5), Axis("Z", 30)
|
|
739
|
-
|
|
740
|
-
with pytest.raises(ValueError):
|
|
741
|
-
a[{"Y": ind1}] # error, "X" is not eliminated by the indexing operation
|
|
742
|
-
|
|
743
|
-
a[{"X": ind2, "Y": ind1}] # ok, because X and Y are eliminated by the indexing operation
|
|
744
|
-
|
|
745
|
-
|
|
746
|
-
def test_slice_old_style():
|
|
747
|
-
H = Axis("H", 10)
|
|
748
|
-
W = Axis("W", 20)
|
|
749
|
-
D = Axis("D", 30)
|
|
750
|
-
|
|
751
|
-
named1 = hax.random.randint(PRNGKey(0), (H, W, D), minval=0, maxval=10)
|
|
752
|
-
|
|
753
|
-
assert jnp.all(named1.slice("H", start=4, length=2).array == named1.array[4:6, :, :])
|
|
754
|
-
assert jnp.all(named1.slice("W", start=4, length=2).array == named1.array[:, 4:6, :])
|
|
755
|
-
assert jnp.all(named1.slice("D", start=4, length=2).array == named1.array[:, :, 4:6])
|
|
756
|
-
|
|
757
|
-
H2 = Axis("H2", 5)
|
|
758
|
-
W2 = Axis("W2", 10)
|
|
759
|
-
D2 = Axis("D2", 15)
|
|
760
|
-
|
|
761
|
-
assert jnp.all(named1.slice("H", H2, start=4).array == named1.array[4 : 4 + H2.size, :, :])
|
|
762
|
-
assert jnp.all(named1.slice("W", W2, start=4).array == named1.array[:, 4 : 4 + W2.size, :])
|
|
763
|
-
assert jnp.all(named1.slice("D", D2, start=4).array == named1.array[:, :, 4 : 4 + D2.size])
|
|
764
|
-
|
|
765
|
-
|
|
766
|
-
def test_slice_new_style():
|
|
767
|
-
H = Axis("H", 10)
|
|
768
|
-
W = Axis("W", 20)
|
|
769
|
-
D = Axis("D", 30)
|
|
770
|
-
|
|
771
|
-
named1 = hax.random.randint(PRNGKey(0), (H, W, D), minval=0, maxval=10)
|
|
772
|
-
|
|
773
|
-
x1 = named1.slice({"H": 4, "W": 5, "D": 7}, length={"H": 2, "W": 3, "D": 4})
|
|
774
|
-
assert jnp.all(x1.array == named1.array[4:6, 5:8, 7:11])
|
|
775
|
-
|
|
776
|
-
with pytest.raises(TypeError):
|
|
777
|
-
named1.slice({"H": 4, "W": 5, "D": 7}, length={"H": 2, "W": 3, "D": 4}, start={"H": 1, "W": 2, "D": 3})
|
|
778
|
-
|
|
779
|
-
with pytest.raises(ValueError):
|
|
780
|
-
named1.slice({"H": 4, "W": 5, "D": 7}, length={"H": 2, "W": 3})
|
|
781
|
-
|
|
782
|
-
H2 = Axis("H2", 5)
|
|
783
|
-
W2 = Axis("W2", 10)
|
|
784
|
-
D2 = Axis("D2", 15)
|
|
785
|
-
|
|
786
|
-
x2 = named1.slice({"H": 4, "W": 5, "D": 7}, length={"H": H2, "W": W2, "D": D2})
|
|
787
|
-
assert jnp.all(x2.array == named1.array[4 : 4 + H2.size, 5 : 5 + W2.size, 7 : 7 + D2.size])
|
|
788
|
-
|
|
789
|
-
|
|
790
|
-
def test_updated_slice():
|
|
791
|
-
H = Axis("H", 10)
|
|
792
|
-
W = Axis("W", 20)
|
|
793
|
-
D = Axis("D", 30)
|
|
794
|
-
|
|
795
|
-
H2 = H.resize(5)
|
|
796
|
-
W2 = W.resize(10)
|
|
797
|
-
D2 = D.resize(15)
|
|
798
|
-
|
|
799
|
-
named1 = hax.random.randint(PRNGKey(0), (H, W, D), minval=0, maxval=10)
|
|
800
|
-
named2 = hax.random.randint(PRNGKey(0), (H2, W2, D2), minval=10, maxval=30)
|
|
801
|
-
|
|
802
|
-
named1_updated = named1.updated_slice({"H": 0, "W": 0, "D": 0}, named2)
|
|
803
|
-
|
|
804
|
-
assert named1_updated.axes == named1.axes
|
|
805
|
-
assert jnp.all(named1_updated["H", 0 : H2.size, "W", 0 : W2.size, "D", 0 : D2.size].array == named2.array)
|
|
806
|
-
|
|
807
|
-
# test broadcasting
|
|
808
|
-
for pair in [(H2, D2), (H2, W2), (W2, D2), (D2, H2), (D2, W2), (W2, H2)]:
|
|
809
|
-
n3 = hax.random.randint(PRNGKey(0), pair, minval=10, maxval=30)
|
|
810
|
-
named1_updated = named1.updated_slice({ax.name: 0 for ax in pair}, n3)
|
|
811
|
-
assert named1_updated.axes == named1.axes
|
|
812
|
-
assert jnp.all((named1_updated[{ax.name: slice(0, ax.size) for ax in pair}] == n3).array)
|
|
813
|
-
# check that the array outside the slice is unchanged
|
|
814
|
-
assert jnp.all(
|
|
815
|
-
(
|
|
816
|
-
named1_updated[{ax.name: slice(ax.size, None) for ax in pair}]
|
|
817
|
-
== named1[{ax.name: slice(ax.size, None) for ax in pair}]
|
|
818
|
-
).array
|
|
819
|
-
)
|
|
820
|
-
|
|
821
|
-
|
|
822
|
-
def test_updated_slice_extra_update_axis_errors():
|
|
823
|
-
H = Axis("H", 10)
|
|
824
|
-
W = Axis("W", 20)
|
|
825
|
-
D = Axis("D", 30)
|
|
826
|
-
|
|
827
|
-
named1 = hax.random.randint(PRNGKey(0), (H, W, D), minval=0, maxval=10)
|
|
828
|
-
named2 = hax.random.randint(PRNGKey(0), (H, W, D), minval=10, maxval=30)
|
|
829
|
-
|
|
830
|
-
with pytest.raises(ValueError):
|
|
831
|
-
named1.updated_slice({"H": 0, "W": 0, "D": 0, "extra": 0}, named2)
|
|
832
|
-
|
|
833
|
-
with pytest.raises(ValueError):
|
|
834
|
-
named3 = hax.random.randint(PRNGKey(0), (H, W), minval=10, maxval=30)
|
|
835
|
-
named3.updated_slice({"H": 0, "W": 0}, named2)
|
|
836
|
-
|
|
837
|
-
|
|
838
|
-
def test_order_of_transpose_add():
|
|
839
|
-
H = Axis("H", 10)
|
|
840
|
-
W = Axis("W", 20)
|
|
841
|
-
|
|
842
|
-
named1 = hax.random.randint(PRNGKey(0), (H, W), minval=0, maxval=10)
|
|
843
|
-
named2 = hax.random.randint(PRNGKey(0), (W, H), minval=10, maxval=30)
|
|
844
|
-
|
|
845
|
-
assert (named1 + named2).axes == (H, W)
|
|
846
|
-
assert jnp.all((named1 + named2).array == named1.array + named2.array.T)
|
|
847
|
-
|
|
848
|
-
|
|
849
|
-
def test_nice_short_string_in_named_array():
|
|
850
|
-
H = Axis("H", 10)
|
|
851
|
-
W = Axis("W", 20)
|
|
852
|
-
|
|
853
|
-
named1 = hax.random.randint(PRNGKey(0), (H, W), minval=0, maxval=10)
|
|
854
|
-
|
|
855
|
-
assert str(named1).startswith("NamedArray(int32{'H': 10, 'W': 20}")
|
|
856
|
-
|
|
857
|
-
|
|
858
|
-
def test_nice_short_string_in_named_array_in_eqx_module():
|
|
859
|
-
H = Axis("H", 10)
|
|
860
|
-
W = Axis("W", 20)
|
|
861
|
-
|
|
862
|
-
named1 = hax.random.randint(PRNGKey(0), (H, W), minval=0, maxval=10)
|
|
863
|
-
|
|
864
|
-
class TestModule(eqx.Module):
|
|
865
|
-
named1: NamedArray
|
|
866
|
-
|
|
867
|
-
mod = TestModule(named1)
|
|
868
|
-
|
|
869
|
-
assert str(mod).startswith("TestModule(named1=Named(int32{'H': 10, 'W': 20}))")
|
|
870
|
-
|
|
871
|
-
|
|
872
|
-
def test_named_arrays_work_in_eqxi_while_loop():
|
|
873
|
-
H = Axis("H", 10)
|
|
874
|
-
W = Axis("W", 20)
|
|
875
|
-
|
|
876
|
-
named1 = hax.random.uniform(PRNGKey(0), (H, W))
|
|
877
|
-
|
|
878
|
-
import equinox.internal as eqxi
|
|
879
|
-
|
|
880
|
-
def body_fun(t):
|
|
881
|
-
i, named1 = t
|
|
882
|
-
return i + 1, named1 + named1
|
|
883
|
-
|
|
884
|
-
def cond_fun(t):
|
|
885
|
-
i, named1 = t
|
|
886
|
-
return i < 10
|
|
887
|
-
|
|
888
|
-
def loss_fun(named1):
|
|
889
|
-
i, named1 = eqxi.while_loop(cond_fun, body_fun, (0, named1), kind="checkpointed", max_steps=10)
|
|
890
|
-
return named1.sum().scalar()
|
|
891
|
-
|
|
892
|
-
grad_fun = eqx.filter_value_and_grad(loss_fun)
|
|
893
|
-
|
|
894
|
-
grad_fun(named1)
|
|
895
|
-
|
|
896
|
-
|
|
897
|
-
def test_at_for_in_placeish():
|
|
898
|
-
H = Axis("H", 10)
|
|
899
|
-
W = Axis("W", 20)
|
|
900
|
-
|
|
901
|
-
named1 = hax.random.uniform(PRNGKey(0), (H, W))
|
|
902
|
-
|
|
903
|
-
named1_at = named1.at[H, 0].set(0)
|
|
904
|
-
|
|
905
|
-
assert jnp.all(jnp.equal(named1_at[H, 0].array, 0))
|
|
906
|
-
assert jnp.all(named1_at[H, 1:].array == named1[H, 1:].array)
|
|
907
|
-
|
|
908
|
-
# test add, multiply, power, etc.
|
|
909
|
-
named1_at = named1.at[H, 0].add(1)
|
|
910
|
-
assert jnp.all(named1_at.array == named1.array.at[0].add(1))
|
|
911
|
-
|
|
912
|
-
named1_at = named1.at[H, 0].multiply(2)
|
|
913
|
-
assert jnp.all(named1_at.array == named1.array.at[0].multiply(2))
|
|
914
|
-
|
|
915
|
-
named1_at = named1.at[H, 0].power(2)
|
|
916
|
-
assert jnp.all(named1_at.array == named1.array.at[0].power(2))
|
|
917
|
-
|
|
918
|
-
named1_at = named1.at[H, 0].divide(2)
|
|
919
|
-
assert jnp.all(named1_at.array == named1.array.at[0].divide(2))
|
|
920
|
-
|
|
921
|
-
named1_at = named1.at[H, 0].apply(hax.square)
|
|
922
|
-
assert jnp.all(named1_at.array == named1.array.at[0].apply(jnp.square))
|
|
923
|
-
|
|
924
|
-
named1_at = named1.at[H, 0].max(0.5)
|
|
925
|
-
assert jnp.all(named1_at.array == named1.array.at[0].max(0.5))
|
|
926
|
-
|
|
927
|
-
named1_at = named1.at[H, 0].min(0.5)
|
|
928
|
-
assert jnp.all(named1_at.array == named1.array.at[0].min(0.5))
|
|
929
|
-
|
|
930
|
-
|
|
931
|
-
def test_at_with_fancy_indexing():
|
|
932
|
-
H = Axis("H", 10)
|
|
933
|
-
W = Axis("W", 20)
|
|
934
|
-
I0 = Axis("I0", 5)
|
|
935
|
-
I1 = Axis("I1", 5)
|
|
936
|
-
|
|
937
|
-
named1 = hax.random.uniform(PRNGKey(0), (H, W))
|
|
938
|
-
ind1 = hax.random.randint(PRNGKey(0), (I0,), 0, H.size)
|
|
939
|
-
ind2 = hax.random.randint(PRNGKey(0), (I1,), 0, W.size)
|
|
940
|
-
|
|
941
|
-
named1_at = named1.at[H, ind1].set(0)
|
|
942
|
-
assert jnp.all(named1_at.array == named1.array.at[ind1.array].set(0))
|
|
943
|
-
|
|
944
|
-
named1_at = named1.at[H, ind1].add(1, mode="clip")
|
|
945
|
-
assert jnp.all(named1_at.array == named1.array.at[ind1.array].add(1, mode="clip"))
|
|
946
|
-
|
|
947
|
-
named1_at = named1.at[H, ind1, W, ind2].set(0)
|
|
948
|
-
assert jnp.all(named1_at.array == named1.array.at[ind1.array.reshape(-1, 1), ind2.array.reshape(1, -1)].set(0))
|
|
949
|
-
|
|
950
|
-
# dslices
|
|
951
|
-
from haliax import ds
|
|
952
|
-
|
|
953
|
-
named1_at = named1.at[H, ds(3, 5)].set(0)
|
|
954
|
-
assert jnp.all(named1_at.array == named1.array.at[3:8].set(0))
|
|
955
|
-
|
|
956
|
-
named1_at = named1.at[H, ds(3, 5), W, ind2].power(2)
|
|
957
|
-
assert jnp.all(named1_at.array == named1.array.at[3:8, ind2.array].power(2))
|
|
958
|
-
|
|
959
|
-
|
|
960
|
-
def test_slice_dslice_and_array():
|
|
961
|
-
H = Axis("H", 10)
|
|
962
|
-
W = Axis("W", 20)
|
|
963
|
-
I0 = Axis("I0", 5)
|
|
964
|
-
|
|
965
|
-
named1 = hax.random.uniform(PRNGKey(0), (H, W))
|
|
966
|
-
ind2 = hax.random.randint(PRNGKey(0), (I0,), 0, W.size)
|
|
967
|
-
|
|
968
|
-
from haliax import ds
|
|
969
|
-
|
|
970
|
-
named1.array.at[3:8, ind2.array].add(jnp.full((5, 5), 2))
|
|
971
|
-
|
|
972
|
-
named1_at = named1.at[H, ds(3, 5), W, ind2].add(2)
|
|
973
|
-
assert jnp.all(named1_at.array == named1.array.at[3:8, ind2.array].add(2))
|
|
@@ -93,8 +93,7 @@ def test_alibi_attention_compared_to_hf():
|
|
|
93
93
|
import torch
|
|
94
94
|
from transformers.models.bloom.modeling_bloom import build_alibi_tensor
|
|
95
95
|
|
|
96
|
-
L = hax.
|
|
97
|
-
H = hax.Axis("NumHeads", 16)
|
|
96
|
+
L, H = hax.make_axes(L=1, H=16)
|
|
98
97
|
|
|
99
98
|
# Returns tensor shaped (batch_size * num_heads, 1, max_seq_len)
|
|
100
99
|
torch_tensor = (
|
|
@@ -107,7 +106,7 @@ def test_alibi_attention_compared_to_hf():
|
|
|
107
106
|
|
|
108
107
|
|
|
109
108
|
def test_fcm_attention_mask():
|
|
110
|
-
KeyPos = hax.
|
|
109
|
+
KeyPos, QueryPos, Head = hax.make_axes(KeyPos=20, QueryPos=10, Head=8)
|
|
111
110
|
|
|
112
111
|
mask = forgetful_causal_mask(KeyPos, mask_prob=0.6, sample_prob=False, key=PRNGKey(0))
|
|
113
112
|
|
|
@@ -116,9 +115,6 @@ def test_fcm_attention_mask():
|
|
|
116
115
|
|
|
117
116
|
assert mask.astype(float).sum().item() <= KeyPos.size
|
|
118
117
|
|
|
119
|
-
QueryPos = hax.Axis("QueryPos", 10)
|
|
120
|
-
Head = hax.Axis("Head", 8)
|
|
121
|
-
|
|
122
118
|
query = hax.arange(QueryPos).broadcast_axis(Head)
|
|
123
119
|
key = hax.arange(KeyPos).broadcast_axis(Head)
|
|
124
120
|
|
|
@@ -1,12 +1,10 @@
|
|
|
1
1
|
import pytest
|
|
2
2
|
|
|
3
|
-
from haliax.axis import Axis, eliminate_axes, rearrange_for_partial_order
|
|
3
|
+
from haliax.axis import Axis, eliminate_axes, make_axes, rearrange_for_partial_order
|
|
4
4
|
|
|
5
5
|
|
|
6
6
|
def test_eliminate_axes():
|
|
7
|
-
H =
|
|
8
|
-
W = Axis("W", 4)
|
|
9
|
-
C = Axis("C", 5)
|
|
7
|
+
H, W, C = make_axes(H=3, W=4, C=5)
|
|
10
8
|
|
|
11
9
|
assert eliminate_axes((H, W), (H,)) == (W,)
|
|
12
10
|
assert eliminate_axes((H, W), (W,)) == (H,)
|
|
@@ -273,12 +273,7 @@ def test_einsum_various_errors():
|
|
|
273
273
|
|
|
274
274
|
|
|
275
275
|
def test_einsum_examples():
|
|
276
|
-
|
|
277
|
-
Batch = hax.Axis("batch", 32)
|
|
278
|
-
Embed = hax.Axis("embed", 64)
|
|
279
|
-
H = hax.Axis("h", 16)
|
|
280
|
-
W = hax.Axis("w", 16)
|
|
281
|
-
C = hax.Axis("c", 3)
|
|
276
|
+
Batch, Embed, H, W, C = hax.make_axes(batch=32, embed=64, h=16, w=16, c=3)
|
|
282
277
|
|
|
283
278
|
# for jax
|
|
284
279
|
im = jnp.zeros((32, 16, 16, 3))
|
|
@@ -46,8 +46,7 @@ def test_dropout():
|
|
|
46
46
|
|
|
47
47
|
|
|
48
48
|
def test_one_hot():
|
|
49
|
-
i =
|
|
50
|
-
c = Axis("c", 3)
|
|
49
|
+
i, c = hax.make_axes(i=3, c=3)
|
|
51
50
|
actual = hax.nn.one_hot(hax.NamedArray(jnp.array([0, 1, 2]), (i,)), c)
|
|
52
51
|
expected = jnp.array([[1.0, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 1.0]])
|
|
53
52
|
|
|
@@ -61,16 +60,14 @@ def test_one_hot():
|
|
|
61
60
|
|
|
62
61
|
|
|
63
62
|
def test_one_hot_out_of_bound():
|
|
64
|
-
i =
|
|
65
|
-
c = Axis("c", 3)
|
|
63
|
+
i, c = hax.make_axes(i=2, c=3)
|
|
66
64
|
actual = hax.nn.one_hot(hax.NamedArray(jnp.array([-1, 3]), (i,)), c)
|
|
67
65
|
expected = jnp.array([[0.0, 0.0, 0.0], [0.0, 0.0, 0.0]])
|
|
68
66
|
assert jnp.all(jnp.isclose(actual.array, expected))
|
|
69
67
|
|
|
70
68
|
|
|
71
69
|
def test_standardize():
|
|
72
|
-
b =
|
|
73
|
-
c = Axis("c", 3)
|
|
70
|
+
b, c = hax.make_axes(b=2, c=3)
|
|
74
71
|
actual = hax.nn.standardize(hax.NamedArray(jnp.array([0, 1, 2]), (c,)), c)
|
|
75
72
|
expected = jax.nn.standardize(jnp.array([0, 1, 2]), axis=0)
|
|
76
73
|
|
|
@@ -113,11 +110,7 @@ def test_standardize():
|
|
|
113
110
|
@pytest.mark.parametrize("depth", [0, 1, 2, 3, 4, 5])
|
|
114
111
|
def test_mlp(depth):
|
|
115
112
|
key = jrandom.PRNGKey(0)
|
|
116
|
-
H =
|
|
117
|
-
C = Axis("C", 12)
|
|
118
|
-
W = Axis("W", 14)
|
|
119
|
-
|
|
120
|
-
E = Axis("E", 16)
|
|
113
|
+
H, C, W, E = hax.make_axes(H=10, C=12, W=14, E=16)
|
|
121
114
|
|
|
122
115
|
hax_mlp = hax.nn.MLP.init((H, C, W), E, width=8, depth=depth, key=key)
|
|
123
116
|
x = hax.random.uniform(key, (H, C, W))
|
|
@@ -150,10 +143,7 @@ def test_mlp(depth):
|
|
|
150
143
|
|
|
151
144
|
|
|
152
145
|
def test_linear_has_no_function_leaves_by_default():
|
|
153
|
-
H =
|
|
154
|
-
C = Axis("C", 12)
|
|
155
|
-
W = Axis("W", 14)
|
|
156
|
-
E = Axis("E", 16)
|
|
146
|
+
H, C, W, E = hax.make_axes(H=10, C=12, W=14, E=16)
|
|
157
147
|
|
|
158
148
|
hax_linear = hax.nn.Linear.init((H, C, W), E, key=jrandom.PRNGKey(0))
|
|
159
149
|
assert all(not isinstance(v, Callable) for v in jax.tree_util.tree_leaves(hax_linear)) # type: ignore
|
|
@@ -7,11 +7,7 @@ from haliax._src.rearrange import einops_rearrange
|
|
|
7
7
|
|
|
8
8
|
|
|
9
9
|
# some axes
|
|
10
|
-
W =
|
|
11
|
-
H = Axis("H", 6)
|
|
12
|
-
C = Axis("C", 3)
|
|
13
|
-
D = Axis("D", 2)
|
|
14
|
-
B = Axis("B", 5)
|
|
10
|
+
H, W, C, D, B = hax.make_axes(H=6, W=4, C=3, D=2, B=5)
|
|
15
11
|
Q = hax.Axis("Q", B.size * H.size)
|
|
16
12
|
E = Axis("E", H.size * W.size * D.size)
|
|
17
13
|
|
|
@@ -1,14 +1,13 @@
|
|
|
1
1
|
import jax
|
|
2
2
|
import jax.numpy as jnp
|
|
3
3
|
|
|
4
|
+
import haliax
|
|
4
5
|
import haliax.specialized_fns as hfns
|
|
5
|
-
from haliax import
|
|
6
|
+
from haliax import NamedArray
|
|
6
7
|
|
|
7
8
|
|
|
8
9
|
def test_top_k():
|
|
9
|
-
H =
|
|
10
|
-
W = Axis("W", 6)
|
|
11
|
-
D = Axis("D", 7)
|
|
10
|
+
H, W, D = haliax.make_axes(H=3, W=4, D=5)
|
|
12
11
|
|
|
13
12
|
rand = jax.random.uniform(jax.random.PRNGKey(0), (H.size, W.size, D.size))
|
|
14
13
|
n_rand = NamedArray(rand, (H, W, D))
|
|
@@ -1 +0,0 @@
|
|
|
1
|
-
__version__ = "1.4.dev315"
|
|
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
|