haliax 1.4.dev420__tar.gz → 1.4.dev438__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.dev420 → haliax-1.4.dev438}/PKG-INFO +1 -1
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/__about__.py +1 -1
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/nn/embedding.py +22 -8
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/nn/linear.py +51 -7
- haliax-1.4.dev438/src/haliax/nn/mup.py +206 -0
- haliax-1.4.dev438/tests/test_mup_coordinate_check.py +164 -0
- haliax-1.4.dev438/tests/test_mup_embedding.py +48 -0
- haliax-1.4.dev438/tests/test_mup_linear.py +120 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/.agents/projects/api_parity.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/.coveragerc +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/.flake8 +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/.gitignore +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/.playbooks/wrap-non-named.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/AGENTS.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/AUTHORS.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/CONTRIBUTORS.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/LICENSE +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/README.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/api.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/css/material.css +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/faq.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/fp8.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/index.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/indexing.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/matmul.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/nn.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/partitioning.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/primer.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/rearrange.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/requirements.txt +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/scan.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/state-dict.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/tutorial.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/typing.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/docs/vmap.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/etc/license_header.txt +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/mkdocs.yml +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/pyproject.toml +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/__init__.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/core.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/fft.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/field.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/haxtyping.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/poly.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/random.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/types.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/util.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/core_test.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_attention.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_axis.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_bitwise_ops.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_conv.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_debug.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_dot.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_fft.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_field.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_hof.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_int8.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_moe_linear.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_nan_reductions.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_nn.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_ops.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_poly_ops.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_pool.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_random.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_scan.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_utils.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/tests/test_visualize_sharding.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev438}/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.dev438
|
|
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/
|
|
@@ -11,21 +11,31 @@ from jaxtyping import PRNGKeyArray
|
|
|
11
11
|
|
|
12
12
|
import haliax as hax
|
|
13
13
|
|
|
14
|
+
from .mup import AbstractEmbeddingReparam, ReparamEnabled, EmbeddingStandardParam
|
|
14
15
|
from ..axis import Axis, AxisSpec, concat_axes
|
|
15
16
|
from ..core import NamedArray
|
|
16
17
|
from ..jax_utils import named_call
|
|
17
18
|
from ..tree_util import resize_axis
|
|
18
19
|
|
|
19
20
|
|
|
20
|
-
class Embedding(eqx.Module):
|
|
21
|
+
class Embedding(eqx.Module, ReparamEnabled):
|
|
21
22
|
weight: NamedArray
|
|
22
23
|
|
|
23
24
|
# axes
|
|
24
25
|
Vocab: Axis = eqx.field(static=True)
|
|
25
26
|
Embed: AxisSpec = eqx.field(static=True)
|
|
27
|
+
reparam: AbstractEmbeddingReparam = eqx.field(static=True)
|
|
26
28
|
|
|
27
29
|
@staticmethod
|
|
28
|
-
def init(
|
|
30
|
+
def init(
|
|
31
|
+
Vocab: Axis,
|
|
32
|
+
Embed: AxisSpec,
|
|
33
|
+
*,
|
|
34
|
+
init_scale: float = 1,
|
|
35
|
+
key,
|
|
36
|
+
initializer_range: float | None = None,
|
|
37
|
+
reparam_cls: type[AbstractEmbeddingReparam] = EmbeddingStandardParam,
|
|
38
|
+
):
|
|
29
39
|
"""
|
|
30
40
|
Initialize an Embedding module.
|
|
31
41
|
|
|
@@ -41,13 +51,17 @@ class Embedding(eqx.Module):
|
|
|
41
51
|
initializer_range: Deprecated. Use init_scale instead.
|
|
42
52
|
"""
|
|
43
53
|
if initializer_range is not None:
|
|
44
|
-
warnings.warn(
|
|
54
|
+
warnings.warn(
|
|
55
|
+
"initializer_range is deprecated. Use init_std instead.",
|
|
56
|
+
DeprecationWarning,
|
|
57
|
+
)
|
|
45
58
|
init_scale = initializer_range
|
|
46
59
|
|
|
47
60
|
all_axes = concat_axes(Vocab, Embed)
|
|
48
|
-
|
|
49
|
-
|
|
50
|
-
|
|
61
|
+
weight = hax.random.truncated_normal(key, all_axes, -3, 3) * (
|
|
62
|
+
init_scale * reparam_cls.init_scale(Vocab, Embed)
|
|
63
|
+
)
|
|
64
|
+
return Embedding(weight=weight, Vocab=Vocab, Embed=Embed, reparam=reparam_cls(Embed, Vocab))
|
|
51
65
|
|
|
52
66
|
def __call__(self, input_ids: NamedArray, *, key: PRNGKeyArray | None = None):
|
|
53
67
|
"""Alias for `embed`. key is ignored."""
|
|
@@ -60,7 +74,7 @@ class Embedding(eqx.Module):
|
|
|
60
74
|
input_ids: token IDs with shape > {Vocab}
|
|
61
75
|
"""
|
|
62
76
|
input_embeds = self.weight.take(self.Vocab, input_ids)
|
|
63
|
-
return input_embeds
|
|
77
|
+
return input_embeds * self.reparam.active_scale
|
|
64
78
|
|
|
65
79
|
def unembed(self, input_embeds: NamedArray):
|
|
66
80
|
"""
|
|
@@ -68,7 +82,7 @@ class Embedding(eqx.Module):
|
|
|
68
82
|
|
|
69
83
|
Equivalent to `input_embeds.dot(self.weight, axis=self.Embed)`.
|
|
70
84
|
"""
|
|
71
|
-
return input_embeds.dot(self.weight, axis=self.Embed)
|
|
85
|
+
return input_embeds.dot(self.weight, axis=self.Embed) * self.reparam.unembed_active_scale
|
|
72
86
|
|
|
73
87
|
def resize_embeddings(self, new_size: int, key: PRNGKeyArray | None = None):
|
|
74
88
|
"""
|
|
@@ -6,6 +6,7 @@
|
|
|
6
6
|
import dataclasses
|
|
7
7
|
import math
|
|
8
8
|
from functools import partial
|
|
9
|
+
from typing import Optional
|
|
9
10
|
|
|
10
11
|
import equinox as eqx
|
|
11
12
|
import jax
|
|
@@ -17,7 +18,16 @@ from jaxtyping import PRNGKeyArray
|
|
|
17
18
|
|
|
18
19
|
import haliax as hax
|
|
19
20
|
|
|
20
|
-
|
|
21
|
+
|
|
22
|
+
from . import mup
|
|
23
|
+
from .mup import AbstractLinearReparam, ReparamEnabled, LinearStandardParam
|
|
24
|
+
from .._src.state_dict import (
|
|
25
|
+
Mod,
|
|
26
|
+
ModuleWithStateDictSerialization,
|
|
27
|
+
StateDict,
|
|
28
|
+
default_eqx_module_from_state_dict,
|
|
29
|
+
default_eqx_module_to_state_dict,
|
|
30
|
+
)
|
|
21
31
|
from ..axis import Axis, AxisSpec
|
|
22
32
|
from ..core import NamedArray
|
|
23
33
|
from ..jax_utils import named_call
|
|
@@ -26,7 +36,7 @@ from ..quantization import DotGeneralOp
|
|
|
26
36
|
from ..util import ensure_tuple
|
|
27
37
|
|
|
28
38
|
|
|
29
|
-
class Linear(ModuleWithStateDictSerialization):
|
|
39
|
+
class Linear(ModuleWithStateDictSerialization, ReparamEnabled):
|
|
30
40
|
"""A named Linear layer. This module allows you to specify multiple named axes for both input
|
|
31
41
|
and output, which is occasionally useful."""
|
|
32
42
|
|
|
@@ -35,6 +45,7 @@ class Linear(ModuleWithStateDictSerialization):
|
|
|
35
45
|
|
|
36
46
|
In: AxisSpec = eqx.field(static=True)
|
|
37
47
|
Out: AxisSpec = eqx.field(static=True)
|
|
48
|
+
reparam: AbstractLinearReparam = eqx.field(static=True)
|
|
38
49
|
dot_general: DotGeneralOp = eqx.field(default_factory=DotGeneralOp.default)
|
|
39
50
|
|
|
40
51
|
@staticmethod
|
|
@@ -47,6 +58,7 @@ class Linear(ModuleWithStateDictSerialization):
|
|
|
47
58
|
out_first: bool = True,
|
|
48
59
|
dot_general: DotGeneralOp | None = None,
|
|
49
60
|
init_scale: float = 1.0,
|
|
61
|
+
reparam_cls: type[AbstractLinearReparam] = LinearStandardParam,
|
|
50
62
|
) -> "Linear":
|
|
51
63
|
"""
|
|
52
64
|
Args:
|
|
@@ -59,14 +71,13 @@ class Linear(ModuleWithStateDictSerialization):
|
|
|
59
71
|
init_scale: float: The scale to use for initialization. We scale init by 1/sqrt(Input.size)*init_scale
|
|
60
72
|
"""
|
|
61
73
|
joint_spec = hax.concat_axis_specs(Out, In) if out_first else hax.concat_axis_specs(In, Out)
|
|
62
|
-
|
|
63
|
-
weight = hax.random.truncated_normal(key, joint_spec, -3, 3) * (init_scale / math.sqrt(input_size))
|
|
74
|
+
weight = hax.random.truncated_normal(key, joint_spec, -3, 3) * (init_scale * reparam_cls.init_scale(In, Out))
|
|
64
75
|
bias = hax.zeros(Out) if use_bias else None
|
|
65
76
|
|
|
66
77
|
if dot_general is None:
|
|
67
78
|
dot_general = DotGeneralOp.default()
|
|
68
79
|
|
|
69
|
-
return Linear(weight, bias, In, Out, dot_general=dot_general)
|
|
80
|
+
return Linear(weight, bias, In, Out, dot_general=dot_general, reparam=reparam_cls(In, Out))
|
|
70
81
|
|
|
71
82
|
@named_call
|
|
72
83
|
def __call__(self, inputs, *, key: PRNGKeyArray | None = None):
|
|
@@ -76,7 +87,11 @@ class Linear(ModuleWithStateDictSerialization):
|
|
|
76
87
|
key: Not used, but there for compat with other modules
|
|
77
88
|
"""
|
|
78
89
|
del key
|
|
79
|
-
q = inputs.dot(
|
|
90
|
+
q = inputs.dot(
|
|
91
|
+
self.weight * self.reparam.active_scale,
|
|
92
|
+
axis=self.In,
|
|
93
|
+
dot_general=self.dot_general,
|
|
94
|
+
)
|
|
80
95
|
q = hax.auto_sharded(q)
|
|
81
96
|
|
|
82
97
|
if self.bias is not None:
|
|
@@ -137,6 +152,32 @@ class Linear(ModuleWithStateDictSerialization):
|
|
|
137
152
|
else:
|
|
138
153
|
return self.weight.axes[-len(self.Out) :] != self.Out
|
|
139
154
|
|
|
155
|
+
def to_state_dict(self, prefix: Optional[str] = None) -> StateDict:
|
|
156
|
+
scaled = dataclasses.replace(self, weight=self.weight * self.reparam.active_scale)
|
|
157
|
+
return default_eqx_module_to_state_dict(scaled, prefix)
|
|
158
|
+
|
|
159
|
+
def from_state_dict(self: Mod, state_dict: StateDict, prefix: Optional[str] = None) -> Mod:
|
|
160
|
+
unscaled = default_eqx_module_from_state_dict(self, state_dict, prefix)
|
|
161
|
+
return dataclasses.replace(unscaled, weight=unscaled.weight / self.reparam.active_scale)
|
|
162
|
+
|
|
163
|
+
@staticmethod
|
|
164
|
+
def input_reparam(use_mup: bool = True) -> type[AbstractLinearReparam]:
|
|
165
|
+
"""Return the reparameterization class for an input linear layer."""
|
|
166
|
+
|
|
167
|
+
return mup.InputLinearMup if use_mup else mup.LinearStandardParam
|
|
168
|
+
|
|
169
|
+
@staticmethod
|
|
170
|
+
def hidden_reparam(use_mup: bool = True) -> type[AbstractLinearReparam]:
|
|
171
|
+
"""Return the reparameterization class for a hidden linear layer."""
|
|
172
|
+
|
|
173
|
+
return mup.HiddenLinearMup if use_mup else mup.LinearStandardParam
|
|
174
|
+
|
|
175
|
+
@staticmethod
|
|
176
|
+
def output_reparam(use_mup: bool = True) -> type[AbstractLinearReparam]:
|
|
177
|
+
"""Return the reparameterization class for an output linear layer."""
|
|
178
|
+
|
|
179
|
+
return mup.OutputLinearMup if use_mup else mup.LinearStandardParam
|
|
180
|
+
|
|
140
181
|
|
|
141
182
|
class MoELinear(eqx.Module):
|
|
142
183
|
"""A named Linear layer for MoE. This module allows you to specify multiple named axes for both input
|
|
@@ -197,7 +238,10 @@ class MoELinear(eqx.Module):
|
|
|
197
238
|
dim_numbers = jax.lax.RaggedDotDimensionNumbers(
|
|
198
239
|
dot_dimension_numbers=(
|
|
199
240
|
# contracting
|
|
200
|
-
(
|
|
241
|
+
(
|
|
242
|
+
ensure_tuple(inputs.axis_indices(self.In)),
|
|
243
|
+
ensure_tuple(self.weight.axis_indices(self.In)),
|
|
244
|
+
),
|
|
201
245
|
# batch
|
|
202
246
|
((), ()),
|
|
203
247
|
),
|
|
@@ -0,0 +1,206 @@
|
|
|
1
|
+
# Copyright 2025 The Levanter Authors
|
|
2
|
+
#
|
|
3
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
4
|
+
|
|
5
|
+
import math
|
|
6
|
+
from abc import ABC, abstractmethod
|
|
7
|
+
from dataclasses import dataclass
|
|
8
|
+
|
|
9
|
+
import haliax as hax
|
|
10
|
+
import equinox as eqx
|
|
11
|
+
|
|
12
|
+
from ..axis import AxisSpec
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class AbstractReparam(ABC):
|
|
16
|
+
"""Abstract base class for abc-parameterization rules.
|
|
17
|
+
|
|
18
|
+
Defines the interface for active scaling of parameters (a),
|
|
19
|
+
computing initialization scales (b), and learning rate scaling (c)
|
|
20
|
+
|
|
21
|
+
See: https://arxiv.org/abs/2011.14522
|
|
22
|
+
"""
|
|
23
|
+
|
|
24
|
+
@staticmethod
|
|
25
|
+
@abstractmethod
|
|
26
|
+
def init_scale(In: AxisSpec, Out: AxisSpec):
|
|
27
|
+
"""Return the scaling factor for initializing weights
|
|
28
|
+
given input and output axes."""
|
|
29
|
+
raise NotImplementedError
|
|
30
|
+
|
|
31
|
+
@property
|
|
32
|
+
@abstractmethod
|
|
33
|
+
def lr_scale(self):
|
|
34
|
+
"""Return the learning-rate scaling factor."""
|
|
35
|
+
raise NotImplementedError
|
|
36
|
+
|
|
37
|
+
@property
|
|
38
|
+
@abstractmethod
|
|
39
|
+
def active_scale(self):
|
|
40
|
+
"""Return the scaling applied to activations."""
|
|
41
|
+
raise NotImplementedError
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
@dataclass
|
|
45
|
+
class AbstractLinearReparam(AbstractReparam):
|
|
46
|
+
"""Base class for linear-layer reparameterizations.
|
|
47
|
+
|
|
48
|
+
Stores input and output axis specifications, and inherits
|
|
49
|
+
the reparameterization interface.
|
|
50
|
+
"""
|
|
51
|
+
|
|
52
|
+
In: AxisSpec
|
|
53
|
+
Out: AxisSpec
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
class LinearStandardParam(AbstractLinearReparam):
|
|
57
|
+
"""Standard (non-muP) parameterization for linear layers.
|
|
58
|
+
|
|
59
|
+
Uses the usual fan-in scaling for initialization and
|
|
60
|
+
leaves learning rate and activation scaling unchanged.
|
|
61
|
+
"""
|
|
62
|
+
|
|
63
|
+
@staticmethod
|
|
64
|
+
def init_scale(In: AxisSpec, Out: AxisSpec):
|
|
65
|
+
return 1 / math.sqrt(hax.axis_size(In))
|
|
66
|
+
|
|
67
|
+
@property
|
|
68
|
+
def active_scale(self):
|
|
69
|
+
return 1
|
|
70
|
+
|
|
71
|
+
@property
|
|
72
|
+
def lr_scale(self):
|
|
73
|
+
return 1
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
class InputLinearMup(AbstractLinearReparam):
|
|
77
|
+
"""muP-style parameterization for input linear layers.
|
|
78
|
+
|
|
79
|
+
Uses no scaling on initialization or learning rate.
|
|
80
|
+
See: https://arxiv.org/abs/2011.14522 (Maximal Update Parametrization)
|
|
81
|
+
"""
|
|
82
|
+
|
|
83
|
+
@staticmethod
|
|
84
|
+
def init_scale(In: AxisSpec, Out: AxisSpec):
|
|
85
|
+
return 1
|
|
86
|
+
|
|
87
|
+
@property
|
|
88
|
+
def active_scale(self):
|
|
89
|
+
return 1
|
|
90
|
+
|
|
91
|
+
@property
|
|
92
|
+
def lr_scale(self):
|
|
93
|
+
return 1
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
class HiddenLinearMup(AbstractLinearReparam):
|
|
97
|
+
"""muP-style parameterization for hidden linear layers.
|
|
98
|
+
|
|
99
|
+
Applies fan-in scaling at initialization and scales
|
|
100
|
+
learning rate inversely with layer width.
|
|
101
|
+
"""
|
|
102
|
+
|
|
103
|
+
@staticmethod
|
|
104
|
+
def init_scale(In: AxisSpec, Out: AxisSpec):
|
|
105
|
+
return 1 / math.sqrt(hax.axis_size(In))
|
|
106
|
+
|
|
107
|
+
@property
|
|
108
|
+
def active_scale(self):
|
|
109
|
+
return 1
|
|
110
|
+
|
|
111
|
+
@property
|
|
112
|
+
def lr_scale(self):
|
|
113
|
+
return 1 / hax.axis_size(self.In)
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
class OutputLinearMup(AbstractLinearReparam):
|
|
117
|
+
"""muP-style parameterization for output linear layers.
|
|
118
|
+
|
|
119
|
+
Uses unit initialization and applies inverse-width
|
|
120
|
+
scaling to the output activations.
|
|
121
|
+
"""
|
|
122
|
+
|
|
123
|
+
@staticmethod
|
|
124
|
+
def init_scale(In: AxisSpec, Out: AxisSpec):
|
|
125
|
+
return 1
|
|
126
|
+
|
|
127
|
+
@property
|
|
128
|
+
def active_scale(self):
|
|
129
|
+
return 1 / hax.axis_size(self.In)
|
|
130
|
+
|
|
131
|
+
@property
|
|
132
|
+
def lr_scale(self):
|
|
133
|
+
return 1
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
@dataclass
|
|
137
|
+
class AbstractEmbeddingReparam(AbstractReparam):
|
|
138
|
+
"""Base class for embedding-layer reparameterizations.
|
|
139
|
+
|
|
140
|
+
Defines the interface for both embedding and unembedding
|
|
141
|
+
scaling rules.
|
|
142
|
+
"""
|
|
143
|
+
|
|
144
|
+
Embed: AxisSpec
|
|
145
|
+
Vocab: AxisSpec
|
|
146
|
+
|
|
147
|
+
@property
|
|
148
|
+
@abstractmethod
|
|
149
|
+
def unembed_active_scale(self):
|
|
150
|
+
"""Scaling factor applied when unembedding embeddings."""
|
|
151
|
+
raise NotImplementedError
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
class EmbeddingStandardParam(AbstractEmbeddingReparam):
|
|
155
|
+
"""Standard embedding parameterization."""
|
|
156
|
+
|
|
157
|
+
@staticmethod
|
|
158
|
+
def init_scale(In: AxisSpec, Out: AxisSpec):
|
|
159
|
+
return 1 / hax.axis_size(Out)
|
|
160
|
+
|
|
161
|
+
@property
|
|
162
|
+
def active_scale(self):
|
|
163
|
+
return 1
|
|
164
|
+
|
|
165
|
+
@property
|
|
166
|
+
def lr_scale(self):
|
|
167
|
+
return 1
|
|
168
|
+
|
|
169
|
+
@property
|
|
170
|
+
def unembed_active_scale(self):
|
|
171
|
+
return 1 / hax.axis_size(self.Embed)
|
|
172
|
+
|
|
173
|
+
|
|
174
|
+
class EmbeddingMup(AbstractEmbeddingReparam):
|
|
175
|
+
"""muP-style parameterization for embeddings.
|
|
176
|
+
|
|
177
|
+
Keeps initialization and learning-rate scaling neutral,
|
|
178
|
+
but applies inverse-width scaling to unembedding outputs for tied weights.
|
|
179
|
+
See: https://www.cerebras.ai/blog/the-practitioners-guide-to-the-maximal-update-parameterization
|
|
180
|
+
"""
|
|
181
|
+
|
|
182
|
+
@staticmethod
|
|
183
|
+
def init_scale(In: AxisSpec, Out: AxisSpec):
|
|
184
|
+
return 1
|
|
185
|
+
|
|
186
|
+
@property
|
|
187
|
+
def active_scale(self):
|
|
188
|
+
return 1
|
|
189
|
+
|
|
190
|
+
@property
|
|
191
|
+
def lr_scale(self):
|
|
192
|
+
return 1
|
|
193
|
+
|
|
194
|
+
@property
|
|
195
|
+
def unembed_active_scale(self):
|
|
196
|
+
return 1 / hax.axis_size(self.Embed)
|
|
197
|
+
|
|
198
|
+
|
|
199
|
+
class ReparamEnabled(ABC):
|
|
200
|
+
"""Mixin for modules that support reparameterization.
|
|
201
|
+
|
|
202
|
+
Stores an abstract `reparam` attribute that specifies
|
|
203
|
+
how initialization and scaling are handled.
|
|
204
|
+
"""
|
|
205
|
+
|
|
206
|
+
reparam: eqx.AbstractVar[AbstractReparam]
|
|
@@ -0,0 +1,164 @@
|
|
|
1
|
+
# Copyright 2025 The Levanter Authors
|
|
2
|
+
#
|
|
3
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
4
|
+
|
|
5
|
+
"""Coordinate check for µP modules built on Haliax primitives."""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import dataclasses
|
|
10
|
+
from typing import Any, Iterable
|
|
11
|
+
|
|
12
|
+
import equinox as eqx
|
|
13
|
+
import jax
|
|
14
|
+
import jax.random as jrandom
|
|
15
|
+
|
|
16
|
+
import haliax as hax
|
|
17
|
+
from haliax import Axis, NamedArray
|
|
18
|
+
from haliax.nn import Linear, activations
|
|
19
|
+
from haliax.nn.mup import InputLinearMup, HiddenLinearMup, OutputLinearMup
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class TinyMLP(eqx.Module):
|
|
23
|
+
"""Minimal 2-hidden-layer MLP composed of Haliax Linear variants."""
|
|
24
|
+
|
|
25
|
+
first: Linear
|
|
26
|
+
second: Linear
|
|
27
|
+
third: Linear
|
|
28
|
+
|
|
29
|
+
@staticmethod
|
|
30
|
+
def init(width: int, *, key: jax.Array, use_mup: bool) -> TinyMLP:
|
|
31
|
+
in_axis = Axis("in", 2)
|
|
32
|
+
hidden = Axis("hidden", width)
|
|
33
|
+
hidden2 = hidden.alias("hidden2")
|
|
34
|
+
out_axis = Axis("out", 1)
|
|
35
|
+
|
|
36
|
+
k1, k2, k3 = jrandom.split(key, 3)
|
|
37
|
+
|
|
38
|
+
if use_mup:
|
|
39
|
+
first = Linear.init((in_axis,), hidden, key=k1, reparam_cls=InputLinearMup)
|
|
40
|
+
second = Linear.init(hidden, hidden2, key=k2, reparam_cls=HiddenLinearMup)
|
|
41
|
+
third = Linear.init(hidden2, (out_axis,), key=k3, reparam_cls=OutputLinearMup)
|
|
42
|
+
else:
|
|
43
|
+
first = Linear.init((in_axis,), hidden, key=k1)
|
|
44
|
+
second = Linear.init(hidden, hidden2, key=k2)
|
|
45
|
+
third = Linear.init(hidden2, (out_axis,), key=k3)
|
|
46
|
+
|
|
47
|
+
return TinyMLP(first=first, second=second, third=third)
|
|
48
|
+
|
|
49
|
+
def __call__(self, x: NamedArray) -> NamedArray:
|
|
50
|
+
h = activations.relu(self.first(x))
|
|
51
|
+
h = activations.relu(self.second(h))
|
|
52
|
+
return self.third(h)
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def _loss_fn(params: TinyMLP, x: NamedArray, y: NamedArray) -> jax.Array:
|
|
56
|
+
preds = params(x)
|
|
57
|
+
diff = preds - y
|
|
58
|
+
return hax.mean(diff * diff).scalar()
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
_loss_and_grad = eqx.filter_jit(eqx.filter_value_and_grad(_loss_fn))
|
|
62
|
+
_loss_value = jax.jit(_loss_fn)
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
def _apply_sgd(module: TinyMLP, grads: TinyMLP, *, base_lr: float, use_mup: bool) -> TinyMLP:
|
|
66
|
+
def update_linear(layer: Linear, grad_layer: Linear) -> Linear:
|
|
67
|
+
lr_scale = layer.reparam.lr_scale
|
|
68
|
+
new_weight = layer.weight - (base_lr * lr_scale) * grad_layer.weight
|
|
69
|
+
if layer.bias is None or grad_layer.bias is None:
|
|
70
|
+
new_bias = layer.bias
|
|
71
|
+
else:
|
|
72
|
+
new_bias = layer.bias - base_lr * grad_layer.bias
|
|
73
|
+
return dataclasses.replace(layer, weight=new_weight, bias=new_bias)
|
|
74
|
+
|
|
75
|
+
return TinyMLP(
|
|
76
|
+
first=update_linear(module.first, grads.first),
|
|
77
|
+
second=update_linear(module.second, grads.second),
|
|
78
|
+
third=update_linear(module.third, grads.third),
|
|
79
|
+
)
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
def _make_dataset(key: jax.Array, *, n_points: int = 2048) -> tuple[NamedArray, NamedArray]:
|
|
83
|
+
data_axis = Axis("data", n_points)
|
|
84
|
+
feature_axis = Axis("in", 2)
|
|
85
|
+
out_axis = Axis("out", 1)
|
|
86
|
+
|
|
87
|
+
xy = jrandom.uniform(key, (n_points, 2), minval=-1.0, maxval=1.0)
|
|
88
|
+
inputs = hax.named(xy, (data_axis, feature_axis))
|
|
89
|
+
targets = hax.named(xy[:, :1], (data_axis, out_axis))
|
|
90
|
+
return inputs, targets
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
def _run_once(
|
|
94
|
+
key: jax.Array,
|
|
95
|
+
*,
|
|
96
|
+
width: int,
|
|
97
|
+
use_mup: bool,
|
|
98
|
+
steps: int = 120,
|
|
99
|
+
batch_size: int = 256,
|
|
100
|
+
base_lr: float = 3e-3,
|
|
101
|
+
) -> float:
|
|
102
|
+
data_key, model_key = jrandom.split(key)
|
|
103
|
+
inputs, targets = _make_dataset(data_key)
|
|
104
|
+
params = TinyMLP.init(width, key=model_key, use_mup=use_mup)
|
|
105
|
+
|
|
106
|
+
def train_step(state: TinyMLP, xb: NamedArray, yb: NamedArray):
|
|
107
|
+
loss, grads = _loss_and_grad(state, xb, yb)
|
|
108
|
+
new_state = _apply_sgd(state, grads, base_lr=base_lr, use_mup=use_mup)
|
|
109
|
+
return new_state, loss
|
|
110
|
+
|
|
111
|
+
data_axis = inputs.axes[0]
|
|
112
|
+
n = data_axis.size
|
|
113
|
+
|
|
114
|
+
state = params
|
|
115
|
+
for t in range(steps):
|
|
116
|
+
start = (t * batch_size) % n
|
|
117
|
+
end = start + batch_size
|
|
118
|
+
batch_idx = (data_axis, slice(start, end))
|
|
119
|
+
xb = inputs[batch_idx]
|
|
120
|
+
yb = targets[batch_idx]
|
|
121
|
+
state, _ = train_step(state, xb, yb)
|
|
122
|
+
|
|
123
|
+
final_loss = _loss_value(state, inputs, targets)
|
|
124
|
+
return float(final_loss)
|
|
125
|
+
|
|
126
|
+
|
|
127
|
+
def _span(values: Iterable[float]) -> float:
|
|
128
|
+
seq = list(values)
|
|
129
|
+
return max(seq) - min(seq)
|
|
130
|
+
|
|
131
|
+
|
|
132
|
+
def coord_check(
|
|
133
|
+
widths: tuple[int, ...] = (32, 128, 512),
|
|
134
|
+
*,
|
|
135
|
+
steps: int = 120,
|
|
136
|
+
base_lr: float = 3e-3,
|
|
137
|
+
) -> dict[str, Any]:
|
|
138
|
+
seed = 0
|
|
139
|
+
keys = jrandom.split(jrandom.PRNGKey(seed), len(widths))
|
|
140
|
+
mup_losses = [
|
|
141
|
+
_run_once(key, width=width, use_mup=True, steps=steps, base_lr=base_lr) for key, width in zip(keys, widths)
|
|
142
|
+
]
|
|
143
|
+
ctrl_losses = [
|
|
144
|
+
_run_once(key, width=width, use_mup=False, steps=steps, base_lr=base_lr) for key, width in zip(keys, widths)
|
|
145
|
+
]
|
|
146
|
+
|
|
147
|
+
return {
|
|
148
|
+
"widths": list(widths),
|
|
149
|
+
"mup_losses": mup_losses,
|
|
150
|
+
"ctrl_losses": ctrl_losses,
|
|
151
|
+
"mup_span": _span(mup_losses),
|
|
152
|
+
"ctrl_span": _span(ctrl_losses),
|
|
153
|
+
}
|
|
154
|
+
|
|
155
|
+
|
|
156
|
+
def test_mup_coordinate_check_is_width_invariant():
|
|
157
|
+
result = coord_check(widths=(32, 128, 512), steps=120, base_lr=3e-3)
|
|
158
|
+
mup_span = result["mup_span"]
|
|
159
|
+
ctrl_span = result["ctrl_span"]
|
|
160
|
+
|
|
161
|
+
if ctrl_span < 1e-5:
|
|
162
|
+
assert mup_span <= ctrl_span + 1e-6, f"μP not at least as invariant: {result}"
|
|
163
|
+
else:
|
|
164
|
+
assert mup_span <= 0.6 * ctrl_span, f"μP did not improve width invariance enough.\n{result}"
|
|
@@ -0,0 +1,48 @@
|
|
|
1
|
+
# Copyright 2025 The Levanter Authors
|
|
2
|
+
#
|
|
3
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
4
|
+
|
|
5
|
+
import jax.numpy as jnp
|
|
6
|
+
import jax.random as jrandom
|
|
7
|
+
import pytest
|
|
8
|
+
|
|
9
|
+
import haliax as hax
|
|
10
|
+
from haliax.nn import Embedding
|
|
11
|
+
from haliax.nn.mup import EmbeddingMup
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
@pytest.mark.parametrize("init_scale", [0.5, 1.0])
|
|
15
|
+
def test_mup_embedding_init_matches_embedding(init_scale: float):
|
|
16
|
+
Vocab = hax.Axis("V", 8)
|
|
17
|
+
Embed = (hax.Axis("E", 4),)
|
|
18
|
+
|
|
19
|
+
key = jrandom.PRNGKey(0)
|
|
20
|
+
|
|
21
|
+
baseline = Embedding.init(Vocab, Embed, key=key, init_scale=init_scale)
|
|
22
|
+
mup = Embedding.init(Vocab, Embed, key=key, init_scale=init_scale, reparam_cls=EmbeddingMup)
|
|
23
|
+
|
|
24
|
+
assert mup.weight.axes == baseline.weight.axes
|
|
25
|
+
scale_factor = hax.axis_size(Embed)
|
|
26
|
+
assert jnp.allclose(mup.weight.array, baseline.weight.array * scale_factor)
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def test_mup_embedding_unembedding_scale():
|
|
30
|
+
Vocab = hax.Axis("V", 6)
|
|
31
|
+
Embed = (hax.Axis("E", 3),)
|
|
32
|
+
|
|
33
|
+
weight = hax.ones(hax.concat_axis_specs(Vocab, Embed))
|
|
34
|
+
layer = Embedding(weight=weight, Vocab=Vocab, Embed=Embed, reparam=EmbeddingMup(Embed, Vocab))
|
|
35
|
+
|
|
36
|
+
scale = layer.reparam.unembed_active_scale
|
|
37
|
+
assert scale == pytest.approx(1.0 / hax.axis_size(Embed))
|
|
38
|
+
assert layer.reparam.active_scale == pytest.approx(1.0)
|
|
39
|
+
assert jnp.allclose((layer.weight * scale).array, weight.array * scale)
|
|
40
|
+
|
|
41
|
+
Batch = hax.Axis("B", 2)
|
|
42
|
+
inputs = hax.ones((Batch, *Embed))
|
|
43
|
+
|
|
44
|
+
logits = layer.unembed(inputs)
|
|
45
|
+
expected = hax.ones((Batch, Vocab))
|
|
46
|
+
|
|
47
|
+
assert logits.axes == expected.axes
|
|
48
|
+
assert jnp.allclose(logits.array, expected.array)
|
|
@@ -0,0 +1,120 @@
|
|
|
1
|
+
# Copyright 2025 The Levanter Authors
|
|
2
|
+
#
|
|
3
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
import dataclasses
|
|
7
|
+
|
|
8
|
+
import jax.numpy as jnp
|
|
9
|
+
import jax.random as jrandom
|
|
10
|
+
import pytest
|
|
11
|
+
|
|
12
|
+
import haliax as hax
|
|
13
|
+
from haliax.nn import Linear
|
|
14
|
+
from haliax.nn.mup import InputLinearMup, LinearStandardParam, HiddenLinearMup, OutputLinearMup
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
@pytest.mark.parametrize("out_first", [True, False])
|
|
18
|
+
def test_mup_linear_init_matches_linear_axes(out_first: bool):
|
|
19
|
+
In = (hax.Axis("I", 3), hax.Axis("J", 2))
|
|
20
|
+
Out = (hax.Axis("O", 5),)
|
|
21
|
+
key_linear, key_mup = jrandom.split(jrandom.PRNGKey(0))
|
|
22
|
+
|
|
23
|
+
linear = Linear.init(In, Out, key=key_linear, out_first=out_first)
|
|
24
|
+
mup = Linear.init(In, Out, key=key_mup, out_first=out_first, reparam_cls=InputLinearMup)
|
|
25
|
+
|
|
26
|
+
assert linear.weight.axes == mup.weight.axes
|
|
27
|
+
if linear.bias is not None:
|
|
28
|
+
assert mup.bias is not None
|
|
29
|
+
assert linear.bias.axes == mup.bias.axes
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def test_mup_linear_call_matches_linear():
|
|
33
|
+
Batch = hax.Axis("B", 2)
|
|
34
|
+
In = (hax.Axis("I", 3),)
|
|
35
|
+
Out = (hax.Axis("O", 4),)
|
|
36
|
+
|
|
37
|
+
weight = hax.ones(hax.concat_axis_specs(Out, In)) * 0.5
|
|
38
|
+
bias = hax.full(Out, 0.25)
|
|
39
|
+
|
|
40
|
+
linear = Linear(weight, bias, In, Out, reparam=LinearStandardParam(In, Out))
|
|
41
|
+
mup = Linear(weight, bias, In, Out, reparam=InputLinearMup(In, Out))
|
|
42
|
+
|
|
43
|
+
inputs = hax.full(hax.concat_axis_specs(Batch, In), 2.0)
|
|
44
|
+
|
|
45
|
+
expected = linear(inputs)
|
|
46
|
+
actual = mup(inputs)
|
|
47
|
+
|
|
48
|
+
assert actual.axes == expected.axes
|
|
49
|
+
assert jnp.allclose(actual.array, expected.array)
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
@pytest.mark.parametrize("out_first", [True, False])
|
|
53
|
+
def test_hidden_linear_init_matches_linear_scaling(out_first: bool):
|
|
54
|
+
In = hax.Axis("I", 6)
|
|
55
|
+
Out = hax.Axis("O", 5)
|
|
56
|
+
key = jrandom.PRNGKey(0)
|
|
57
|
+
|
|
58
|
+
linear = Linear.init(In, Out, key=key, use_bias=False, out_first=out_first)
|
|
59
|
+
hidden = Linear.init(
|
|
60
|
+
In,
|
|
61
|
+
Out,
|
|
62
|
+
key=key,
|
|
63
|
+
use_bias=False,
|
|
64
|
+
out_first=out_first,
|
|
65
|
+
reparam_cls=HiddenLinearMup,
|
|
66
|
+
)
|
|
67
|
+
|
|
68
|
+
assert jnp.allclose(hidden.weight.array, linear.weight.array)
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def test_output_linear_scales_activation_by_input_size():
|
|
72
|
+
Batch = hax.Axis("B", 2)
|
|
73
|
+
In = hax.Axis("I", 3)
|
|
74
|
+
Out = hax.Axis("O", 2)
|
|
75
|
+
|
|
76
|
+
layer = Linear.init(In, Out, key=jrandom.PRNGKey(2), use_bias=False, reparam_cls=OutputLinearMup)
|
|
77
|
+
layer = dataclasses.replace(layer, weight=hax.ones(layer.weight.axes))
|
|
78
|
+
|
|
79
|
+
inputs = hax.full((Batch, In), 2.0)
|
|
80
|
+
expected = hax.full(hax.concat_axis_specs(Batch, layer.Out), 2.0)
|
|
81
|
+
|
|
82
|
+
actual = layer(inputs)
|
|
83
|
+
|
|
84
|
+
assert actual.axes == expected.axes
|
|
85
|
+
assert jnp.allclose(actual.array, expected.array)
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
def test_output_linear_state_dict_scales_weight():
|
|
89
|
+
In = hax.Axis("I", 4)
|
|
90
|
+
Out = hax.Axis("O", 3)
|
|
91
|
+
|
|
92
|
+
layer = Linear.init(In, Out, key=jrandom.PRNGKey(3), use_bias=False, reparam_cls=OutputLinearMup)
|
|
93
|
+
layer = dataclasses.replace(layer, weight=hax.ones(layer.weight.axes))
|
|
94
|
+
|
|
95
|
+
state = layer.to_state_dict()
|
|
96
|
+
assert jnp.allclose(state["weight"], layer.weight.array * layer.reparam.active_scale)
|
|
97
|
+
|
|
98
|
+
template = Linear.init(In, Out, key=jrandom.PRNGKey(4), use_bias=False, reparam_cls=OutputLinearMup)
|
|
99
|
+
restored = template.from_state_dict(state)
|
|
100
|
+
|
|
101
|
+
assert jnp.allclose(restored.weight.array, layer.weight.array)
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
def test_input_linear_behaves_like_base_linear():
|
|
105
|
+
Batch = hax.Axis("B", 2)
|
|
106
|
+
In = hax.Axis("I", 3)
|
|
107
|
+
Out = hax.Axis("O", 4)
|
|
108
|
+
|
|
109
|
+
weight = hax.ones((Out, In)) * 0.1
|
|
110
|
+
bias = hax.zeros(Out)
|
|
111
|
+
|
|
112
|
+
linear = Linear(weight, bias, In, Out, reparam=LinearStandardParam(In, Out))
|
|
113
|
+
input_linear = Linear(weight, bias, In, Out, reparam=InputLinearMup(In, Out))
|
|
114
|
+
|
|
115
|
+
inputs = hax.random.normal(jrandom.PRNGKey(5), (Batch, In))
|
|
116
|
+
|
|
117
|
+
expected = linear(inputs)
|
|
118
|
+
actual = input_linear(inputs)
|
|
119
|
+
|
|
120
|
+
assert jnp.allclose(actual.array, expected.array)
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|