anyinit 0.1.0__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 (90) hide show
  1. anyinit-0.1.0/.gitignore +14 -0
  2. anyinit-0.1.0/CHANGELOG.md +30 -0
  3. anyinit-0.1.0/LICENSE +21 -0
  4. anyinit-0.1.0/PKG-INFO +161 -0
  5. anyinit-0.1.0/README.md +129 -0
  6. anyinit-0.1.0/docs/api.md +44 -0
  7. anyinit-0.1.0/docs/changelog.md +1 -0
  8. anyinit-0.1.0/docs/contributing.md +40 -0
  9. anyinit-0.1.0/docs/design.md +104 -0
  10. anyinit-0.1.0/docs/experiments/README.md +21 -0
  11. anyinit-0.1.0/docs/experiments/depth_sweep.py +80 -0
  12. anyinit-0.1.0/docs/experiments/draw_drift.py +51 -0
  13. anyinit-0.1.0/docs/experiments/gain_table.py +84 -0
  14. anyinit-0.1.0/docs/experiments/gaussian_breakdown.py +57 -0
  15. anyinit-0.1.0/docs/experiments/quadrature_accuracy.py +63 -0
  16. anyinit-0.1.0/docs/experiments/stability_table.py +86 -0
  17. anyinit-0.1.0/docs/guide/frameworks.md +39 -0
  18. anyinit-0.1.0/docs/guide/modes.md +44 -0
  19. anyinit-0.1.0/docs/guide/options.md +54 -0
  20. anyinit-0.1.0/docs/guide/report.md +48 -0
  21. anyinit-0.1.0/docs/guide/stability.md +61 -0
  22. anyinit-0.1.0/docs/index.md +1 -0
  23. anyinit-0.1.0/docs/reference.md +51 -0
  24. anyinit-0.1.0/docs/requirements.txt +4 -0
  25. anyinit-0.1.0/docs/tutorials/custom-activation.md +31 -0
  26. anyinit-0.1.0/docs/tutorials/depth.md +14 -0
  27. anyinit-0.1.0/docs/tutorials/multiple-frameworks.md +17 -0
  28. anyinit-0.1.0/docs/tutorials/pinn.md +19 -0
  29. anyinit-0.1.0/docs/tutorials/resnet.md +16 -0
  30. anyinit-0.1.0/docs/why.md +45 -0
  31. anyinit-0.1.0/examples/custom_activation.py +69 -0
  32. anyinit-0.1.0/examples/depth_matters.py +93 -0
  33. anyinit-0.1.0/examples/multibackend.py +120 -0
  34. anyinit-0.1.0/examples/pinn_relu2.py +123 -0
  35. anyinit-0.1.0/examples/resnet_cifar.py +119 -0
  36. anyinit-0.1.0/pyproject.toml +107 -0
  37. anyinit-0.1.0/src/anyinit/__init__.py +298 -0
  38. anyinit-0.1.0/src/anyinit/_run.py +429 -0
  39. anyinit-0.1.0/src/anyinit/backends/__init__.py +88 -0
  40. anyinit-0.1.0/src/anyinit/backends/base.py +254 -0
  41. anyinit-0.1.0/src/anyinit/backends/jax.py +664 -0
  42. anyinit-0.1.0/src/anyinit/backends/keras.py +674 -0
  43. anyinit-0.1.0/src/anyinit/backends/pytorch.py +1170 -0
  44. anyinit-0.1.0/src/anyinit/config.py +86 -0
  45. anyinit-0.1.0/src/anyinit/core/__init__.py +1 -0
  46. anyinit-0.1.0/src/anyinit/core/activations.py +172 -0
  47. anyinit-0.1.0/src/anyinit/core/analytic.py +497 -0
  48. anyinit-0.1.0/src/anyinit/core/distributions.py +124 -0
  49. anyinit-0.1.0/src/anyinit/core/empirical.py +177 -0
  50. anyinit-0.1.0/src/anyinit/core/fan.py +76 -0
  51. anyinit-0.1.0/src/anyinit/core/graph.py +154 -0
  52. anyinit-0.1.0/src/anyinit/core/moments.py +72 -0
  53. anyinit-0.1.0/src/anyinit/core/profile.py +285 -0
  54. anyinit-0.1.0/src/anyinit/core/quadrature.py +237 -0
  55. anyinit-0.1.0/src/anyinit/core/registry.py +204 -0
  56. anyinit-0.1.0/src/anyinit/core/solve.py +145 -0
  57. anyinit-0.1.0/src/anyinit/core/stability.py +208 -0
  58. anyinit-0.1.0/src/anyinit/core/topology.py +184 -0
  59. anyinit-0.1.0/src/anyinit/core/transfer.py +170 -0
  60. anyinit-0.1.0/src/anyinit/errors.py +27 -0
  61. anyinit-0.1.0/src/anyinit/py.typed +0 -0
  62. anyinit-0.1.0/src/anyinit/report.py +348 -0
  63. anyinit-0.1.0/tests/__init__.py +0 -0
  64. anyinit-0.1.0/tests/conftest.py +112 -0
  65. anyinit-0.1.0/tests/helpers.py +76 -0
  66. anyinit-0.1.0/tests/test_agreement.py +56 -0
  67. anyinit-0.1.0/tests/test_api.py +125 -0
  68. anyinit-0.1.0/tests/test_backend_jax.py +225 -0
  69. anyinit-0.1.0/tests/test_backend_keras.py +153 -0
  70. anyinit-0.1.0/tests/test_backend_pytorch.py +206 -0
  71. anyinit-0.1.0/tests/test_core_isolation.py +66 -0
  72. anyinit-0.1.0/tests/test_cross_backend.py +126 -0
  73. anyinit-0.1.0/tests/test_custom_activation.py +200 -0
  74. anyinit-0.1.0/tests/test_distributions.py +148 -0
  75. anyinit-0.1.0/tests/test_docs_tables.py +31 -0
  76. anyinit-0.1.0/tests/test_errors.py +71 -0
  77. anyinit-0.1.0/tests/test_fan.py +61 -0
  78. anyinit-0.1.0/tests/test_graph.py +53 -0
  79. anyinit-0.1.0/tests/test_idempotence.py +97 -0
  80. anyinit-0.1.0/tests/test_input_moments.py +76 -0
  81. anyinit-0.1.0/tests/test_optional_dependencies.py +75 -0
  82. anyinit-0.1.0/tests/test_pooling.py +171 -0
  83. anyinit-0.1.0/tests/test_profile.py +92 -0
  84. anyinit-0.1.0/tests/test_property.py +93 -0
  85. anyinit-0.1.0/tests/test_quadrature.py +104 -0
  86. anyinit-0.1.0/tests/test_regressions.py +452 -0
  87. anyinit-0.1.0/tests/test_report.py +175 -0
  88. anyinit-0.1.0/tests/test_stability.py +125 -0
  89. anyinit-0.1.0/tests/test_topology.py +157 -0
  90. anyinit-0.1.0/tests/test_transfer.py +128 -0
