haliax 1.4.dev362__tar.gz → 1.4.dev364__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.dev362 → haliax-1.4.dev364}/PKG-INFO +1 -1
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/api.md +0 -1
- haliax-1.4.dev364/docs/typing.md +61 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/mkdocs.yml +1 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/pyproject.toml +16 -0
- haliax-1.4.dev364/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/__init__.py +11 -3
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/dot.py +3 -2
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/axis.py +306 -147
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/core.py +343 -115
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/hof.py +6 -3
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/activations.py +1 -1
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/attention.py +3 -52
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/conv.py +14 -5
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/embedding.py +2 -2
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/pool.py +10 -11
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/partitioning.py +2 -2
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/random.py +35 -103
- haliax-1.4.dev364/src/haliax/typing.py +88 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/wrap.py +7 -6
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/core_test.py +44 -4
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_attention.py +0 -22
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_axis.py +106 -15
- haliax-1.4.dev364/tests/test_dtype_typing.py +41 -0
- haliax-1.4.dev364/tests/test_namedarray_typing.py +64 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_random.py +35 -21
- haliax-1.4.dev362/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev362 → haliax-1.4.dev364}/.coveragerc +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/.flake8 +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/.gitignore +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/LICENSE +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/README.md +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/css/material.css +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/faq.md +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/fp8.md +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/index.md +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/indexing.md +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/matmul.md +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/nn.md +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/partitioning.md +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/rearrange.md +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/requirements.txt +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/scan.md +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/state-dict.md +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/tutorial.md +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/docs/vmap.md +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/types.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/src/haliax/util.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_conv.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_debug.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_dot.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_hof.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_int8.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_nn.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_ops.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_pool.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_scan.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev362 → haliax-1.4.dev364}/tests/test_utils.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: haliax
|
|
3
|
-
Version: 1.4.
|
|
3
|
+
Version: 1.4.dev364
|
|
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,61 @@
|
|
|
1
|
+
# NamedArray Type Annotations
|
|
2
|
+
|
|
3
|
+
Haliax supports a lightweight syntax for specifying the axes of a `NamedArray`
|
|
4
|
+
in type annotations. Internally, `Named[...]` expands to
|
|
5
|
+
`Annotated[NamedArray, axes]`, so it works well with static type checkers like
|
|
6
|
+
``mypy``. The syntax mirrors normal indexing with axis names. Some examples:
|
|
7
|
+
|
|
8
|
+
```python
|
|
9
|
+
from haliax import Named
|
|
10
|
+
|
|
11
|
+
arr: Named["batch", "embed"]
|
|
12
|
+
arr: Named["batch embed ..."] # starts with these axes
|
|
13
|
+
arr: Named["... embed"] # ends with this axis
|
|
14
|
+
arr: Named["batch ... embed"] # contains these axes in order
|
|
15
|
+
arr: Named[{"batch", "embed"}] # has exactly these axes, order ignored
|
|
16
|
+
arr: Named[{"batch", "embed", ...}] # has at least these axes
|
|
17
|
+
```
|
|
18
|
+
|
|
19
|
+
At runtime you can verify that a `NamedArray` conforms to a particular
|
|
20
|
+
annotation using `matches_axes`:
|
|
21
|
+
|
|
22
|
+
```python
|
|
23
|
+
if not arr.matches_axes(Named["batch embed ..."]):
|
|
24
|
+
raise ValueError("unexpected axes")
|
|
25
|
+
```
|
|
26
|
+
|
|
27
|
+
## DType-aware annotations
|
|
28
|
+
|
|
29
|
+
Sometimes it is useful to express both the axes **and** the dtype in the type
|
|
30
|
+
annotation. The :mod:`haliax.typing` module defines symbolic types for all of
|
|
31
|
+
JAX's common dtypes that can be indexed just like ``Named``. In documentation
|
|
32
|
+
examples we'll use ``import haliax.typing as ht``:
|
|
33
|
+
|
|
34
|
+
```python
|
|
35
|
+
import haliax.typing as ht
|
|
36
|
+
|
|
37
|
+
def foo(x: ht.f32["batch"]):
|
|
38
|
+
...
|
|
39
|
+
|
|
40
|
+
def bar(x: ht.i32["batch"]):
|
|
41
|
+
...
|
|
42
|
+
```
|
|
43
|
+
|
|
44
|
+
For convenience the module also provides aggregate categories ``Float``,
|
|
45
|
+
``Complex``, ``Int`` and ``UInt`` that match any floating point, complex,
|
|
46
|
+
signed integer or unsigned integer dtype respectively:
|
|
47
|
+
|
|
48
|
+
```python
|
|
49
|
+
def baz(x: ht.Float["batch"]):
|
|
50
|
+
...
|
|
51
|
+
```
|
|
52
|
+
|
|
53
|
+
At runtime ``matches_axes`` also checks the dtype when one is present:
|
|
54
|
+
|
|
55
|
+
```python
|
|
56
|
+
from haliax import Axis, zeros
|
|
57
|
+
import haliax.typing as ht
|
|
58
|
+
|
|
59
|
+
arr = zeros({"batch": 4})
|
|
60
|
+
assert arr.matches_axes(ht.f32["batch"]) # dtype and axes both match
|
|
61
|
+
```
|
|
@@ -71,3 +71,19 @@ src_paths = ["src", "tests"]
|
|
|
71
71
|
"Homepage" = "https://github.com/stanford-crfm/haliax"
|
|
72
72
|
"Bug Tracker" = "https://github.com/stanford-crfm/haliax/issues/"
|
|
73
73
|
"Documentation" = "https://haliax.readthedocs.io/en/latest/"
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
[tool.coverage.report]
|
|
77
|
+
exclude_also = [
|
|
78
|
+
"def __repr__",
|
|
79
|
+
"if self.debug:",
|
|
80
|
+
"if settings.DEBUG",
|
|
81
|
+
"raise AssertionError",
|
|
82
|
+
"raise NotImplementedError",
|
|
83
|
+
"if 0:",
|
|
84
|
+
"if __name__ == .__main__.:",
|
|
85
|
+
"if TYPE_CHECKING:",
|
|
86
|
+
"class .*\\bProtocol\\):",
|
|
87
|
+
"@(abc\\.)?abstractmethod",
|
|
88
|
+
"[.][.][.]"
|
|
89
|
+
]
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "1.4.dev364"
|
|
@@ -29,6 +29,7 @@ from .axis import (
|
|
|
29
29
|
AxisSpec,
|
|
30
30
|
axis_name,
|
|
31
31
|
axis_size,
|
|
32
|
+
axis_spec_to_tuple,
|
|
32
33
|
concat_axes,
|
|
33
34
|
dblock,
|
|
34
35
|
ds,
|
|
@@ -38,9 +39,13 @@ from .axis import (
|
|
|
38
39
|
replace_axis,
|
|
39
40
|
resolve_axis,
|
|
40
41
|
selects_axis,
|
|
42
|
+
to_jax_shape,
|
|
41
43
|
)
|
|
42
44
|
from .core import (
|
|
45
|
+
Named,
|
|
43
46
|
NamedArray,
|
|
47
|
+
NamedArrayAxes,
|
|
48
|
+
NamedArrayAxesSpec,
|
|
44
49
|
NamedOrNumeric,
|
|
45
50
|
are_shape_checks_enabled,
|
|
46
51
|
broadcast_arrays,
|
|
@@ -102,8 +107,8 @@ def full(shape: AxisSpec, fill_value: T, dtype: Optional[DTypeLike] = None) -> N
|
|
|
102
107
|
if isinstance(shape, Axis):
|
|
103
108
|
return NamedArray(jnp.full(shape=shape.size, fill_value=fill_value, dtype=dtype), (shape,))
|
|
104
109
|
else:
|
|
105
|
-
x_shape =
|
|
106
|
-
return NamedArray(jnp.full(shape=x_shape, fill_value=fill_value, dtype=dtype),
|
|
110
|
+
x_shape = to_jax_shape(shape)
|
|
111
|
+
return NamedArray(jnp.full(shape=x_shape, fill_value=fill_value, dtype=dtype), shape)
|
|
107
112
|
|
|
108
113
|
|
|
109
114
|
def zeros_like(a: NamedArray, dtype=None) -> NamedArray:
|
|
@@ -152,7 +157,7 @@ def arange(axis: AxisSpec, *, start=0, step=1, dtype: Optional[DTypeLike] = None
|
|
|
152
157
|
|
|
153
158
|
arr = jax.lax.iota(dtype=dtype or jnp.result_type(start), size=size) * step + start
|
|
154
159
|
arr = arr.reshape(to_jax_shape(axis))
|
|
155
|
-
return NamedArray(arr,
|
|
160
|
+
return NamedArray(arr, axis_spec_to_tuple(axis))
|
|
156
161
|
|
|
157
162
|
|
|
158
163
|
# TODO: add overrides for arraylike start/stop to linspace, logspace, geomspace
|
|
@@ -926,6 +931,9 @@ __all__ = [
|
|
|
926
931
|
"make_axes",
|
|
927
932
|
"axis_name",
|
|
928
933
|
"axis_size",
|
|
934
|
+
"NamedArrayAxesSpec",
|
|
935
|
+
"NamedArrayAxes",
|
|
936
|
+
"Named",
|
|
929
937
|
"NamedArray",
|
|
930
938
|
"broadcast_to",
|
|
931
939
|
"broadcast_axis",
|
|
@@ -12,6 +12,7 @@ from haliax.axis import (
|
|
|
12
12
|
AxisSelection,
|
|
13
13
|
PartialAxisSpec,
|
|
14
14
|
axis_name,
|
|
15
|
+
axis_spec_to_shape_dict,
|
|
15
16
|
eliminate_axes,
|
|
16
17
|
rearrange_for_partial_order,
|
|
17
18
|
union_axes,
|
|
@@ -140,8 +141,8 @@ def dot(
|
|
|
140
141
|
if axis is None:
|
|
141
142
|
jax_str = f"contract {', '.join(axis_name(ax) for ax in all_axes)} -> <scalar>"
|
|
142
143
|
else:
|
|
143
|
-
axis =
|
|
144
|
-
jax_str = f"contract {', '.join(
|
|
144
|
+
axis = axis_spec_to_shape_dict(axis)
|
|
145
|
+
jax_str = f"contract {', '.join(axis)} -> {', '.join(a.name for a in output_axes)}"
|
|
145
146
|
|
|
146
147
|
with jax.named_scope(jax_str):
|
|
147
148
|
output = _jittable_dg_einsum(
|