haliax 1.4.dev441__tar.gz → 1.4.dev443__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (134) hide show
  1. {haliax-1.4.dev441 → haliax-1.4.dev443}/PKG-INFO +1 -1
  2. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/__about__.py +1 -1
  3. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/embedding.py +7 -2
  4. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/linear.py +7 -2
  5. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_mup_embedding.py +1 -1
  6. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_mup_linear.py +10 -5
  7. {haliax-1.4.dev441 → haliax-1.4.dev443}/.agents/projects/api_parity.md +0 -0
  8. {haliax-1.4.dev441 → haliax-1.4.dev443}/.coveragerc +0 -0
  9. {haliax-1.4.dev441 → haliax-1.4.dev443}/.flake8 +0 -0
  10. {haliax-1.4.dev441 → haliax-1.4.dev443}/.github/workflows/publish_dev.yaml +0 -0
  11. {haliax-1.4.dev441 → haliax-1.4.dev443}/.github/workflows/run_pre_commit.yaml +0 -0
  12. {haliax-1.4.dev441 → haliax-1.4.dev443}/.github/workflows/run_quick_levanter_tests.yaml +0 -0
  13. {haliax-1.4.dev441 → haliax-1.4.dev443}/.github/workflows/run_tests.yaml +0 -0
  14. {haliax-1.4.dev441 → haliax-1.4.dev443}/.gitignore +0 -0
  15. {haliax-1.4.dev441 → haliax-1.4.dev443}/.playbooks/add-types.md +0 -0
  16. {haliax-1.4.dev441 → haliax-1.4.dev443}/.playbooks/wrap-non-named.md +0 -0
  17. {haliax-1.4.dev441 → haliax-1.4.dev443}/.pre-commit-config.yaml +0 -0
  18. {haliax-1.4.dev441 → haliax-1.4.dev443}/.readthedocs.yaml +0 -0
  19. {haliax-1.4.dev441 → haliax-1.4.dev443}/AGENTS.md +0 -0
  20. {haliax-1.4.dev441 → haliax-1.4.dev443}/AUTHORS.md +0 -0
  21. {haliax-1.4.dev441 → haliax-1.4.dev443}/CONTRIBUTING.md +0 -0
  22. {haliax-1.4.dev441 → haliax-1.4.dev443}/CONTRIBUTORS.md +0 -0
  23. {haliax-1.4.dev441 → haliax-1.4.dev443}/LICENSE +0 -0
  24. {haliax-1.4.dev441 → haliax-1.4.dev443}/README.md +0 -0
  25. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/api.md +0 -0
  26. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/broadcasting.md +0 -0
  27. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/cheatsheet.md +0 -0
  28. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/css/material.css +0 -0
  29. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/css/mkdocstrings.css +0 -0
  30. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/faq.md +0 -0
  31. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/data_parallel_mesh.png +0 -0
  32. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/data_parallel_mesh_replicated.png +0 -0
  33. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/device_mesh_1d.png +0 -0
  34. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/device_mesh_1d_zero.png +0 -0
  35. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/device_mesh_2d.png +0 -0
  36. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/device_mesh_2d_batch_partitioned.png +0 -0
  37. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/device_mesh_2d_data_replicated.png +0 -0
  38. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/device_mesh_2d_data_replicated_mlp_partitioned.png +0 -0
  39. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/device_mesh_2d_intermediate_fully_partitioned.png +0 -0
  40. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/figures/device_mesh_2d_zero.png +0 -0
  41. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/fp8.md +0 -0
  42. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/index.md +0 -0
  43. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/indexing.md +0 -0
  44. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/matmul.md +0 -0
  45. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/nn.md +0 -0
  46. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/partitioning.md +0 -0
  47. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/primer.md +0 -0
  48. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/rearrange.ipynb +0 -0
  49. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/rearrange.md +0 -0
  50. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/requirements.txt +0 -0
  51. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/scan.md +0 -0
  52. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/state-dict.md +0 -0
  53. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/tutorial.md +0 -0
  54. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/typing.md +0 -0
  55. {haliax-1.4.dev441 → haliax-1.4.dev443}/docs/vmap.md +0 -0
  56. {haliax-1.4.dev441 → haliax-1.4.dev443}/etc/license_header.txt +0 -0
  57. {haliax-1.4.dev441 → haliax-1.4.dev443}/mkdocs.yml +0 -0
  58. {haliax-1.4.dev441 → haliax-1.4.dev443}/pyproject.toml +0 -0
  59. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/__init__.py +0 -0
  60. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/__init__.py +0 -0
  61. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/compile_utils.py +0 -0
  62. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/dot.py +0 -0
  63. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/einsum.py +0 -0
  64. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/fp8.py +0 -0
  65. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/parsing.py +0 -0
  66. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/rearrange.py +0 -0
  67. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/scan.py +0 -0
  68. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/state_dict.py +0 -0
  69. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/_src/util.py +0 -0
  70. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/axis.py +0 -0
  71. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/core.py +0 -0
  72. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/debug.py +0 -0
  73. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/fft.py +0 -0
  74. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/field.py +0 -0
  75. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/haxtyping.py +0 -0
  76. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/hof.py +0 -0
  77. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/jax_utils.py +0 -0
  78. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/__init__.py +0 -0
  79. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/activations.py +0 -0
  80. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/attention.py +0 -0
  81. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/conv.py +0 -0
  82. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/dropout.py +0 -0
  83. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/loss.py +0 -0
  84. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/mlp.py +0 -0
  85. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/mup.py +0 -0
  86. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/normalization.py +0 -0
  87. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/pool.py +0 -0
  88. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/nn/scan.py +0 -0
  89. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/ops.py +0 -0
  90. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/partitioning.py +0 -0
  91. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/poly.py +0 -0
  92. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/quantization.py +0 -0
  93. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/random.py +0 -0
  94. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/specialized_fns.py +0 -0
  95. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/state_dict.py +0 -0
  96. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/tree.py +0 -0
  97. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/tree_util.py +0 -0
  98. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/types.py +0 -0
  99. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/util.py +0 -0
  100. {haliax-1.4.dev441 → haliax-1.4.dev443}/src/haliax/wrap.py +0 -0
  101. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/core_test.py +0 -0
  102. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_attention.py +0 -0
  103. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_axis.py +0 -0
  104. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_bitwise_ops.py +0 -0
  105. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_conv.py +0 -0
  106. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_debug.py +0 -0
  107. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_dot.py +0 -0
  108. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_dtype_typing.py +0 -0
  109. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_einsum.py +0 -0
  110. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_fft.py +0 -0
  111. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_field.py +0 -0
  112. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_fp8.py +0 -0
  113. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_hof.py +0 -0
  114. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_int8.py +0 -0
  115. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_moe_linear.py +0 -0
  116. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_mup_coordinate_check.py +0 -0
  117. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_namedarray_typing.py +0 -0
  118. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_nan_reductions.py +0 -0
  119. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_nn.py +0 -0
  120. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_ops.py +0 -0
  121. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_parsing.py +0 -0
  122. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_partitioning.py +0 -0
  123. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_poly_ops.py +0 -0
  124. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_pool.py +0 -0
  125. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_random.py +0 -0
  126. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_rearrange.py +0 -0
  127. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_scan.py +0 -0
  128. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_scatter_gather.py +0 -0
  129. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_specialized_fns.py +0 -0
  130. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_state_dict.py +0 -0
  131. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_tree_util.py +0 -0
  132. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_utils.py +0 -0
  133. {haliax-1.4.dev441 → haliax-1.4.dev443}/tests/test_visualize_sharding.py +0 -0
  134. {haliax-1.4.dev441 → haliax-1.4.dev443}/uv.lock +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: haliax
