haliax 1.4.dev363__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.dev363 → haliax-1.4.dev364}/PKG-INFO +1 -1
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/api.md +0 -1
- {haliax-1.4.dev363 → haliax-1.4.dev364}/pyproject.toml +16 -0
- haliax-1.4.dev364/src/haliax/__about__.py +1 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/__init__.py +5 -3
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/dot.py +3 -2
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/axis.py +306 -147
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/core.py +170 -114
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/hof.py +6 -3
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/activations.py +1 -1
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/attention.py +3 -52
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/conv.py +14 -5
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/embedding.py +2 -2
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/pool.py +10 -11
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/partitioning.py +2 -2
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/random.py +35 -103
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/wrap.py +7 -6
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/core_test.py +44 -4
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_attention.py +0 -22
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_axis.py +106 -15
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_random.py +35 -21
- haliax-1.4.dev363/src/haliax/__about__.py +0 -1
- {haliax-1.4.dev363 → haliax-1.4.dev364}/.coveragerc +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/.flake8 +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/.gitignore +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/LICENSE +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/README.md +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/css/material.css +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/faq.md +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/fp8.md +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/index.md +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/indexing.md +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/matmul.md +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/nn.md +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/partitioning.md +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/rearrange.md +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/requirements.txt +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/scan.md +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/state-dict.md +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/tutorial.md +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/typing.md +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/docs/vmap.md +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/mkdocs.yml +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/linear.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/types.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/typing.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/src/haliax/util.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_conv.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_debug.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_dot.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_hof.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_int8.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_nn.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_ops.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_pool.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_scan.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev363 → haliax-1.4.dev364}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev363 → 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/
|
|
@@ -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,6 +39,7 @@ 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 (
|
|
43
45
|
Named,
|
|
@@ -105,8 +107,8 @@ def full(shape: AxisSpec, fill_value: T, dtype: Optional[DTypeLike] = None) -> N
|
|
|
105
107
|
if isinstance(shape, Axis):
|
|
106
108
|
return NamedArray(jnp.full(shape=shape.size, fill_value=fill_value, dtype=dtype), (shape,))
|
|
107
109
|
else:
|
|
108
|
-
x_shape =
|
|
109
|
-
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)
|
|
110
112
|
|
|
111
113
|
|
|
112
114
|
def zeros_like(a: NamedArray, dtype=None) -> NamedArray:
|
|
@@ -155,7 +157,7 @@ def arange(axis: AxisSpec, *, start=0, step=1, dtype: Optional[DTypeLike] = None
|
|
|
155
157
|
|
|
156
158
|
arr = jax.lax.iota(dtype=dtype or jnp.result_type(start), size=size) * step + start
|
|
157
159
|
arr = arr.reshape(to_jax_shape(axis))
|
|
158
|
-
return NamedArray(arr,
|
|
160
|
+
return NamedArray(arr, axis_spec_to_tuple(axis))
|
|
159
161
|
|
|
160
162
|
|
|
161
163
|
# TODO: add overrides for arraylike start/stop to linspace, logspace, geomspace
|
|
@@ -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(
|