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.
Files changed (134) hide show
  1. {haliax-1.4.dev420 → haliax-1.4.dev439}/PKG-INFO +1 -1
  2. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/api.md +31 -0
  3. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/__about__.py +1 -1
  4. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/__init__.py +1 -0
  5. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/embedding.py +22 -8
  6. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/linear.py +51 -7
  7. haliax-1.4.dev439/src/haliax/nn/mup.py +206 -0
  8. haliax-1.4.dev439/src/haliax/tree.py +59 -0
  9. haliax-1.4.dev439/tests/test_mup_coordinate_check.py +164 -0
  10. haliax-1.4.dev439/tests/test_mup_embedding.py +48 -0
  11. haliax-1.4.dev439/tests/test_mup_linear.py +120 -0
  12. {haliax-1.4.dev420 → haliax-1.4.dev439}/.agents/projects/api_parity.md +0 -0
  13. {haliax-1.4.dev420 → haliax-1.4.dev439}/.coveragerc +0 -0
  14. {haliax-1.4.dev420 → haliax-1.4.dev439}/.flake8 +0 -0
  15. {haliax-1.4.dev420 → haliax-1.4.dev439}/.github/workflows/publish_dev.yaml +0 -0
  16. {haliax-1.4.dev420 → haliax-1.4.dev439}/.github/workflows/run_pre_commit.yaml +0 -0
  17. {haliax-1.4.dev420 → haliax-1.4.dev439}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  18. {haliax-1.4.dev420 → haliax-1.4.dev439}/.github/workflows/run_tests.yaml +0 -0
  19. {haliax-1.4.dev420 → haliax-1.4.dev439}/.gitignore +0 -0
  20. {haliax-1.4.dev420 → haliax-1.4.dev439}/.playbooks/add-types.md +0 -0
  21. {haliax-1.4.dev420 → haliax-1.4.dev439}/.playbooks/wrap-non-named.md +0 -0
  22. {haliax-1.4.dev420 → haliax-1.4.dev439}/.pre-commit-config.yaml +0 -0
  23. {haliax-1.4.dev420 → haliax-1.4.dev439}/.readthedocs.yaml +0 -0
  24. {haliax-1.4.dev420 → haliax-1.4.dev439}/AGENTS.md +0 -0
  25. {haliax-1.4.dev420 → haliax-1.4.dev439}/AUTHORS.md +0 -0
  26. {haliax-1.4.dev420 → haliax-1.4.dev439}/CONTRIBUTING.md +0 -0
  27. {haliax-1.4.dev420 → haliax-1.4.dev439}/CONTRIBUTORS.md +0 -0
  28. {haliax-1.4.dev420 → haliax-1.4.dev439}/LICENSE +0 -0
  29. {haliax-1.4.dev420 → haliax-1.4.dev439}/README.md +0 -0
  30. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/broadcasting.md +0 -0
  31. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/cheatsheet.md +0 -0
  32. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/css/material.css +0 -0
  33. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/css/mkdocstrings.css +0 -0
  34. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/faq.md +0 -0
  35. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/data_parallel_mesh.png +0 -0
  36. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  37. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/device_mesh_1d.png +0 -0
  38. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/device_mesh_1d_zero.png +0 -0
  39. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/device_mesh_2d.png +0 -0
  40. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  41. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  42. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  43. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  44. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/figures/device_mesh_2d_zero.png +0 -0
  45. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/fp8.md +0 -0
  46. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/index.md +0 -0
  47. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/indexing.md +0 -0
  48. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/matmul.md +0 -0
  49. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/nn.md +0 -0
  50. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/partitioning.md +0 -0
  51. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/primer.md +0 -0
  52. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/rearrange.ipynb +0 -0
  53. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/rearrange.md +0 -0
  54. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/requirements.txt +0 -0
  55. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/scan.md +0 -0
  56. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/state-dict.md +0 -0
  57. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/tutorial.md +0 -0
  58. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/typing.md +0 -0
  59. {haliax-1.4.dev420 → haliax-1.4.dev439}/docs/vmap.md +0 -0
  60. {haliax-1.4.dev420 → haliax-1.4.dev439}/etc/license_header.txt +0 -0
  61. {haliax-1.4.dev420 → haliax-1.4.dev439}/mkdocs.yml +0 -0
  62. {haliax-1.4.dev420 → haliax-1.4.dev439}/pyproject.toml +0 -0
  63. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/__init__.py +0 -0
  64. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/compile_utils.py +0 -0
  65. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/dot.py +0 -0
  66. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/einsum.py +0 -0
  67. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/fp8.py +0 -0
  68. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/parsing.py +0 -0
  69. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/rearrange.py +0 -0
  70. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/scan.py +0 -0
  71. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/state_dict.py +0 -0
  72. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/_src/util.py +0 -0
  73. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/axis.py +0 -0
  74. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/core.py +0 -0
  75. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/debug.py +0 -0
  76. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/fft.py +0 -0
  77. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/field.py +0 -0
  78. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/haxtyping.py +0 -0
  79. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/hof.py +0 -0
  80. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/jax_utils.py +0 -0
  81. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/__init__.py +0 -0
  82. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/activations.py +0 -0
  83. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/attention.py +0 -0
  84. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/conv.py +0 -0
  85. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/dropout.py +0 -0
  86. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/loss.py +0 -0
  87. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/mlp.py +0 -0
  88. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/normalization.py +0 -0
  89. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/pool.py +0 -0
  90. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/nn/scan.py +0 -0
  91. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/ops.py +0 -0
  92. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/partitioning.py +0 -0
  93. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/poly.py +0 -0
  94. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/quantization.py +0 -0
  95. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/random.py +0 -0
  96. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/specialized_fns.py +0 -0
  97. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/state_dict.py +0 -0
  98. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/tree_util.py +0 -0
  99. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/types.py +0 -0
  100. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/util.py +0 -0
  101. {haliax-1.4.dev420 → haliax-1.4.dev439}/src/haliax/wrap.py +0 -0
  102. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/core_test.py +0 -0
  103. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_attention.py +0 -0
  104. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_axis.py +0 -0
  105. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_bitwise_ops.py +0 -0
  106. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_conv.py +0 -0
  107. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_debug.py +0 -0
  108. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_dot.py +0 -0
  109. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_dtype_typing.py +0 -0
  110. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_einsum.py +0 -0
  111. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_fft.py +0 -0
  112. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_field.py +0 -0
  113. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_fp8.py +0 -0
  114. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_hof.py +0 -0
  115. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_int8.py +0 -0
  116. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_moe_linear.py +0 -0
  117. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_namedarray_typing.py +0 -0
  118. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_nan_reductions.py +0 -0
  119. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_nn.py +0 -0
  120. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_ops.py +0 -0
  121. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_parsing.py +0 -0
  122. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_partitioning.py +0 -0
  123. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_poly_ops.py +0 -0
  124. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_pool.py +0 -0
  125. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_random.py +0 -0
  126. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_rearrange.py +0 -0
  127. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_scan.py +0 -0
  128. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_scatter_gather.py +0 -0
  129. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_specialized_fns.py +0 -0
  130. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_state_dict.py +0 -0
  131. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_tree_util.py +0 -0
  132. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_utils.py +0 -0
  133. {haliax-1.4.dev420 → haliax-1.4.dev439}/tests/test_visualize_sharding.py +0 -0
  134. {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.dev420
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
@@ -3,4 +3,4 @@
3
3
  # SPDX-License-Identifier: Apache-2.0
4
4
 
5
5
 
6
- __version__ = "1.4.dev420"
6
+ __version__ = "1.4.dev439"
@@ -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(Vocab: Axis, Embed: AxisSpec, *, init_scale: float = 1, key, initializer_range: float | None = None):
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("initializer_range is deprecated. Use init_std instead.", DeprecationWarning)
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
- output_size = hax.axis_size(Embed)
49
- weight = hax.random.truncated_normal(key, all_axes, -3, 3) * (init_scale / output_size)
50
- return Embedding(weight=weight, Vocab=Vocab, Embed=Embed)
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
- from .._src.state_dict import Mod, ModuleWithStateDictSerialization
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
- input_size = hax.axis_size(In)
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(self.weight, axis=self.In, dot_general=self.dot_general)
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
- (ensure_tuple(inputs.axis_indices(self.In)), ensure_tuple(self.weight.axis_indices(self.In))),
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}"