3
- Version: 1.4.dev441
3
+ Version: 1.4.dev443
4
4
  Summary: Named Tensors for Legible Deep Learning in JAX
5
5
  Project-URL: Homepage, https://github.com/stanford-crfm/haliax
6
6
  Project-URL: Bug Tracker, https://github.com/stanford-crfm/haliax/issues/
@@ -3,4 +3,4 @@
3
3
  # SPDX-License-Identifier: Apache-2.0
4
4
 
5
5
 
6
- __version__ = "1.4.dev441"
6
+ __version__ = "1.4.dev443"
@@ -24,7 +24,12 @@ class Embedding(eqx.Module, ReparamEnabled):
24
24
  # axes
25
25
  Vocab: Axis = eqx.field(static=True)
26
26
  Embed: AxisSpec = eqx.field(static=True)
27
- reparam: AbstractEmbeddingReparam = eqx.field(static=True)
27
+
28
+ _reparam_cls: type[AbstractEmbeddingReparam] = eqx.field(static=True, default=EmbeddingStandardParam)
29
+
30
+ @property
31
+ def reparam(self) -> AbstractEmbeddingReparam:
32
+ return self._reparam_cls(self.Embed, self.Vocab)
28
33
 
29
34
  @staticmethod
30
35
  def init(
@@ -61,7 +66,7 @@ class Embedding(eqx.Module, ReparamEnabled):
61
66
  weight = hax.random.truncated_normal(key, all_axes, -3, 3) * (
62
67
  init_scale * reparam_cls.init_scale(Vocab, Embed)
63
68
  )
64
- return Embedding(weight=weight, Vocab=Vocab, Embed=Embed, reparam=reparam_cls(Embed, Vocab))
69
+ return Embedding(weight=weight, Vocab=Vocab, Embed=Embed, _reparam_cls=reparam_cls)
65
70
 
66
71
  def __call__(self, input_ids: NamedArray, *, key: PRNGKeyArray | None = None):
67
72
  """Alias for `embed`. key is ignored."""
@@ -45,9 +45,14 @@ class Linear(ModuleWithStateDictSerialization, ReparamEnabled):
45
45
 
46
46
  In: AxisSpec = eqx.field(static=True)
47
47
  Out: AxisSpec = eqx.field(static=True)
48
- reparam: AbstractLinearReparam = eqx.field(static=True)
49
48
  dot_general: DotGeneralOp = eqx.field(default_factory=DotGeneralOp.default)
50
49
 
50
+ _reparam_cls: type[AbstractLinearReparam] = eqx.field(static=True, default=LinearStandardParam)
51
+
52
+ @property
53
+ def reparam(self) -> AbstractLinearReparam:
54
+ return self._reparam_cls(self.In, self.Out)
55
+
51
56
  @staticmethod
52
57
  def init(
53
58
  In: AxisSpec,
@@ -77,7 +82,7 @@ class Linear(ModuleWithStateDictSerialization, ReparamEnabled):
77
82
  if dot_general is None:
78
83
  dot_general = DotGeneralOp.default()
79
84
 
80
- return Linear(weight, bias, In, Out, dot_general=dot_general, reparam=reparam_cls(In, Out))
85
+ return Linear(weight, bias, In, Out, dot_general=dot_general, _reparam_cls=reparam_cls)
81
86
 
82
87
  @named_call
83
88
  def __call__(self, inputs, *, key: PRNGKeyArray | None = None):
@@ -31,7 +31,7 @@ def test_mup_embedding_unembedding_scale():
31
31
  Embed = (hax.Axis("E", 3),)
32
32
 
33
33
  weight = hax.ones(hax.concat_axis_specs(Vocab, Embed))
34
- layer = Embedding(weight=weight, Vocab=Vocab, Embed=Embed, reparam=EmbeddingMup(Embed, Vocab))
34
+ layer = Embedding(weight=weight, Vocab=Vocab, Embed=Embed, _reparam_cls=EmbeddingMup)
35
35
 
36
36
  scale = layer.reparam.unembed_active_scale
37
37
  assert scale == pytest.approx(1.0 / hax.axis_size(Embed))
@@ -11,7 +11,12 @@ import pytest
11
11
 
12
12
  import haliax as hax
13
13
  from haliax.nn import Linear
14
- from haliax.nn.mup import InputLinearMup, LinearStandardParam, HiddenLinearMup, OutputLinearMup
14
+ from haliax.nn.mup import (
15
+ InputLinearMup,
16
+ LinearStandardParam,
17
+ HiddenLinearMup,
18
+ OutputLinearMup,
19
+ )
15
20
 
16
21
 
17
22
  @pytest.mark.parametrize("out_first", [True, False])
@@ -37,8 +42,8 @@ def test_mup_linear_call_matches_linear():
37
42
  weight = hax.ones(hax.concat_axis_specs(Out, In)) * 0.5
38
43
  bias = hax.full(Out, 0.25)
39
44
 
40
- linear = Linear(weight, bias, In, Out, reparam=LinearStandardParam(In, Out))
41
- mup = Linear(weight, bias, In, Out, reparam=InputLinearMup(In, Out))
45
+ linear = Linear(weight, bias, In, Out, _reparam_cls=LinearStandardParam)
46
+ mup = Linear(weight, bias, In, Out, _reparam_cls=InputLinearMup)
42
47
 
43
48
  inputs = hax.full(hax.concat_axis_specs(Batch, In), 2.0)
44
49
 
@@ -109,8 +114,8 @@ def test_input_linear_behaves_like_base_linear():
109
114
  weight = hax.ones((Out, In)) * 0.1
110
115
  bias = hax.zeros(Out)
111
116
 
112
- linear = Linear(weight, bias, In, Out, reparam=LinearStandardParam(In, Out))
113
- input_linear = Linear(weight, bias, In, Out, reparam=InputLinearMup(In, Out))
117
+ linear = Linear(weight, bias, In, Out, _reparam_cls=LinearStandardParam)
118
+ input_linear = Linear(weight, bias, In, Out, _reparam_cls=InputLinearMup)
114
119
 
115
120
  inputs = hax.random.normal(jrandom.PRNGKey(5), (Batch, In))
116
121
 
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes