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 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