haliax 1.4.dev420__tar.gz → 1.4.dev439__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.dev439}/PKG-INFO +1 -1
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/api.md +31 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/__about__.py +1 -1
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/__init__.py +1 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/embedding.py +22 -8
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/linear.py +51 -7
- haliax-1.4.dev439/src/haliax/nn/mup.py +206 -0
- haliax-1.4.dev439/src/haliax/tree.py +59 -0
- haliax-1.4.dev439/tests/test_mup_coordinate_check.py +164 -0
- haliax-1.4.dev439/tests/test_mup_embedding.py +48 -0
- haliax-1.4.dev439/tests/test_mup_linear.py +120 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/.agents/projects/api_parity.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/.coveragerc +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/.flake8 +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/.github/workflows/publish_dev.yaml +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/.github/workflows/run_pre_commit.yaml +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/.github/workflows/run_tests.yaml +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/.gitignore +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/.playbooks/add-types.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/.playbooks/wrap-non-named.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/.pre-commit-config.yaml +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/.readthedocs.yaml +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/AGENTS.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/AUTHORS.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/CONTRIBUTING.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/CONTRIBUTORS.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/LICENSE +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/README.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/broadcasting.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/cheatsheet.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/css/material.css +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/css/mkdocstrings.css +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/faq.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/data_parallel_mesh.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/data_parallel_mesh_replicated.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/device_mesh_1d.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/device_mesh_1d_zero.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/device_mesh_2d.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/device_mesh_2d_zero.png +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/fp8.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/index.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/indexing.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/matmul.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/nn.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/partitioning.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/primer.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/rearrange.ipynb +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/rearrange.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/requirements.txt +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/scan.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/state-dict.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/tutorial.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/typing.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/vmap.md +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/etc/license_header.txt +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/mkdocs.yml +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/pyproject.toml +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/__init__.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/compile_utils.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/dot.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/einsum.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/fp8.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/parsing.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/rearrange.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/scan.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/state_dict.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/util.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/axis.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/core.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/debug.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/fft.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/field.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/haxtyping.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/hof.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/jax_utils.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/__init__.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/activations.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/attention.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/conv.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/dropout.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/loss.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/mlp.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/normalization.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/pool.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/scan.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/ops.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/partitioning.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/poly.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/quantization.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/random.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/specialized_fns.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/state_dict.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/tree_util.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/types.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/util.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/wrap.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/core_test.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_attention.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_axis.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_bitwise_ops.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_conv.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_debug.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_dot.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_dtype_typing.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_einsum.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_fft.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_field.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_fp8.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_hof.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_int8.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_moe_linear.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_namedarray_typing.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_nan_reductions.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_nn.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_ops.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_parsing.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_partitioning.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_poly_ops.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_pool.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_random.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_rearrange.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_scan.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_scatter_gather.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_specialized_fns.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_state_dict.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_tree_util.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_utils.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_visualize_sharding.py +0 -0
- {haliax-1.4.dev420 → haliax-1.4.dev439}/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.dev439
|
|
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/
|
|
@@ -4,6 +4,37 @@ that we use names (either strings or [haliax.Axis][] objects) to specify axes in
|
|
|
4
4
|
arrays (see [haliax.zeros][] and [haliax.ones][]) as well as things like reductions (see [haliax.sum][] and
|
|
5
5
|
[haliax.mean][]).
|
|
6
6
|
|
|
7
|
+
## PyTree Helpers
|
|
8
|
+
|
|
9
|
+
PyTrees are the lingua franca for composing state in JAX ecosystems. Haliax provides drop-in replacements for the
|
|
10
|
+
[`jax.tree`][] helpers that are aware of [`NamedArray`][haliax.NamedArray] semantics. They preserve axis metadata across
|
|
11
|
+
transformations while interoperating with standard JAX containers, so you can use them anywhere you would have reached
|
|
12
|
+
for JAX's versions.
|
|
13
|
+
|
|
14
|
+
Use these helpers whenever you need to map, flatten, or rebuild PyTrees that might include `NamedArray` instances:
|
|
15
|
+
|
|
16
|
+
* [`haliax.tree.map`][] mirrors [`jax.tree.map`][] but forwards to Haliax's [`haliax.tree_util.tree_map`][] so axis names remain
|
|
17
|
+
intact.
|
|
18
|
+
* [`haliax.tree.scan_aware_map`][] descends into [`haliax.nn.Stacked`][haliax.nn.Stacked] modules so that each layer is
|
|
19
|
+
transformed individually, effectively treating them as if they were unrolled when applying
|
|
20
|
+
[`haliax.tree_util.scan_aware_tree_map`][].
|
|
21
|
+
* [`haliax.tree.flatten`][] / [`haliax.tree.unflatten`][] match the familiar flattening API while handling `NamedArray`
|
|
22
|
+
payloads safely.
|
|
23
|
+
* [`haliax.tree.leaves`][] and [`haliax.tree.structure`][] provide direct access to the leaves and PyTree structure.
|
|
24
|
+
|
|
25
|
+
All of these helpers accept the same `is_leaf` hook you might already use with JAX's utilities. They should be the first
|
|
26
|
+
tools you reach for when you need deterministic tree transforms that understand named axes.
|
|
27
|
+
|
|
28
|
+
::: haliax.tree.map
|
|
29
|
+
::: haliax.tree.scan_aware_map
|
|
30
|
+
::: haliax.tree.flatten
|
|
31
|
+
::: haliax.tree.unflatten
|
|
32
|
+
::: haliax.tree.leaves
|
|
33
|
+
::: haliax.tree.structure
|
|
34
|
+
|
|
35
|
+
[`jax.tree`]: https://jax.readthedocs.io/en/latest/_autosummary/jax.tree.html
|
|
36
|
+
[`jax.tree.map`]: https://jax.readthedocs.io/en/latest/_autosummary/jax.tree.map.html
|
|
37
|
+
|
|
7
38
|
## Axis Types
|
|
8
39
|
|
|
9
40
|
If you already speak NumPy or `jax.numpy`, think of Haliax as swapping positional axes (`axis=0`) for named axes
|
|
@@ -16,6 +16,7 @@ import haliax.nn as nn
|
|
|
16
16
|
import haliax.quantization as quantization
|
|
17
17
|
import haliax.random as random
|
|
18
18
|
import haliax.state_dict as state_dict
|
|
19
|
+
import haliax.tree as tree # noqa: F401
|
|
19
20
|
import haliax.tree_util as tree_util
|
|
20
21
|
import haliax.util as util
|
|
21
22
|
from .field import field
|
|
@@ -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,59 @@
|
|
|
1
|
+
# Copyright 2025 The Levanter Authors
|
|
2
|
+
#
|
|
3
|
+
# SPDX-License-Identifier: Apache-2.0
|
|
4
|
+
|
|
5
|
+
"""Convenience wrappers for :mod:`haliax.tree_util` that mirror :mod:`jax.tree`."""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
from typing import Any, Callable, Iterable, Sequence, TypeVar
|
|
10
|
+
|
|
11
|
+
from . import tree_util
|
|
12
|
+
|
|
13
|
+
T = TypeVar("T")
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def map(fn: Callable[..., T], tree: Any, *rest: Any, is_leaf: Callable[[Any], bool] | None = None) -> Any:
|
|
17
|
+
"""Alias for :func:`haliax.tree_util.tree_map` matching :func:`jax.tree.map`."""
|
|
18
|
+
|
|
19
|
+
return tree_util.tree_map(fn, tree, *rest, is_leaf=is_leaf)
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def scan_aware_map(fn: Callable[..., T], tree: Any, *rest: Any, is_leaf: Callable[[Any], bool] | None = None) -> Any:
|
|
23
|
+
"""Alias for :func:`haliax.tree_util.scan_aware_tree_map` with :mod:`jax.tree` style naming."""
|
|
24
|
+
|
|
25
|
+
return tree_util.scan_aware_tree_map(fn, tree, *rest, is_leaf=is_leaf)
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def flatten(tree: Any, *, is_leaf: Callable[[Any], bool] | None = None) -> tuple[Sequence[Any], Any]:
|
|
29
|
+
"""Alias for :func:`haliax.tree_util.tree_flatten` matching :func:`jax.tree.flatten`."""
|
|
30
|
+
|
|
31
|
+
return tree_util.tree_flatten(tree, is_leaf=is_leaf)
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def unflatten(treedef: Any, leaves: Iterable[Any]) -> Any:
|
|
35
|
+
"""Alias for :func:`haliax.tree_util.tree_unflatten` matching :func:`jax.tree.unflatten`."""
|
|
36
|
+
|
|
37
|
+
return tree_util.tree_unflatten(treedef, leaves)
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def leaves(tree: Any, *, is_leaf: Callable[[Any], bool] | None = None) -> Sequence[Any]:
|
|
41
|
+
"""Alias for :func:`haliax.tree_util.tree_leaves` matching :func:`jax.tree.leaves`."""
|
|
42
|
+
|
|
43
|
+
return tree_util.tree_leaves(tree, is_leaf=is_leaf)
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def structure(tree: Any, *, is_leaf: Callable[[Any], bool] | None = None) -> Any:
|
|
47
|
+
"""Alias for :func:`haliax.tree_util.tree_structure` matching :func:`jax.tree.structure`."""
|
|
48
|
+
|
|
49
|
+
return tree_util.tree_structure(tree, is_leaf=is_leaf)
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
__all__ = [
|
|
53
|
+
"map",
|
|
54
|
+
"scan_aware_map",
|
|
55
|
+
"flatten",
|
|
56
|
+
"unflatten",
|
|
57
|
+
"leaves",
|
|
58
|
+
"structure",
|
|
59
|
+
]
|
|
@@ -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}"
|