@@ -0,0 +1,14 @@
1
+ __pycache__/
2
+ *.py[cod]
3
+ .pytest_cache/
4
+ .mypy_cache/
5
+ .ruff_cache/
6
+ .coverage
7
+ htmlcov/
8
+ build/
9
+ dist/
10
+ *.egg-info/
11
+ .venv/
12
+ data/
13
+ _old_code/
14
+ site/
@@ -0,0 +1,30 @@
1
+ # Changelog
2
+
3
+ ## 0.1.0
4
+
5
+ First release.
6
+
7
+ ### Added
8
+
9
+ - `initialize(model, mode, input_spec, ...)` — one call initializes any model, with
10
+ the framework detected from it. The remaining options are explicit keywords:
11
+ `distribution`, `center`, `gains`, `seed` and `params`. PyTorch, Keras 3 (on any of
12
+ its backends) and JAX (via Flax), each an optional dependency, so installing one never
13
+ imports another.
14
+ - `initialize_params(model, params, ...)` for functional frameworks, which return a new
15
+ parameter tree rather than mutating one.
16
+ - Two solve modes. `analytic` propagates moments through the graph and never runs the
17
+ model; `empirical` measures real batches and assumes nothing distributional. They share
18
+ the graph, the layer adapters and the target policy.
19
+ - `register_activation` for arbitrary activations. Their moment maps are computed by
20
+ Gauss–Legendre quadrature over the Gaussian, exact to machine precision across kinks.
21
+ - `anyinit.gain(name, **params)`: the gain AnyInit gives an activation, `√2` for ReLU.
22
+ `initialize(gains={...})` fixes the gain for an activation instead of solving for it.
23
+ - Depth-stability diagnosis. Every activation gets `χ = d log E[f²] / d log Var[z]`, which
24
+ equals the homogeneity degree for homogeneous activations and so identifies the ones no
25
+ scalar initialization can stabilize.
26
+ - `InitReport`: the scale chosen for every layer, the graph fidelity actually achieved, a
27
+ measured validation pass sized against its own sampling noise, and anything AnyInit
28
+ could not do. `assert_healthy()` raises on the findings.
29
+ - Weights are drawn in NumPy in a canonical layout from the run's seed, so the same
30
+ architecture and seed produce bit-identical weights under all three frameworks.
anyinit-0.1.0/LICENSE ADDED
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2026 Jose I. Mestre
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
anyinit-0.1.0/PKG-INFO ADDED
@@ -0,0 +1,161 @@
1
+ Metadata-Version: 2.5
2
+ Name: anyinit
3
+ Version: 0.1.0
4
+ Summary: Initialize any model, in any framework, correctly — with one call
5
+ Project-URL: Homepage, https://github.com/jmiravet/AnyInit
6
+ Project-URL: Documentation, https://jmiravet.github.io/AnyInit/
7
+ Project-URL: Repository, https://github.com/jmiravet/AnyInit
8
+ Project-URL: Issues, https://github.com/jmiravet/AnyInit/issues
9
+ Project-URL: Changelog, https://github.com/jmiravet/AnyInit/blob/main/CHANGELOG.md
10
+ Author-email: "Jose I. Mestre" <jmiravet@uji.es>, Alberto Fernández-Hernández <a.fernandez@upv.es>
11
+ Maintainer-email: "Jose I. Mestre" <jmiravet@uji.es>
12
+ License-Expression: MIT
13
+ License-File: LICENSE
14
+ Keywords: deep-learning,flax,initialization,jax,keras,pytorch,signal-propagation,tensorflow,weight-initialization
15
+ Classifier: Development Status :: 4 - Beta
16
+ Classifier: Intended Audience :: Science/Research
17
+ Classifier: Programming Language :: Python :: 3
18
+ Classifier: Programming Language :: Python :: 3.10
19
+ Classifier: Programming Language :: Python :: 3.11
20
+ Classifier: Programming Language :: Python :: 3.12
21
+ Classifier: Programming Language :: Python :: 3.13
22
+ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
23
+ Classifier: Typing :: Typed
24
+ Requires-Python: >=3.10
25
+ Requires-Dist: numpy>=1.21
26
+ Provides-Extra: dev
27
+ Requires-Dist: mypy; extra == 'dev'
28
+ Requires-Dist: pytest-cov; extra == 'dev'
29
+ Requires-Dist: pytest>=7; extra == 'dev'
30
+ Requires-Dist: ruff; extra == 'dev'
31
+ Description-Content-Type: text/markdown
32
+
33
+ # AnyInit
34
+
35
+ [![CI](https://github.com/jmiravet/AnyInit/actions/workflows/ci.yml/badge.svg)](https://github.com/jmiravet/AnyInit/actions/workflows/ci.yml)
36
+ [![codecov](https://codecov.io/gh/jmiravet/AnyInit/graph/badge.svg)](https://codecov.io/gh/jmiravet/AnyInit)
37
+ [![PyPI](https://img.shields.io/pypi/v/anyinit)](https://pypi.org/project/anyinit/)
38
+ [![Python](https://img.shields.io/pypi/pyversions/anyinit)](https://pypi.org/project/anyinit/)
39
+ [![Docs](https://img.shields.io/badge/docs-jmiravet.github.io%2FAnyInit-blue)](https://jmiravet.github.io/AnyInit/)
40
+ [![License: MIT](https://img.shields.io/badge/license-MIT-blue)](https://github.com/jmiravet/AnyInit/blob/main/LICENSE)
41
+
42
+ Initialize any model, in any framework, correctly — with one call.
43
+
44
+ ```python
45
+ import anyinit
46
+
47
+ report = anyinit.initialize(model)
48
+ print(report)
49
+ ```
50
+
51
+ `model` can be a PyTorch module, a Keras model or a Flax module. AnyInit traces it, finds
52
+ which activation follows each layer, and scales every weight so the signal neither dies
53
+ nor explodes on its way through. Nothing to configure, nothing to look up.
54
+
55
+ ## Install
56
+
57
+ ```bash
58
+ pip install anyinit
59
+ ```
60
+
61
+ AnyInit depends on NumPy alone and uses whichever framework your model comes from.
62
+
63
+ ## What "any" means
64
+
65
+ **Any framework.** PyTorch, Keras 3 (on TensorFlow, JAX or PyTorch) and Flax, detected
66
+ from the model. The same seed gives the same weights in all three.
67
+
68
+ **Any activation.** AnyInit measures the function itself instead of looking it up in a
69
+ table, so a new activation is one decorator away:
70
+
71
+ ```python
72
+ @anyinit.register_activation
73
+ class ReLUCubed(torch.nn.Module):
74
+ def forward(self, x):
75
+ return torch.relu(x) ** 3
76
+ ```
77
+
78
+ **Any architecture.** Branches, residual additions, concatenations, normalization,
79
+ pooling, dropout, attention and embeddings.
80
+
81
+ **Theory or measurement.** Scale from theory, with no data, or from real batches:
82
+
83
+ ```python
84
+ anyinit.initialize(model, "analytic", (32, 3, 224, 224)) # shapes only, milliseconds
85
+ anyinit.initialize(model, "empirical", batch) # measures, assumes nothing
86
+ ```
87
+
88
+ ## It tells you when it cannot help
89
+
90
+ Some activations cannot be kept stable across depth by any initialization. AnyInit says
91
+ so in the report instead of handing back a network that will not train:
92
+
93
+ ```
94
+ Stability
95
+ relu3 chi= 3.000 sigma*= 0.7148 unstable degree=3
96
+ ! relu3 is homogeneous of degree 3 (chi=3.000), so a relative error grows by
97
+ 3.00x per layer and reaches 3.49e+09x over 20 layers. No scalar initialization
98
+ is depth-stable here: reduce depth, insert normalization, or use a degree-one
99
+ activation
100
+ ```
101
+
102
+ Call `report.assert_healthy()` to turn that into an exception, for example in CI.
103
+
104
+ ## Learn more
105
+
106
+ The [documentation](https://jmiravet.github.io/AnyInit/) covers:
107
+
108
+ - [Reading the report](https://jmiravet.github.io/AnyInit/guide/report/)
109
+ - [Analytic or empirical](https://jmiravet.github.io/AnyInit/guide/modes/)
110
+ - [Stability across depth](https://jmiravet.github.io/AnyInit/guide/stability/)
111
+ - [Frameworks](https://jmiravet.github.io/AnyInit/guide/frameworks/)
112
+ - [Options](https://jmiravet.github.io/AnyInit/guide/options/)
113
+ - [Why a gain table is not enough](https://jmiravet.github.io/AnyInit/why/), with measurements
114
+ - [Tutorials](https://jmiravet.github.io/AnyInit/tutorials/custom-activation/) and the
115
+ [API](https://jmiravet.github.io/AnyInit/api/)
116
+
117
+ Contributions are welcome; see
118
+ [Contributing](https://jmiravet.github.io/AnyInit/contributing/).
119
+
120
+ ## References
121
+
122
+ Variance scaling comes from Glorot and Bengio (2010) and He et al. (2015), whose rectifier
123
+ gain falls out of the moment map as the degree-one case. Treating signal propagation as a
124
+ dynamical system, and the order parameter behind the `χ` diagnostic, come from Poole et al.
125
+ (2016) and Schoenholz et al. (2017). The empirical mode generalizes LSUV (Mishkin & Matas,
126
+ 2016) from a sequence of layers to a graph. Fixed points of the variance map, and the SELU
127
+ constants the reference table reproduces, are from Klambauer et al. (2017). The
128
+ `sinusoidal` distribution is from Fernández-Hernández et al. (2025).
129
+
130
+ What is new is evaluating the moment map numerically, by Gaussian quadrature, instead of
131
+ deriving it per activation, which is why AnyInit accepts an activation it has never seen.
132
+
133
+ Fernández-Hernández, A., Mestre, J. I., Dolz, M. F., Duato, J., & Quintana-Ortí, E. S.
134
+ (2025). Sinusoidal initialization, time for a new start. In *Advances in Neural
135
+ Information Processing Systems* (Vol. 38). https://doi.org/10.48550/arXiv.2505.12909
136
+
137
+ Glorot, X., & Bengio, Y. (2010). Understanding the difficulty of training deep feedforward
138
+ neural networks. In *Proceedings of the Thirteenth International Conference on Artificial
139
+ Intelligence and Statistics* (pp. 249–256). PMLR.
140
+ https://proceedings.mlr.press/v9/glorot10a.html
141
+
142
+ He, K., Zhang, X., Ren, S., & Sun, J. (2015). Delving deep into rectifiers: Surpassing
143
+ human-level performance on ImageNet classification. In *Proceedings of the IEEE
144
+ International Conference on Computer Vision* (pp. 1026–1034). IEEE.
145
+ https://doi.org/10.1109/ICCV.2015.123
146
+
147
+ Klambauer, G., Unterthiner, T., Mayr, A., & Hochreiter, S. (2017). Self-normalizing neural
148
+ networks. In *Advances in Neural Information Processing Systems* (Vol. 30, pp. 971–980).
149
+ https://arxiv.org/abs/1706.02515
150
+
151
+ Mishkin, D., & Matas, J. (2016). All you need is a good init. In *International Conference
152
+ on Learning Representations*. https://arxiv.org/abs/1511.06422
153
+
154
+ Poole, B., Lahiri, S., Raghu, M., Sohl-Dickstein, J., & Ganguli, S. (2016). Exponential
155
+ expressivity in deep neural networks through transient chaos. In *Advances in Neural
156
+ Information Processing Systems* (Vol. 29, pp. 3360–3368).
157
+ https://arxiv.org/abs/1606.05340
158
+
159
+ Schoenholz, S. S., Gilmer, J., Ganguli, S., & Sohl-Dickstein, J. (2017). Deep information
160
+ propagation. In *International Conference on Learning Representations*.
161
+ https://arxiv.org/abs/1611.01232
@@ -0,0 +1,129 @@
1
+ # AnyInit
2
+
3
+ [![CI](https://github.com/jmiravet/AnyInit/actions/workflows/ci.yml/badge.svg)](https://github.com/jmiravet/AnyInit/actions/workflows/ci.yml)
4
+ [![codecov](https://codecov.io/gh/jmiravet/AnyInit/graph/badge.svg)](https://codecov.io/gh/jmiravet/AnyInit)
5
+ [![PyPI](https://img.shields.io/pypi/v/anyinit)](https://pypi.org/project/anyinit/)
6
+ [![Python](https://img.shields.io/pypi/pyversions/anyinit)](https://pypi.org/project/anyinit/)
7
+ [![Docs](https://img.shields.io/badge/docs-jmiravet.github.io%2FAnyInit-blue)](https://jmiravet.github.io/AnyInit/)
8
+ [![License: MIT](https://img.shields.io/badge/license-MIT-blue)](https://github.com/jmiravet/AnyInit/blob/main/LICENSE)
9
+
10
+ Initialize any model, in any framework, correctly — with one call.
11
+
12
+ ```python
13
+ import anyinit
14
+
15
+ report = anyinit.initialize(model)
16
+ print(report)
17
+ ```
18
+
19
+ `model` can be a PyTorch module, a Keras model or a Flax module. AnyInit traces it, finds
20
+ which activation follows each layer, and scales every weight so the signal neither dies
21
+ nor explodes on its way through. Nothing to configure, nothing to look up.
22
+
23
+ ## Install
24
+
25
+ ```bash
26
+ pip install anyinit
27
+ ```
28
+
29
+ AnyInit depends on NumPy alone and uses whichever framework your model comes from.
30
+
31
+ ## What "any" means
32
+
33
+ **Any framework.** PyTorch, Keras 3 (on TensorFlow, JAX or PyTorch) and Flax, detected
34
+ from the model. The same seed gives the same weights in all three.
35
+
36
+ **Any activation.** AnyInit measures the function itself instead of looking it up in a
37
+ table, so a new activation is one decorator away:
38
+
39
+ ```python
40
+ @anyinit.register_activation
41
+ class ReLUCubed(torch.nn.Module):
42
+ def forward(self, x):
43
+ return torch.relu(x) ** 3
44
+ ```
45
+
46
+ **Any architecture.** Branches, residual additions, concatenations, normalization,
47
+ pooling, dropout, attention and embeddings.
48
+
49
+ **Theory or measurement.** Scale from theory, with no data, or from real batches:
50
+
51
+ ```python
52
+ anyinit.initialize(model, "analytic", (32, 3, 224, 224)) # shapes only, milliseconds
53
+ anyinit.initialize(model, "empirical", batch) # measures, assumes nothing
54
+ ```
55
+
56
+ ## It tells you when it cannot help
57
+
58
+ Some activations cannot be kept stable across depth by any initialization. AnyInit says
59
+ so in the report instead of handing back a network that will not train:
60
+
61
+ ```
62
+ Stability
63
+ relu3 chi= 3.000 sigma*= 0.7148 unstable degree=3
64
+ ! relu3 is homogeneous of degree 3 (chi=3.000), so a relative error grows by
65
+ 3.00x per layer and reaches 3.49e+09x over 20 layers. No scalar initialization
66
+ is depth-stable here: reduce depth, insert normalization, or use a degree-one
67
+ activation
68
+ ```
69
+
70
+ Call `report.assert_healthy()` to turn that into an exception, for example in CI.
71
+
72
+ ## Learn more
73
+
74
+ The [documentation](https://jmiravet.github.io/AnyInit/) covers:
75
+
76
+ - [Reading the report](https://jmiravet.github.io/AnyInit/guide/report/)
77
+ - [Analytic or empirical](https://jmiravet.github.io/AnyInit/guide/modes/)
78
+ - [Stability across depth](https://jmiravet.github.io/AnyInit/guide/stability/)
79
+ - [Frameworks](https://jmiravet.github.io/AnyInit/guide/frameworks/)
80
+ - [Options](https://jmiravet.github.io/AnyInit/guide/options/)
81
+ - [Why a gain table is not enough](https://jmiravet.github.io/AnyInit/why/), with measurements
82
+ - [Tutorials](https://jmiravet.github.io/AnyInit/tutorials/custom-activation/) and the
83
+ [API](https://jmiravet.github.io/AnyInit/api/)
84
+
85
+ Contributions are welcome; see
86
+ [Contributing](https://jmiravet.github.io/AnyInit/contributing/).
87
+
88
+ ## References
89
+
90
+ Variance scaling comes from Glorot and Bengio (2010) and He et al. (2015), whose rectifier
91
+ gain falls out of the moment map as the degree-one case. Treating signal propagation as a
92
+ dynamical system, and the order parameter behind the `χ` diagnostic, come from Poole et al.
93
+ (2016) and Schoenholz et al. (2017). The empirical mode generalizes LSUV (Mishkin & Matas,
94
+ 2016) from a sequence of layers to a graph. Fixed points of the variance map, and the SELU
95
+ constants the reference table reproduces, are from Klambauer et al. (2017). The
96
+ `sinusoidal` distribution is from Fernández-Hernández et al. (2025).
97
+
98
+ What is new is evaluating the moment map numerically, by Gaussian quadrature, instead of
99
+ deriving it per activation, which is why AnyInit accepts an activation it has never seen.
100
+
101
+ Fernández-Hernández, A., Mestre, J. I., Dolz, M. F., Duato, J., & Quintana-Ortí, E. S.
102
+ (2025). Sinusoidal initialization, time for a new start. In *Advances in Neural
103
+ Information Processing Systems* (Vol. 38). https://doi.org/10.48550/arXiv.2505.12909
104
+
105
+ Glorot, X., & Bengio, Y. (2010). Understanding the difficulty of training deep feedforward
106
+ neural networks. In *Proceedings of the Thirteenth International Conference on Artificial
107
+ Intelligence and Statistics* (pp. 249–256). PMLR.
108
+ https://proceedings.mlr.press/v9/glorot10a.html
109
+
110
+ He, K., Zhang, X., Ren, S., & Sun, J. (2015). Delving deep into rectifiers: Surpassing
111
+ human-level performance on ImageNet classification. In *Proceedings of the IEEE
112
+ International Conference on Computer Vision* (pp. 1026–1034). IEEE.
113
+ https://doi.org/10.1109/ICCV.2015.123
114
+
115
+ Klambauer, G., Unterthiner, T., Mayr, A., & Hochreiter, S. (2017). Self-normalizing neural
116
+ networks. In *Advances in Neural Information Processing Systems* (Vol. 30, pp. 971–980).
117
+ https://arxiv.org/abs/1706.02515
118
+
119
+ Mishkin, D., & Matas, J. (2016). All you need is a good init. In *International Conference
120
+ on Learning Representations*. https://arxiv.org/abs/1511.06422
121
+
122
+ Poole, B., Lahiri, S., Raghu, M., Sohl-Dickstein, J., & Ganguli, S. (2016). Exponential
123
+ expressivity in deep neural networks through transient chaos. In *Advances in Neural
124
+ Information Processing Systems* (Vol. 29, pp. 3360–3368).
125
+ https://arxiv.org/abs/1606.05340
126
+
127
+ Schoenholz, S. S., Gilmer, J., Ganguli, S., & Sohl-Dickstein, J. (2017). Deep information
128
+ propagation. In *International Conference on Learning Representations*.
129
+ https://arxiv.org/abs/1611.01232
@@ -0,0 +1,44 @@
1
+ # API
2
+
3
+ Everything below is importable from the top-level `anyinit` package.
4
+
5
+ ## Initializing
6
+
7
+ ::: anyinit.initialize
8
+
9
+ ::: anyinit.initialize_params
10
+
11
+ ## Activations
12
+
13
+ ::: anyinit.register_activation
14
+
15
+ ::: anyinit.unregister_activation
16
+
17
+ ::: anyinit.registered_activations
18
+
19
+ ::: anyinit.gain
20
+
21
+ ::: anyinit.activation_profile
22
+
23
+ ::: anyinit.ActivationProfile
24
+
25
+ ## The report
26
+
27
+ ::: anyinit.InitReport
28
+
29
+ ::: anyinit.LayerRecord
30
+
31
+ ::: anyinit.StabilityRecord
32
+
33
+ ## Backends
34
+
35
+ ::: anyinit.available_backends
36
+
37
+ ::: anyinit.known_backends
38
+
39
+ ## Errors
40
+
41
+ ::: anyinit.errors
42
+ options:
43
+ show_root_heading: false
44
+ members_order: source
@@ -0,0 +1 @@
1
+ --8<-- "CHANGELOG.md"
@@ -0,0 +1,40 @@
1
+ # Contributing
2
+
3
+ ## Setup
4
+
5
+ ```bash
6
+ pip install -e '.[dev]' torch tensorflow jax flax
7
+ pytest # backend tests skip when a framework is absent
8
+ ruff check src tests examples docs
9
+ ruff format --check src tests examples docs
10
+ mypy src # strict
11
+
12
+ pip install -r docs/requirements.txt
13
+ mkdocs serve # the documentation site, locally
14
+ ```
15
+
16
+ [Design](design.md) explains how the code is laid out.
17
+
18
+ ## Adding a backend
19
+
20
+ Implement `anyinit.backends.base.Backend` and add it to `_BACKENDS` in
21
+ `anyinit/backends/__init__.py`. None of the numerics is touched. The contract is:
22
+
23
+ - `handles(model)` — recognize the model from its class hierarchy, without importing the
24
+ framework. `module_roots` does the inspection.
25
+ - `build_graph(model, input_spec)` — trace into a `ModelGraph`.
26
+ - `eval_elementwise(fn, x)` — apply a native callable to quadrature abscissas.
27
+ - `read_weight` / `write_weight` / `write_bias` / `write_gain` — canonical layout in and
28
+ out. `write_gain` assigns rather than multiplies, so `initialize` stays idempotent.
29
+ - `forward_taps(model, inputs, taps, training=True)` — one forward pass, returning moments
30
+ at the named nodes. `training` must normally be true: at initialization a batch-norm
31
+ layer's running statistics are still (0, 1), so in inference mode it scales by gamma
32
+ instead of normalizing. Running statistics must be restored afterward, and random state
33
+ fixed during the pass, so that dropout draws the same masks every time and the empirical
34
+ solver sees a deterministic function of the scales.
35
+ - `unscaled_weights(model, graph)` — every weight no graph node writes, with the reason.
36
+ - `label` — how the report names the backend; Keras adds what it runs on, `keras[jax]`.
37
+ - `finalize(model)` — new parameters for a functional framework, `None` otherwise.
38
+
39
+ `tests/test_cross_backend.py` then checks that the new backend produces the same weights
40
+ as the others from the same seed.
@@ -0,0 +1,104 @@
1
+ # Design
2
+
3
+ Notes for anyone reading or extending the code. To add a framework, see
4
+ [Contributing](contributing.md#adding-a-backend).
5
+
6
+ ## Two layers
7
+
8
+ ```
9
+ ┌──────────────────────────────────────────────────────┐
10
+ │ anyinit.core — imports no framework │
11
+ │ quadrature · profiles · stability · graph IR · │
12
+ │ topology · transfer · fan · distributions · solvers │
13
+ └───────────────────────┬──────────────────────────────┘
14
+ │ ModelGraph + numpy.ndarray
15
+ ┌───────────────────────┴──────────────────────────────┐
16
+ │ anyinit.backends — one module per framework │
17
+ │ pytorch.py · keras.py · jax.py │
18
+ │ tracing · weight layouts · forward passes · taps │
19
+ └──────────────────────────────────────────────────────┘
20
+ ```
21
+
22
+ The core never holds a tensor. It speaks `ModelGraph` — an IR of typed nodes — and
23
+ `numpy.ndarray` in a canonical layout. Backends trace the native model into that IR,
24
+ translate layouts, read and write parameters, and run forward passes returning
25
+ *statistics* rather than tensors.
26
+
27
+ ## Principles
28
+
29
+ 1. **The core imports no framework.** `tests/test_core_isolation.py` walks the AST of
30
+ every core module to enforce it, and runs a subprocess to check that importing
31
+ `anyinit` loads none of them.
32
+ 2. **Lazy imports.** A framework is imported inside the method that needs it, never at
33
+ module scope, so installing one never pulls in another.
34
+ 3. **Reinitialize, not rescale.** The analytic recursion assumes `W` is i.i.d. zero-mean
35
+ with variance `σ_w²`; drawing the weights makes that true by construction.
36
+ 4. **One graph, two solvers.** `analytic` and `empirical` share the IR, the topology, the
37
+ layer adapters, the target policy and the report. They differ only in where the
38
+ moments come from.
39
+ 5. **Mode purity.** No measurement ever sets an `analytic` scale. The model is run only
40
+ for the validation pass, which runs when `input_spec` is given and whose result
41
+ reaches the report alone, and, under Flax, for a 16-row probe that identifies
42
+ activations (below). Shapes come from fake tensors,
43
+ which carry no data. `empirical` never consults a Gaussian profile to set a scale.
44
+ 6. **Weights drawn in NumPy.** The same seed gives the same weights under every backend.
45
+ 7. **Fail visibly.** Anything AnyInit cannot do correctly goes into the report rather
46
+ than being approximated silently.
47
+
48
+ ## Canonical weight layout
49
+
50
+ The core reasons in `(fan_out, fan_in, *receptive_field)`. Each backend translates:
51
+
52
+ | Layer | Canonical | PyTorch | Keras | Flax |
53
+ |---|---|---|---|---|
54
+ | Dense | `(out, in)` | `(out, in)` | `(in, out)` | `(in, out)` |
55
+ | Conv2D | `(out, in/g, kh, kw)` | `(out, in/g, kh, kw)` | `(kh, kw, in, out)` | `(kh, kw, in, out)` |
56
+ | ConvTranspose2D | `(out, in/g, kh, kw)` | `(in, out/g, kh, kw)` | `(kh, kw, out, in)` | `(kh, kw, in, out)` |
57
+ | Embedding | `(num, dim)` | `(num, dim)` | `(num, dim)` | `(num, dim)` |
58
+
59
+ A transposed convolution stores its input and output axes the other way round from an
60
+ ordinary one, so fans computed on the native shape come out inverted.
61
+
62
+ ## Tracing
63
+
64
+ | Framework | Mechanism | Fidelity |
65
+ |---|---|---|
66
+ | PyTorch | `torch.fx` symbolic trace; shapes from fake-tensor propagation | graph; falls back to a linear chain from `named_modules` when tracing fails |
67
+ | Keras | `model.operations` and each operation's inbound nodes | graph; linear for an unwired Sequential |
68
+ | Flax | Flax layers and `jax.nn` activations wrapped inside a context manager, then the model applied to a 16-row probe | linear |
69
+
70
+ PyTorch's `nn.Transformer*` layers are traced as single modules, so the adapter expands
71
+ each into its real structure — attention, residual additions, normalizations and the
72
+ feed-forward block — with both `norm_first` layouts. The expanded nodes have no FX node of
73
+ their own and are measured through hooks on the submodules that produce or consume them.
74
+ Adaptive pooling records no window; the fake-tensor shapes supply it.
75
+
76
+ A jaxpr is too far from the source to read activations off: most `jax.nn` functions lower
77
+ to opaque wrappers and a bias add is indistinguishable from a residual one. The
78
+ instrumented trace recovers call order instead, and each layer reports `scope.path`, which
79
+ is its key in the parameter tree. That yields an ordered chain rather than a DAG, since
80
+ nothing intercepts `+`, and the graph is marked `linear`. An activation that never passes
81
+ through `jax.nn` — a module field, a closure, a lambda — is not seen being called, so the
82
+ probe's values identify it: wherever one layer's output reaches the next transformed, the
83
+ transform is matched against the known and registered activations.
84
+
85
+ Whatever a backend reaches but cannot scale — recurrent weights, attention layers it has no
86
+ adapter for, bare parameters — is listed in the report with the reason, never skipped
87
+ silently.
88
+
89
+ ## Where things live
90
+
91
+ | Concern | Module |
92
+ |---|---|
93
+ | Gaussian integration, kink detection | `core/quadrature.py` |
94
+ | NumPy reference activations | `core/activations.py` |
95
+ | Moment map of an activation | `core/profile.py` |
96
+ | Activation registry | `core/registry.py` |
97
+ | `χ`, feasibility, verdicts | `core/stability.py` |
98
+ | Typed graph IR | `core/graph.py` |
99
+ | Layer/activation pairing | `core/topology.py` |
100
+ | Moment transfer per node kind | `core/transfer.py` |
101
+ | Canonical shapes and fan arithmetic | `core/fan.py` |
102
+ | Weight samplers | `core/distributions.py` |
103
+ | Solvers | `core/analytic.py`, `core/empirical.py` |
104
+ | Orchestration | `_run.py` |
@@ -0,0 +1,21 @@
1
+ # Reproducing the measurements
2
+
3
+ The figures quoted in the guide, in `../why.md` and in `../reference.md` come from these.
4
+ They are kept out of the test suite because they take minutes rather than seconds.
5
+
6
+ | script | what it measures |
7
+ |---|---|
8
+ | `quadrature_accuracy.py` | Gauss-Legendre against Gauss-Hermite and Monte Carlo on `relu³`, whose moments are known exactly |
9
+ | `gaussian_breakdown.py` | How far a high moment departs from the Gaussian prediction with width, and whether the Gamma correction tracks it |
10
+ | `stability_table.py` | `χ` and `σ*` for every builtin activation and for `relu^p` |
11
+ | `gain_table.py` | `E[a²]` through a 20-layer MLP under the usual gain table and under both AnyInit modes |
12
+ | `draw_drift.py` | How far one draw of a deep ReLU MLP lands from the analytic prediction, with and without `center=True`, against the empirical mode |
13
+ | `depth_sweep.py` | Measured rescaling across depth and width for `relu^p`, showing which degrees survive |
14
+
15
+ Run any of them directly:
16
+
17
+ ```bash
18
+ python docs/experiments/stability_table.py
19
+ ```
20
+
21
+ `depth_sweep.py`, `draw_drift.py` and `gain_table.py` need PyTorch; the others need only NumPy.
@@ -0,0 +1,80 @@
1
+ """Which activations survive depth, measured rather than predicted.
2
+
3
+ Rescales each layer from measurements so that E[a^2] is one after every activation -- the
4
+ best any per-layer scalar can do -- and reports what actually arrives at the end.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import math
10
+
11
+ import torch
12
+ import torch.nn as nn
13
+
14
+ DEPTHS_AND_WIDTHS = ((3, 256), (5, 256), (10, 256), (20, 256), (20, 1024), (20, 4096))
15
+
16
+
17
+ def network(power: int, depth: int, width: int) -> tuple[nn.Module, type[nn.Module]]:
18
+ class Power(nn.Module):
19
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
20
+ return torch.relu(x) ** power
21
+
22
+ layers: list[nn.Module] = []
23
+ for _ in range(depth):
24
+ layers += [nn.Linear(width, width, bias=False), Power()]
25
+ return nn.Sequential(*layers), Power
26
+
27
+
28
+ def rescale(model: nn.Module, activation: type[nn.Module], power: int, rows: int) -> None:
29
+ """Drive E[a^2] to one after every activation, measuring as we go."""
30
+ first = next(m for m in model if isinstance(m, nn.Linear))
31
+ hidden = torch.randn(rows, first.in_features)
32
+ with torch.no_grad():
33
+ for index, layer in enumerate(model):
34
+ if isinstance(layer, nn.Linear):
35
+ follower = model[index + 1]
36
+ for _ in range(40):
37
+ value = float(follower(layer(hidden)).pow(2).mean())
38
+ if not math.isfinite(value) or value <= 0.0:
39
+ layer.weight.mul_(0.5)
40
+ continue
41
+ layer.weight.mul_((1.0 / value) ** (1.0 / (2 * power)))
42
+ if abs(value - 1.0) < 1e-5:
43
+ break
44
+ hidden = layer(hidden)
45
+
46
+
47
+ def main(rows: int = 2048) -> None:
48
+ print(f"{'p':>3s} {'depth':>6s} {'width':>6s} {'final E[a^2]':>13s} {'top-8 energy':>13s}")
49
+ for power in (1, 2, 3):
50
+ for depth, width in DEPTHS_AND_WIDTHS:
51
+ torch.manual_seed(0)
52
+ model, activation = network(power, depth, width)
53
+ for module in model.modules():
54
+ if isinstance(module, nn.Linear):
55
+ nn.init.normal_(module.weight, 0.0, 0.05)
56
+ rescale(model, activation, power, rows)
57
+
58
+ with torch.no_grad():
59
+ output = model(torch.randn(rows, width))
60
+ energy = output.pow(2)
61
+ total = float(energy.sum())
62
+ share = (
63
+ float(energy.sum(0).sort(descending=True).values[:8].sum()) / total
64
+ if total > 0.0
65
+ else float("nan")
66
+ )
67
+ print(
68
+ f"{power:3d} {depth:6d} {width:6d} {float(energy.mean()):13.4g} "
69
+ f"{share * 100:12.1f}%"
70
+ )
71
+
72
+ print(
73
+ "\nDegree one is stable at any depth and widening only spreads the energy further.\n"
74
+ "Above it no amount of per-layer rescaling holds: the signal concentrates into a\n"
75
+ "handful of units and then dies, and extra width does not help."
76
+ )
77
+
78
+
79
+ if __name__ == "__main__":
80
+ main()