anyinit 0.1.0__py3-none-any.whl
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/__init__.py +298 -0
- anyinit/_run.py +429 -0
- anyinit/backends/__init__.py +88 -0
- anyinit/backends/base.py +254 -0
- anyinit/backends/jax.py +664 -0
- anyinit/backends/keras.py +674 -0
- anyinit/backends/pytorch.py +1170 -0
- anyinit/config.py +86 -0
- anyinit/core/__init__.py +1 -0
- anyinit/core/activations.py +172 -0
- anyinit/core/analytic.py +497 -0
- anyinit/core/distributions.py +124 -0
- anyinit/core/empirical.py +177 -0
- anyinit/core/fan.py +76 -0
- anyinit/core/graph.py +154 -0
- anyinit/core/moments.py +72 -0
- anyinit/core/profile.py +285 -0
- anyinit/core/quadrature.py +237 -0
- anyinit/core/registry.py +204 -0
- anyinit/core/solve.py +145 -0
- anyinit/core/stability.py +208 -0
- anyinit/core/topology.py +184 -0
- anyinit/core/transfer.py +170 -0
- anyinit/errors.py +27 -0
- anyinit/py.typed +0 -0
- anyinit/report.py +348 -0
- anyinit-0.1.0.dist-info/METADATA +161 -0
- anyinit-0.1.0.dist-info/RECORD +30 -0
- anyinit-0.1.0.dist-info/WHEEL +4 -0
- anyinit-0.1.0.dist-info/licenses/LICENSE +21 -0
anyinit/__init__.py
ADDED
|
@@ -0,0 +1,298 @@
|
|
|
1
|
+
"""AnyInit: initialize any model, in any framework, correctly.
|
|
2
|
+
|
|
3
|
+
One call traces the architecture, works out which activation follows each layer, and
|
|
4
|
+
scales every weight so the signal neither dies nor explodes with depth::
|
|
5
|
+
|
|
6
|
+
import anyinit
|
|
7
|
+
|
|
8
|
+
report = anyinit.initialize(model) # data-free
|
|
9
|
+
report = anyinit.initialize(model, "empirical", batch) # measured
|
|
10
|
+
print(report)
|
|
11
|
+
|
|
12
|
+
``model`` may be a ``torch.nn.Module``, a Keras model or a Flax module; the framework is
|
|
13
|
+
detected from it, and only the one in use needs to be installed.
|
|
14
|
+
|
|
15
|
+
An activation AnyInit has never seen is a first-class input. Register it and its moment
|
|
16
|
+
map is measured by Gaussian quadrature, then used in either mode::
|
|
17
|
+
|
|
18
|
+
@anyinit.register_activation
|
|
19
|
+
def relu3(x):
|
|
20
|
+
return torch.relu(x) ** 3
|
|
21
|
+
|
|
22
|
+
Some activations cannot be stabilized across depth by any initialization, and the report
|
|
23
|
+
says so rather than returning a dead network: ``relu3`` above is homogeneous of degree
|
|
24
|
+
three, so a relative error in the variance triples at every layer. See
|
|
25
|
+
``report.stability``, or call ``report.assert_healthy()`` to raise on the finding.
|
|
26
|
+
"""
|
|
27
|
+
|
|
28
|
+
from __future__ import annotations
|
|
29
|
+
|
|
30
|
+
from collections.abc import Mapping
|
|
31
|
+
from typing import Any
|
|
32
|
+
|
|
33
|
+
import numpy as np
|
|
34
|
+
|
|
35
|
+
from . import backends as _backends
|
|
36
|
+
from .config import InitConfig
|
|
37
|
+
from .core.activations import BUILTIN as BUILTIN_ACTIVATIONS
|
|
38
|
+
from .core.profile import ActivationProfile
|
|
39
|
+
from .core.registry import REGISTRY, ActivationRef
|
|
40
|
+
from .core.stability import depth_error_factor
|
|
41
|
+
from .errors import (
|
|
42
|
+
AnyInitError,
|
|
43
|
+
BackendNotFoundError,
|
|
44
|
+
BackendUnavailableError,
|
|
45
|
+
ConfigError,
|
|
46
|
+
TraceError,
|
|
47
|
+
)
|
|
48
|
+
from .report import InitReport, LayerRecord, StabilityRecord
|
|
49
|
+
|
|
50
|
+
__version__ = "0.1.0"
|
|
51
|
+
|
|
52
|
+
__all__ = [
|
|
53
|
+
"BUILTIN_ACTIVATIONS",
|
|
54
|
+
"ActivationProfile",
|
|
55
|
+
"AnyInitError",
|
|
56
|
+
"BackendNotFoundError",
|
|
57
|
+
"BackendUnavailableError",
|
|
58
|
+
"ConfigError",
|
|
59
|
+
"InitReport",
|
|
60
|
+
"LayerRecord",
|
|
61
|
+
"StabilityRecord",
|
|
62
|
+
"TraceError",
|
|
63
|
+
"__version__",
|
|
64
|
+
"activation_profile",
|
|
65
|
+
"available_backends",
|
|
66
|
+
"depth_error_factor",
|
|
67
|
+
"gain",
|
|
68
|
+
"initialize",
|
|
69
|
+
"initialize_params",
|
|
70
|
+
"known_backends",
|
|
71
|
+
"register_activation",
|
|
72
|
+
"registered_activations",
|
|
73
|
+
"unregister_activation",
|
|
74
|
+
]
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
def initialize(
|
|
78
|
+
model: Any,
|
|
79
|
+
mode: str = "analytic",
|
|
80
|
+
input_spec: Any = None,
|
|
81
|
+
*,
|
|
82
|
+
distribution: str = "normal",
|
|
83
|
+
center: bool = False,
|
|
84
|
+
gains: Mapping[str, float] | None = None,
|
|
85
|
+
seed: int | None = None,
|
|
86
|
+
params: Any = None,
|
|
87
|
+
) -> InitReport:
|
|
88
|
+
"""Reinitialize ``model`` so its activation statistics hold across depth.
|
|
89
|
+
|
|
90
|
+
Args:
|
|
91
|
+
model: A ``torch.nn.Module``, a Keras model or layer, or a Flax module.
|
|
92
|
+
mode: ``"analytic"`` propagates moments through the graph without running the
|
|
93
|
+
model; ``"empirical"`` measures real batches and assumes nothing.
|
|
94
|
+
input_spec: A shape, a batch, or a callable returning a batch. Optional in
|
|
95
|
+
analytic mode, where it is used only for the validation pass; required in
|
|
96
|
+
empirical mode and for Flax.
|
|
97
|
+
distribution: Shape of the weight draw: ``"normal"``, ``"uniform"`` or
|
|
98
|
+
``"sinusoidal"``.
|
|
99
|
+
center: Remove each output unit's weight mean, so that a layer discards the mean
|
|
100
|
+
of its input instead of passing it on.
|
|
101
|
+
gains: Activation name -> fixed gain, e.g. ``{"relu": 2**0.5}``. Layers feeding
|
|
102
|
+
that activation get ``gain / sqrt(fan_in)`` instead of a solved scale.
|
|
103
|
+
seed: Seed for the weight draw. The same seed gives the same weights in every
|
|
104
|
+
framework.
|
|
105
|
+
params: Parameter tree for functional frameworks. The new tree comes back as
|
|
106
|
+
``report.params``.
|
|
107
|
+
|
|
108
|
+
Returns:
|
|
109
|
+
An :class:`~anyinit.report.InitReport`: the scale chosen for every layer, a
|
|
110
|
+
depth-stability verdict per activation, and anything AnyInit could not do.
|
|
111
|
+
|
|
112
|
+
Raises:
|
|
113
|
+
ConfigError: An option is invalid. Raised before the model is touched.
|
|
114
|
+
BackendNotFoundError: The object is not a model of a supported framework.
|
|
115
|
+
|
|
116
|
+
"""
|
|
117
|
+
config = InitConfig.build(
|
|
118
|
+
mode, input_spec, distribution=distribution, center=center, gains=gains, seed=seed
|
|
119
|
+
)
|
|
120
|
+
from ._run import run
|
|
121
|
+
|
|
122
|
+
return run(model, config, params=params)
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
def initialize_params(
|
|
126
|
+
model: Any, params: Any, input_spec: Any = None, mode: str = "analytic", **options: Any
|
|
127
|
+
) -> tuple[Any, InitReport]:
|
|
128
|
+
"""Functional form for JAX and other immutable-parameter frameworks.
|
|
129
|
+
|
|
130
|
+
As :func:`initialize`, but returns the new parameter tree alongside the report::
|
|
131
|
+
|
|
132
|
+
params, report = anyinit.initialize_params(model, params, (1, 32))
|
|
133
|
+
"""
|
|
134
|
+
report = initialize(model, mode, input_spec, params=params, **options)
|
|
135
|
+
return report.params, report
|
|
136
|
+
|
|
137
|
+
|
|
138
|
+
def register_activation(
|
|
139
|
+
activation: Any = None, *, name: str | None = None, overwrite: bool = False
|
|
140
|
+
) -> Any:
|
|
141
|
+
"""Teach AnyInit an activation it does not know.
|
|
142
|
+
|
|
143
|
+
Usable bare, with arguments, or as a plain call::
|
|
144
|
+
|
|
145
|
+
@anyinit.register_activation
|
|
146
|
+
def relu3(x):
|
|
147
|
+
return torch.relu(x) ** 3
|
|
148
|
+
|
|
149
|
+
@anyinit.register_activation(name="relu_cubed")
|
|
150
|
+
class ReLU3(torch.nn.Module):
|
|
151
|
+
def forward(self, x):
|
|
152
|
+
return torch.relu(x) ** 3
|
|
153
|
+
|
|
154
|
+
anyinit.register_activation(lambda x: np.maximum(x, 0) ** 3, name="relu3")
|
|
155
|
+
|
|
156
|
+
The moment map is computed by Gaussian quadrature rather than looked up, and both
|
|
157
|
+
modes then use it. Registering also makes the activation visible to the tracers, so a
|
|
158
|
+
module subclass becomes a single graph node instead of being inlined into primitives.
|
|
159
|
+
|
|
160
|
+
A function written against NumPy works with every backend and needs no framework; one
|
|
161
|
+
written against a specific framework is evaluated through it. Which it is gets
|
|
162
|
+
detected.
|
|
163
|
+
|
|
164
|
+
Args:
|
|
165
|
+
activation: A callable, or an ``nn.Module``/Keras layer subclass.
|
|
166
|
+
name: Registry name. Defaults to the function or class name, lowercased.
|
|
167
|
+
overwrite: Permit replacing an existing registration of the same name.
|
|
168
|
+
|
|
169
|
+
Returns:
|
|
170
|
+
The activation itself, so this works as a decorator.
|
|
171
|
+
|
|
172
|
+
"""
|
|
173
|
+
if activation is None:
|
|
174
|
+
return lambda target: register_activation(target, name=name, overwrite=overwrite)
|
|
175
|
+
|
|
176
|
+
callable_obj, native_type = _as_callable(activation)
|
|
177
|
+
use_numpy = _is_numpy_callable(callable_obj)
|
|
178
|
+
REGISTRY.register(
|
|
179
|
+
name or _default_name(activation),
|
|
180
|
+
numpy_fn=callable_obj if use_numpy else None,
|
|
181
|
+
native_fn=None if use_numpy else callable_obj,
|
|
182
|
+
native_type=native_type,
|
|
183
|
+
overwrite=overwrite,
|
|
184
|
+
)
|
|
185
|
+
return activation
|
|
186
|
+
|
|
187
|
+
|
|
188
|
+
def unregister_activation(name: str) -> None:
|
|
189
|
+
"""Remove a registration and discard its cached profile."""
|
|
190
|
+
REGISTRY.unregister(name)
|
|
191
|
+
|
|
192
|
+
|
|
193
|
+
def registered_activations() -> tuple[str, ...]:
|
|
194
|
+
"""Names of user-registered activations."""
|
|
195
|
+
return REGISTRY.custom_names
|
|
196
|
+
|
|
197
|
+
|
|
198
|
+
def activation_profile(name: str, backend: Any = None, **params: float) -> ActivationProfile:
|
|
199
|
+
"""Profile of a registered or builtin activation, for inspection.
|
|
200
|
+
|
|
201
|
+
``activation_profile("relu3").chi`` says whether an activation can hold a signal across
|
|
202
|
+
depth before anything is built with it. ``backend`` is needed only for natively
|
|
203
|
+
registered code, and is worked out from the registration when omitted.
|
|
204
|
+
"""
|
|
205
|
+
reference = ActivationRef.of(name, **params)
|
|
206
|
+
if backend is None:
|
|
207
|
+
backend = _backend_for(name)
|
|
208
|
+
profile = REGISTRY.profile(reference, backend)
|
|
209
|
+
if profile is None:
|
|
210
|
+
known = ", ".join(sorted(set(BUILTIN_ACTIVATIONS) | set(REGISTRY.custom_names)))
|
|
211
|
+
raise ConfigError(f"no profile for activation {name!r}; known activations: {known}")
|
|
212
|
+
return profile
|
|
213
|
+
|
|
214
|
+
|
|
215
|
+
def gain(name: str, **params: float) -> float:
|
|
216
|
+
"""Gain AnyInit gives an activation: weights scaled ``gain / sqrt(fan_in)``.
|
|
217
|
+
|
|
218
|
+
The input standard deviation that lands ``E[f(z)^2]`` on one, so ``gain("relu")`` is
|
|
219
|
+
``sqrt(2)``. A bounded activation cannot reach one; its gain is the input scale at the
|
|
220
|
+
middle of its reachable variance range. Parameters such as ``negative_slope`` go in
|
|
221
|
+
as keywords. Pass the result, or any other value, to ``initialize(gains=...)`` to fix
|
|
222
|
+
it instead of solving for it.
|
|
223
|
+
"""
|
|
224
|
+
return activation_profile(name, **params).gain
|
|
225
|
+
|
|
226
|
+
|
|
227
|
+
def available_backends() -> tuple[str, ...]:
|
|
228
|
+
"""Names of the backends whose framework is installed here.
|
|
229
|
+
|
|
230
|
+
Checked without importing any framework.
|
|
231
|
+
"""
|
|
232
|
+
return tuple(cls.name for cls in _backends.installed())
|
|
233
|
+
|
|
234
|
+
|
|
235
|
+
def known_backends() -> tuple[str, ...]:
|
|
236
|
+
"""Names of every backend AnyInit ships, installed or not."""
|
|
237
|
+
return tuple(cls.name for cls in _backends.known())
|
|
238
|
+
|
|
239
|
+
|
|
240
|
+
# ------------------------------------------------------------------- helpers
|
|
241
|
+
|
|
242
|
+
|
|
243
|
+
def _backend_for(name: str) -> Any:
|
|
244
|
+
"""Backend able to evaluate a natively-registered activation, if one is needed."""
|
|
245
|
+
spec = REGISTRY.spec(name)
|
|
246
|
+
if spec is None or not spec.needs_backend:
|
|
247
|
+
return None
|
|
248
|
+
roots = _backends.module_roots(spec.native_type or spec.native_fn)
|
|
249
|
+
for cls in _backends.installed():
|
|
250
|
+
if cls.frameworks and roots & set(cls.frameworks):
|
|
251
|
+
return _instantiate(cls)
|
|
252
|
+
return None
|
|
253
|
+
|
|
254
|
+
|
|
255
|
+
def _instantiate(cls: Any) -> Any:
|
|
256
|
+
"""Build a backend, returning ``None`` if its framework will not import."""
|
|
257
|
+
try:
|
|
258
|
+
return cls()
|
|
259
|
+
except Exception:
|
|
260
|
+
return None
|
|
261
|
+
|
|
262
|
+
|
|
263
|
+
def _default_name(activation: Any) -> str:
|
|
264
|
+
if isinstance(activation, type):
|
|
265
|
+
return activation.__name__.lower()
|
|
266
|
+
explicit = getattr(activation, "__name__", None)
|
|
267
|
+
if explicit and explicit != "<lambda>":
|
|
268
|
+
return str(explicit)
|
|
269
|
+
return type(activation).__name__.lower()
|
|
270
|
+
|
|
271
|
+
|
|
272
|
+
def _as_callable(activation: Any) -> tuple[Any, Any]:
|
|
273
|
+
"""Return ``(callable, native_type)``.
|
|
274
|
+
|
|
275
|
+
A class is instantiated so it can be profiled, and kept so the tracers recognize its
|
|
276
|
+
instances inside a model.
|
|
277
|
+
"""
|
|
278
|
+
if isinstance(activation, type):
|
|
279
|
+
return activation(), activation
|
|
280
|
+
if not callable(activation):
|
|
281
|
+
raise ConfigError(f"activation must be callable, got {type(activation).__name__}")
|
|
282
|
+
native_type = type(activation) if _looks_like_layer(activation) else None
|
|
283
|
+
return activation, native_type
|
|
284
|
+
|
|
285
|
+
|
|
286
|
+
def _looks_like_layer(activation: Any) -> bool:
|
|
287
|
+
roots = _backends.module_roots(activation)
|
|
288
|
+
return bool(roots & {"torch", "keras", "tensorflow", "flax"})
|
|
289
|
+
|
|
290
|
+
|
|
291
|
+
def _is_numpy_callable(fn: Any) -> bool:
|
|
292
|
+
"""Whether ``fn`` can be applied directly to a NumPy array."""
|
|
293
|
+
probe = np.array([-1.0, 0.0, 1.0])
|
|
294
|
+
try:
|
|
295
|
+
out = fn(probe)
|
|
296
|
+
except Exception:
|
|
297
|
+
return False
|
|
298
|
+
return isinstance(out, np.ndarray) and out.shape == probe.shape
|