haliax 1.4.dev441__tar.gz → 1.4.dev443__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.dev441 → haliax-1.4.dev443}/PKG-INFO +1 -1
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/__about__.py +1 -1
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/embedding.py +7 -2
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/linear.py +7 -2
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_mup_embedding.py +1 -1
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_mup_linear.py +10 -5
- {haliax-1.4.dev441 → haliax-1.4.dev443}/.agents/projects/api_parity.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/.coveragerc +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/.flake8 +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/.gitignore +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/.playbooks/wrap-non-named.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/AGENTS.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/AUTHORS.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/CONTRIBUTORS.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/LICENSE +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/README.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/api.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/css/material.css +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/faq.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/fp8.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/index.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/indexing.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/matmul.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/nn.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/partitioning.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/primer.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/rearrange.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/requirements.txt +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/scan.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/state-dict.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/tutorial.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/typing.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/vmap.md +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/etc/license_header.txt +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/mkdocs.yml +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/pyproject.toml +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/core.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/fft.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/field.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/haxtyping.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/mup.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/poly.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/random.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/tree.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/types.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/util.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/core_test.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_attention.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_axis.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_bitwise_ops.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_conv.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_debug.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_dot.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_fft.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_field.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_hof.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_int8.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_moe_linear.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_mup_coordinate_check.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_nan_reductions.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_nn.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_ops.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_poly_ops.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_pool.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_random.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_scan.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_utils.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_visualize_sharding.py +0 -0
- {haliax-1.4.dev441 → haliax-1.4.dev443}/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.dev443
|
|
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/
|
|
@@ -24,7 +24,12 @@ class Embedding(eqx.Module, ReparamEnabled):
|
|
|
24
24
|
# axes
|
|
25
25
|
Vocab: Axis = eqx.field(static=True)
|
|
26
26
|
Embed: AxisSpec = eqx.field(static=True)
|
|
27
|
-
|
|
27
|
+
|
|
28
|
+
_reparam_cls: type[AbstractEmbeddingReparam] = eqx.field(static=True, default=EmbeddingStandardParam)
|
|
29
|
+
|
|
30
|
+
@property
|
|
31
|
+
def reparam(self) -> AbstractEmbeddingReparam:
|
|
32
|
+
return self._reparam_cls(self.Embed, self.Vocab)
|
|
28
33
|
|
|
29
34
|
@staticmethod
|
|
30
35
|
def init(
|
|
@@ -61,7 +66,7 @@ class Embedding(eqx.Module, ReparamEnabled):
|
|
|
61
66
|
weight = hax.random.truncated_normal(key, all_axes, -3, 3) * (
|
|
62
67
|
init_scale * reparam_cls.init_scale(Vocab, Embed)
|
|
63
68
|
)
|
|
64
|
-
return Embedding(weight=weight, Vocab=Vocab, Embed=Embed,
|
|
69
|
+
return Embedding(weight=weight, Vocab=Vocab, Embed=Embed, _reparam_cls=reparam_cls)
|
|
65
70
|
|
|
66
71
|
def __call__(self, input_ids: NamedArray, *, key: PRNGKeyArray | None = None):
|
|
67
72
|
"""Alias for `embed`. key is ignored."""
|
|
@@ -45,9 +45,14 @@ class Linear(ModuleWithStateDictSerialization, ReparamEnabled):
|
|
|
45
45
|
|
|
46
46
|
In: AxisSpec = eqx.field(static=True)
|
|
47
47
|
Out: AxisSpec = eqx.field(static=True)
|
|
48
|
-
reparam: AbstractLinearReparam = eqx.field(static=True)
|
|
49
48
|
dot_general: DotGeneralOp = eqx.field(default_factory=DotGeneralOp.default)
|
|
50
49
|
|
|
50
|
+
_reparam_cls: type[AbstractLinearReparam] = eqx.field(static=True, default=LinearStandardParam)
|
|
51
|
+
|
|
52
|
+
@property
|
|
53
|
+
def reparam(self) -> AbstractLinearReparam:
|
|
54
|
+
return self._reparam_cls(self.In, self.Out)
|
|
55
|
+
|
|
51
56
|
@staticmethod
|
|
52
57
|
def init(
|
|
53
58
|
In: AxisSpec,
|
|
@@ -77,7 +82,7 @@ class Linear(ModuleWithStateDictSerialization, ReparamEnabled):
|
|
|
77
82
|
if dot_general is None:
|
|
78
83
|
dot_general = DotGeneralOp.default()
|
|
79
84
|
|
|
80
|
-
return Linear(weight, bias, In, Out, dot_general=dot_general,
|
|
85
|
+
return Linear(weight, bias, In, Out, dot_general=dot_general, _reparam_cls=reparam_cls)
|
|
81
86
|
|
|
82
87
|
@named_call
|
|
83
88
|
def __call__(self, inputs, *, key: PRNGKeyArray | None = None):
|
|
@@ -31,7 +31,7 @@ def test_mup_embedding_unembedding_scale():
|
|
|
31
31
|
Embed = (hax.Axis("E", 3),)
|
|
32
32
|
|
|
33
33
|
weight = hax.ones(hax.concat_axis_specs(Vocab, Embed))
|
|
34
|
-
layer = Embedding(weight=weight, Vocab=Vocab, Embed=Embed,
|
|
34
|
+
layer = Embedding(weight=weight, Vocab=Vocab, Embed=Embed, _reparam_cls=EmbeddingMup)
|
|
35
35
|
|
|
36
36
|
scale = layer.reparam.unembed_active_scale
|
|
37
37
|
assert scale == pytest.approx(1.0 / hax.axis_size(Embed))
|
|
@@ -11,7 +11,12 @@ import pytest
|
|
|
11
11
|
|
|
12
12
|
import haliax as hax
|
|
13
13
|
from haliax.nn import Linear
|
|
14
|
-
from haliax.nn.mup import
|
|
14
|
+
from haliax.nn.mup import (
|
|
15
|
+
InputLinearMup,
|
|
16
|
+
LinearStandardParam,
|
|
17
|
+
HiddenLinearMup,
|
|
18
|
+
OutputLinearMup,
|
|
19
|
+
)
|
|
15
20
|
|
|
16
21
|
|
|
17
22
|
@pytest.mark.parametrize("out_first", [True, False])
|
|
@@ -37,8 +42,8 @@ def test_mup_linear_call_matches_linear():
|
|
|
37
42
|
weight = hax.ones(hax.concat_axis_specs(Out, In)) * 0.5
|
|
38
43
|
bias = hax.full(Out, 0.25)
|
|
39
44
|
|
|
40
|
-
linear = Linear(weight, bias, In, Out,
|
|
41
|
-
mup = Linear(weight, bias, In, Out,
|
|
45
|
+
linear = Linear(weight, bias, In, Out, _reparam_cls=LinearStandardParam)
|
|
46
|
+
mup = Linear(weight, bias, In, Out, _reparam_cls=InputLinearMup)
|
|
42
47
|
|
|
43
48
|
inputs = hax.full(hax.concat_axis_specs(Batch, In), 2.0)
|
|
44
49
|
|
|
@@ -109,8 +114,8 @@ def test_input_linear_behaves_like_base_linear():
|
|
|
109
114
|
weight = hax.ones((Out, In)) * 0.1
|
|
110
115
|
bias = hax.zeros(Out)
|
|
111
116
|
|
|
112
|
-
linear = Linear(weight, bias, In, Out,
|
|
113
|
-
input_linear = Linear(weight, bias, In, Out,
|
|
117
|
+
linear = Linear(weight, bias, In, Out, _reparam_cls=LinearStandardParam)
|
|
118
|
+
input_linear = Linear(weight, bias, In, Out, _reparam_cls=InputLinearMup)
|
|
114
119
|
|
|
115
120
|
inputs = hax.random.normal(jrandom.PRNGKey(5), (Batch, In))
|
|
116
121
|
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|