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.
- anyinit-0.1.0/.gitignore +14 -0
- anyinit-0.1.0/CHANGELOG.md +30 -0
- anyinit-0.1.0/LICENSE +21 -0
- anyinit-0.1.0/PKG-INFO +161 -0
- anyinit-0.1.0/README.md +129 -0
- anyinit-0.1.0/docs/api.md +44 -0
- anyinit-0.1.0/docs/changelog.md +1 -0
- anyinit-0.1.0/docs/contributing.md +40 -0
- anyinit-0.1.0/docs/design.md +104 -0
- anyinit-0.1.0/docs/experiments/README.md +21 -0
- anyinit-0.1.0/docs/experiments/depth_sweep.py +80 -0
- anyinit-0.1.0/docs/experiments/draw_drift.py +51 -0
- anyinit-0.1.0/docs/experiments/gain_table.py +84 -0
- anyinit-0.1.0/docs/experiments/gaussian_breakdown.py +57 -0
- anyinit-0.1.0/docs/experiments/quadrature_accuracy.py +63 -0
- anyinit-0.1.0/docs/experiments/stability_table.py +86 -0
- anyinit-0.1.0/docs/guide/frameworks.md +39 -0
- anyinit-0.1.0/docs/guide/modes.md +44 -0
- anyinit-0.1.0/docs/guide/options.md +54 -0
- anyinit-0.1.0/docs/guide/report.md +48 -0
- anyinit-0.1.0/docs/guide/stability.md +61 -0
- anyinit-0.1.0/docs/index.md +1 -0
- anyinit-0.1.0/docs/reference.md +51 -0
- anyinit-0.1.0/docs/requirements.txt +4 -0
- anyinit-0.1.0/docs/tutorials/custom-activation.md +31 -0
- anyinit-0.1.0/docs/tutorials/depth.md +14 -0
- anyinit-0.1.0/docs/tutorials/multiple-frameworks.md +17 -0
- anyinit-0.1.0/docs/tutorials/pinn.md +19 -0
- anyinit-0.1.0/docs/tutorials/resnet.md +16 -0
- anyinit-0.1.0/docs/why.md +45 -0
- anyinit-0.1.0/examples/custom_activation.py +69 -0
- anyinit-0.1.0/examples/depth_matters.py +93 -0
- anyinit-0.1.0/examples/multibackend.py +120 -0
- anyinit-0.1.0/examples/pinn_relu2.py +123 -0
- anyinit-0.1.0/examples/resnet_cifar.py +119 -0
- anyinit-0.1.0/pyproject.toml +107 -0
- anyinit-0.1.0/src/anyinit/__init__.py +298 -0
- anyinit-0.1.0/src/anyinit/_run.py +429 -0
- anyinit-0.1.0/src/anyinit/backends/__init__.py +88 -0
- anyinit-0.1.0/src/anyinit/backends/base.py +254 -0
- anyinit-0.1.0/src/anyinit/backends/jax.py +664 -0
- anyinit-0.1.0/src/anyinit/backends/keras.py +674 -0
- anyinit-0.1.0/src/anyinit/backends/pytorch.py +1170 -0
- anyinit-0.1.0/src/anyinit/config.py +86 -0
- anyinit-0.1.0/src/anyinit/core/__init__.py +1 -0
- anyinit-0.1.0/src/anyinit/core/activations.py +172 -0
- anyinit-0.1.0/src/anyinit/core/analytic.py +497 -0
- anyinit-0.1.0/src/anyinit/core/distributions.py +124 -0
- anyinit-0.1.0/src/anyinit/core/empirical.py +177 -0
- anyinit-0.1.0/src/anyinit/core/fan.py +76 -0
- anyinit-0.1.0/src/anyinit/core/graph.py +154 -0
- anyinit-0.1.0/src/anyinit/core/moments.py +72 -0
- anyinit-0.1.0/src/anyinit/core/profile.py +285 -0
- anyinit-0.1.0/src/anyinit/core/quadrature.py +237 -0
- anyinit-0.1.0/src/anyinit/core/registry.py +204 -0
- anyinit-0.1.0/src/anyinit/core/solve.py +145 -0
- anyinit-0.1.0/src/anyinit/core/stability.py +208 -0
- anyinit-0.1.0/src/anyinit/core/topology.py +184 -0
- anyinit-0.1.0/src/anyinit/core/transfer.py +170 -0
- anyinit-0.1.0/src/anyinit/errors.py +27 -0
- anyinit-0.1.0/src/anyinit/py.typed +0 -0
- anyinit-0.1.0/src/anyinit/report.py +348 -0
- anyinit-0.1.0/tests/__init__.py +0 -0
- anyinit-0.1.0/tests/conftest.py +112 -0
- anyinit-0.1.0/tests/helpers.py +76 -0
- anyinit-0.1.0/tests/test_agreement.py +56 -0
- anyinit-0.1.0/tests/test_api.py +125 -0
- anyinit-0.1.0/tests/test_backend_jax.py +225 -0
- anyinit-0.1.0/tests/test_backend_keras.py +153 -0
- anyinit-0.1.0/tests/test_backend_pytorch.py +206 -0
- anyinit-0.1.0/tests/test_core_isolation.py +66 -0
- anyinit-0.1.0/tests/test_cross_backend.py +126 -0
- anyinit-0.1.0/tests/test_custom_activation.py +200 -0
- anyinit-0.1.0/tests/test_distributions.py +148 -0
- anyinit-0.1.0/tests/test_docs_tables.py +31 -0
- anyinit-0.1.0/tests/test_errors.py +71 -0
- anyinit-0.1.0/tests/test_fan.py +61 -0
- anyinit-0.1.0/tests/test_graph.py +53 -0
- anyinit-0.1.0/tests/test_idempotence.py +97 -0
- anyinit-0.1.0/tests/test_input_moments.py +76 -0
- anyinit-0.1.0/tests/test_optional_dependencies.py +75 -0
- anyinit-0.1.0/tests/test_pooling.py +171 -0
- anyinit-0.1.0/tests/test_profile.py +92 -0
- anyinit-0.1.0/tests/test_property.py +93 -0
- anyinit-0.1.0/tests/test_quadrature.py +104 -0
- anyinit-0.1.0/tests/test_regressions.py +452 -0
- anyinit-0.1.0/tests/test_report.py +175 -0
- anyinit-0.1.0/tests/test_stability.py +125 -0
- anyinit-0.1.0/tests/test_topology.py +157 -0
- anyinit-0.1.0/tests/test_transfer.py +128 -0
anyinit-0.1.0/.gitignore
ADDED
|
@@ -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
|
+
[](https://github.com/jmiravet/AnyInit/actions/workflows/ci.yml)
|
|
36
|
+
[](https://codecov.io/gh/jmiravet/AnyInit)
|
|
37
|
+
[](https://pypi.org/project/anyinit/)
|
|
38
|
+
[](https://pypi.org/project/anyinit/)
|
|
39
|
+
[](https://jmiravet.github.io/AnyInit/)
|
|
40
|
+
[](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
|
anyinit-0.1.0/README.md
ADDED
|
@@ -0,0 +1,129 @@
|
|
|
1
|
+
# AnyInit
|
|
2
|
+
|
|
3
|
+
[](https://github.com/jmiravet/AnyInit/actions/workflows/ci.yml)
|
|
4
|
+
[](https://codecov.io/gh/jmiravet/AnyInit)
|
|
5
|
+
[](https://pypi.org/project/anyinit/)
|
|
6
|
+
[](https://pypi.org/project/anyinit/)
|
|
7
|
+
[](https://jmiravet.github.io/AnyInit/)
|
|
8
|
+
[](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()
|