yaxlib 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.
- yax/__init__.py +28 -0
- yax/core.py +346 -0
- yax/image/__init__.py +0 -0
- yax/image/augmentation.py +123 -0
- yax/layers/Conv_layer.py +70 -0
- yax/layers/Dropout.py +63 -0
- yax/layers/Embedding.py +29 -0
- yax/layers/LayerNorm.py +42 -0
- yax/layers/Linear.py +20 -0
- yax/layers/MLP.py +62 -0
- yax/layers/MessagePassing_layer.py +80 -0
- yax/layers/MultiHeadAttention.py +127 -0
- yax/layers/RNN_layer.py +156 -0
- yax/layers/TransformerBlock.py +78 -0
- yax/layers/__init__.py +0 -0
- yax/layers/positional_encoding.py +37 -0
- yax/models/MiniYOLO.py +200 -0
- yax/models/UNet.py +107 -0
- yax/models/__init__.py +0 -0
- yax/training/History.py +52 -0
- yax/training/Trainer.py +240 -0
- yax/training/__init__.py +0 -0
- yax/training/configs.py +20 -0
- yax/training/losses.py +77 -0
- yaxlib-0.1.0.dist-info/METADATA +72 -0
- yaxlib-0.1.0.dist-info/RECORD +27 -0
- yaxlib-0.1.0.dist-info/WHEEL +4 -0
yax/__init__.py
ADDED
|
@@ -0,0 +1,28 @@
|
|
|
1
|
+
"""yax : un mini-framework de réseaux de neurones pour jax, à but pédagogique.
|
|
2
|
+
|
|
3
|
+
Trois principes :
|
|
4
|
+
- un modèle est un pytree (`yax.Module`) dont les feuilles sont exactement les
|
|
5
|
+
paramètres : `jax.grad`, `jax.jit` et `optimizer.init(model)` l'acceptent
|
|
6
|
+
tel quel, sans machinerie de filtrage ;
|
|
7
|
+
- signature uniforme `apply(x, rkey=None)` : `rkey` est la source d'aléatoire
|
|
8
|
+
(dropout, échantillonnage), jamais un mode. Le mode évaluation se bascule
|
|
9
|
+
par `model = model.set_inference(True)` — le Trainer s'en charge ;
|
|
10
|
+
- tout ce qui n'est pas un paramètre est statique (`yax.StaticField`), y
|
|
11
|
+
compris les constantes-tableaux.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
from yax.core import Module, StaticField, tree_at, tree_pprint
|
|
15
|
+
|
|
16
|
+
from yax.layers.Linear import Linear
|
|
17
|
+
from yax.layers.MLP import MLP
|
|
18
|
+
from yax.layers.Dropout import Dropout
|
|
19
|
+
from yax.layers.LayerNorm import LayerNorm
|
|
20
|
+
from yax.layers.Embedding import Embedding
|
|
21
|
+
from yax.layers.Conv_layer import Conv_layer
|
|
22
|
+
from yax.layers.RNN_layer import RNN_layer, GRUCell, LSTMCell
|
|
23
|
+
from yax.layers.MultiHeadAttention import MultiHeadAttention, causal_mask, padding_mask
|
|
24
|
+
from yax.layers.TransformerBlock import TransformerBlock
|
|
25
|
+
from yax.layers.MessagePassing_layer import MessagePassing_layer
|
|
26
|
+
from yax.layers.positional_encoding import sinusoidal_positional_encoding
|
|
27
|
+
|
|
28
|
+
__version__ = "0.1.0"
|
yax/core.py
ADDED
|
@@ -0,0 +1,346 @@
|
|
|
1
|
+
"""Le coeur de yax : la classe Module, sans dépendance à equinox.
|
|
2
|
+
|
|
3
|
+
Un yax.Module est un pytree jax. Ses champs se déclarent par annotations de
|
|
4
|
+
classe et se rangent en deux familles :
|
|
5
|
+
|
|
6
|
+
- les champs **dynamiques** (annotation seule) : les feuilles du pytree — les
|
|
7
|
+
paramètres. Ils ne peuvent contenir QUE des tableaux jax, des sous-modules,
|
|
8
|
+
ou des conteneurs (list/tuple/dict) de ceux-ci. Tout le reste est refusé à
|
|
9
|
+
la construction, avec un message explicite : c'est la discipline qui permet
|
|
10
|
+
`jax.grad(loss)(model)`, `jax.jit` et `optimizer.init(model)` sans aucune
|
|
11
|
+
machinerie de filtrage ;
|
|
12
|
+
- les champs **statiques** (`= StaticField(default_value=...)`) : des
|
|
13
|
+
métadonnées rangées dans la STRUCTURE du pytree, invisibles pour `grad`,
|
|
14
|
+
`tree.map` et les optimiseurs. Ils font partie de l'identité du modèle
|
|
15
|
+
(changer un statique = changer de structure = recompilation). Un statique
|
|
16
|
+
peut être un tableau jax : il devient alors une CONSTANTE du modèle
|
|
17
|
+
(encodage positionnel, grille figée...), cuite dans le code compilé.
|
|
18
|
+
|
|
19
|
+
Les modules sont immuables : on « modifie » avec yax.tree_at (reconstruction
|
|
20
|
+
d'une copie). Chaque module porte un drapeau statique `inference` (False par
|
|
21
|
+
défaut) ; `model.set_inference(True)` rend une copie où le drapeau est basculé
|
|
22
|
+
récursivement dans tous les sous-modules. Convention d'usage : le mode se
|
|
23
|
+
change rarement (le Trainer s'en occupe), l'aléatoire passe par l'argument
|
|
24
|
+
`rkey` des `apply`.
|
|
25
|
+
"""
|
|
26
|
+
|
|
27
|
+
import functools
|
|
28
|
+
|
|
29
|
+
import jax
|
|
30
|
+
import jax.tree_util as jtu
|
|
31
|
+
import numpy as np
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
class StaticField:
|
|
35
|
+
"""Marqueur de champ statique, avec valeur par défaut optionnelle.
|
|
36
|
+
|
|
37
|
+
S'utilise comme valeur de classe : `layer_sizes: tuple = StaticField()`.
|
|
38
|
+
"""
|
|
39
|
+
|
|
40
|
+
def __init__(self, default_value=None):
|
|
41
|
+
self.default_value = default_value
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
class _HashableArray:
|
|
45
|
+
"""Enveloppe hachable pour un tableau rangé dans un champ statique.
|
|
46
|
+
|
|
47
|
+
La structure d'un pytree doit être hachable (le cache de jit la hache à
|
|
48
|
+
chaque appel). Les octets et le hash sont précalculés une fois — le
|
|
49
|
+
tableau est immuable. À réserver aux petites constantes : l'égalité est
|
|
50
|
+
comparée à chaque dispatch jitté.
|
|
51
|
+
"""
|
|
52
|
+
|
|
53
|
+
__slots__ = ("value", "_bytes", "_hash")
|
|
54
|
+
|
|
55
|
+
def __init__(self, value):
|
|
56
|
+
self.value = value
|
|
57
|
+
arr = np.asarray(value)
|
|
58
|
+
self._bytes = (arr.shape, str(arr.dtype), arr.tobytes())
|
|
59
|
+
self._hash = hash(self._bytes)
|
|
60
|
+
|
|
61
|
+
def __eq__(self, other):
|
|
62
|
+
return type(other) is _HashableArray and self._bytes == other._bytes
|
|
63
|
+
|
|
64
|
+
def __hash__(self):
|
|
65
|
+
return self._hash
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def _est_tableau(x):
|
|
69
|
+
return isinstance(x, jax.Array)
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def _valide_dynamique(cls_name, name, value):
|
|
73
|
+
"""Un champ dynamique : tableau jax, Module, ou conteneur de ceux-ci."""
|
|
74
|
+
if _est_tableau(value) or isinstance(value, Module):
|
|
75
|
+
return
|
|
76
|
+
if isinstance(value, (list, tuple)):
|
|
77
|
+
for v in value:
|
|
78
|
+
_valide_dynamique(cls_name, name, v)
|
|
79
|
+
return
|
|
80
|
+
if isinstance(value, dict):
|
|
81
|
+
for v in value.values():
|
|
82
|
+
_valide_dynamique(cls_name, name, v)
|
|
83
|
+
return
|
|
84
|
+
if isinstance(value, np.ndarray):
|
|
85
|
+
raise TypeError(
|
|
86
|
+
f"{cls_name}.{name} : tableau numpy dans un champ dynamique — "
|
|
87
|
+
f"convertissez avec jnp.asarray(...), ou déclarez le champ "
|
|
88
|
+
f"StaticField si c'est une constante.")
|
|
89
|
+
raise TypeError(
|
|
90
|
+
f"{cls_name}.{name} : un champ dynamique ne peut contenir que des "
|
|
91
|
+
f"tableaux jax, des sous-modules, ou des list/tuple/dict de ceux-ci — "
|
|
92
|
+
f"reçu {type(value).__name__} ({value!r:.60}). Si ce n'est pas un "
|
|
93
|
+
f"paramètre, déclarez-le : {name}: ... = StaticField().")
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
def _valide_statique(cls_name, name, value):
|
|
97
|
+
"""Un champ statique : hachable, ou tableau (enveloppé au vol)."""
|
|
98
|
+
if _est_tableau(value) or isinstance(value, np.ndarray):
|
|
99
|
+
return
|
|
100
|
+
if isinstance(value, list):
|
|
101
|
+
raise TypeError(
|
|
102
|
+
f"{cls_name}.{name} : une liste dans un champ statique — un "
|
|
103
|
+
f"statique fait partie de l'identité du modèle et doit être "
|
|
104
|
+
f"immuable : utilisez un tuple.")
|
|
105
|
+
try:
|
|
106
|
+
hash(value)
|
|
107
|
+
except TypeError:
|
|
108
|
+
raise TypeError(
|
|
109
|
+
f"{cls_name}.{name} : un champ statique doit être hachable "
|
|
110
|
+
f"(ou être un tableau) — reçu {type(value).__name__}.")
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
class Module:
|
|
114
|
+
"""Classe de base : voir la docstring du module."""
|
|
115
|
+
|
|
116
|
+
inference: bool = StaticField(default_value=False)
|
|
117
|
+
|
|
118
|
+
_initializing = False # défaut de classe : les instances unflatten le gardent
|
|
119
|
+
|
|
120
|
+
# ------------------------------------------------------------- champs
|
|
121
|
+
@classmethod
|
|
122
|
+
def _collecte_champs(cls):
|
|
123
|
+
statics, dynamics = {}, []
|
|
124
|
+
for klass in reversed(cls.__mro__):
|
|
125
|
+
for name in getattr(klass, "__annotations__", {}):
|
|
126
|
+
default = klass.__dict__.get(name, None)
|
|
127
|
+
if isinstance(default, StaticField):
|
|
128
|
+
statics[name] = default.default_value
|
|
129
|
+
if name in dynamics:
|
|
130
|
+
dynamics.remove(name)
|
|
131
|
+
elif name not in statics and name not in dynamics:
|
|
132
|
+
if default is not None:
|
|
133
|
+
raise TypeError(
|
|
134
|
+
f"{cls.__name__}.{name} : valeur par défaut sur un "
|
|
135
|
+
f"champ dynamique — non supporté (les dynamiques "
|
|
136
|
+
f"s'initialisent dans __init__ ; pour une "
|
|
137
|
+
f"métadonnée, utilisez StaticField(default_value=...)).")
|
|
138
|
+
dynamics.append(name)
|
|
139
|
+
cls._yax_statics = statics # nom -> valeur par défaut
|
|
140
|
+
cls._yax_dynamics = tuple(dynamics) # ordre de déclaration
|
|
141
|
+
|
|
142
|
+
def __init_subclass__(cls, **kwargs):
|
|
143
|
+
super().__init_subclass__(**kwargs)
|
|
144
|
+
cls._collecte_champs()
|
|
145
|
+
|
|
146
|
+
init = cls.__init__
|
|
147
|
+
if not getattr(init, "_yax_wrapped", False):
|
|
148
|
+
@functools.wraps(init)
|
|
149
|
+
def init_encadre(self, *args, **kw):
|
|
150
|
+
object.__setattr__(self, "_initializing", True)
|
|
151
|
+
init(self, *args, **kw)
|
|
152
|
+
object.__setattr__(self, "_initializing", False)
|
|
153
|
+
self._fin_de_construction()
|
|
154
|
+
|
|
155
|
+
init_encadre._yax_wrapped = True
|
|
156
|
+
cls.__init__ = init_encadre
|
|
157
|
+
|
|
158
|
+
jtu.register_pytree_with_keys(
|
|
159
|
+
cls,
|
|
160
|
+
flatten_with_keys=cls._flatten_with_keys,
|
|
161
|
+
flatten_func=cls._flatten,
|
|
162
|
+
unflatten_func=cls._unflatten)
|
|
163
|
+
|
|
164
|
+
def _fin_de_construction(self):
|
|
165
|
+
cls = type(self)
|
|
166
|
+
for name, default in cls._yax_statics.items():
|
|
167
|
+
if name not in self.__dict__:
|
|
168
|
+
object.__setattr__(self, name, default)
|
|
169
|
+
for name in cls._yax_dynamics:
|
|
170
|
+
if name not in self.__dict__:
|
|
171
|
+
raise AttributeError(
|
|
172
|
+
f"{cls.__name__} : le champ dynamique '{name}' n'a pas été "
|
|
173
|
+
f"initialisé dans __init__.")
|
|
174
|
+
|
|
175
|
+
# -------------------------------------------------------- immutabilité
|
|
176
|
+
def __setattr__(self, name, value):
|
|
177
|
+
cls = type(self)
|
|
178
|
+
if not self._initializing:
|
|
179
|
+
raise AttributeError(
|
|
180
|
+
f"{cls.__name__} est immuable : impossible d'assigner "
|
|
181
|
+
f"'{name}'. Pour « modifier » un module, reconstruire une "
|
|
182
|
+
f"copie avec yax.tree_at.")
|
|
183
|
+
if name in cls._yax_statics:
|
|
184
|
+
_valide_statique(cls.__name__, name, value)
|
|
185
|
+
elif name in cls._yax_dynamics:
|
|
186
|
+
_valide_dynamique(cls.__name__, name, value)
|
|
187
|
+
else:
|
|
188
|
+
raise AttributeError(
|
|
189
|
+
f"{cls.__name__}.{name} : champ non déclaré — ajoutez "
|
|
190
|
+
f"l'annotation de classe correspondante.")
|
|
191
|
+
object.__setattr__(self, name, value)
|
|
192
|
+
|
|
193
|
+
# ------------------------------------------------------------- pytree
|
|
194
|
+
def _aux(self):
|
|
195
|
+
cls = type(self)
|
|
196
|
+
out = []
|
|
197
|
+
for name in cls._yax_statics:
|
|
198
|
+
v = self.__dict__.get(name)
|
|
199
|
+
if _est_tableau(v) or isinstance(v, np.ndarray):
|
|
200
|
+
v = self.__dict__.setdefault("_wrapped_" + name, _HashableArray(v))
|
|
201
|
+
out.append((name, v))
|
|
202
|
+
return tuple(out)
|
|
203
|
+
|
|
204
|
+
def _flatten(self):
|
|
205
|
+
cls = type(self)
|
|
206
|
+
children = tuple(self.__dict__.get(n) for n in cls._yax_dynamics)
|
|
207
|
+
return children, self._aux()
|
|
208
|
+
|
|
209
|
+
def _flatten_with_keys(self):
|
|
210
|
+
cls = type(self)
|
|
211
|
+
children = tuple((jtu.GetAttrKey(n), self.__dict__.get(n))
|
|
212
|
+
for n in cls._yax_dynamics)
|
|
213
|
+
return children, self._aux()
|
|
214
|
+
|
|
215
|
+
@classmethod
|
|
216
|
+
def _unflatten(cls, aux, children):
|
|
217
|
+
# aucun __init__, aucune validation : jax passe ici avec des tracers,
|
|
218
|
+
# des gradients, ou les sorties de n'importe quel tree.map
|
|
219
|
+
obj = object.__new__(cls)
|
|
220
|
+
for name, value in zip(cls._yax_dynamics, children):
|
|
221
|
+
object.__setattr__(obj, name, value)
|
|
222
|
+
for name, value in aux:
|
|
223
|
+
if type(value) is _HashableArray:
|
|
224
|
+
object.__setattr__(obj, "_wrapped_" + name, value)
|
|
225
|
+
value = value.value
|
|
226
|
+
object.__setattr__(obj, name, value)
|
|
227
|
+
return obj
|
|
228
|
+
|
|
229
|
+
# ------------------------------------------------------------ services
|
|
230
|
+
def set_inference(self, value: bool):
|
|
231
|
+
"""Copie du module où le drapeau statique `inference` vaut `value`,
|
|
232
|
+
récursivement dans tous les sous-modules. L'original est intact."""
|
|
233
|
+
|
|
234
|
+
def reconstruit(x):
|
|
235
|
+
if isinstance(x, Module):
|
|
236
|
+
klass = type(x)
|
|
237
|
+
obj = object.__new__(klass)
|
|
238
|
+
for name in klass._yax_dynamics:
|
|
239
|
+
object.__setattr__(obj, name, reconstruit(x.__dict__.get(name)))
|
|
240
|
+
for name in klass._yax_statics:
|
|
241
|
+
object.__setattr__(obj, name, x.__dict__.get(name))
|
|
242
|
+
object.__setattr__(obj, "inference", bool(value))
|
|
243
|
+
return obj
|
|
244
|
+
if isinstance(x, list):
|
|
245
|
+
return [reconstruit(v) for v in x]
|
|
246
|
+
if isinstance(x, tuple):
|
|
247
|
+
return tuple(reconstruit(v) for v in x)
|
|
248
|
+
if isinstance(x, dict):
|
|
249
|
+
return {k: reconstruit(v) for k, v in x.items()}
|
|
250
|
+
return x
|
|
251
|
+
|
|
252
|
+
return reconstruit(self)
|
|
253
|
+
|
|
254
|
+
# -------------------------------------------------------------- repr
|
|
255
|
+
def __repr__(self):
|
|
256
|
+
return _pformat(self, indent=0)
|
|
257
|
+
|
|
258
|
+
|
|
259
|
+
def _pformat(x, indent):
|
|
260
|
+
pad = " " * indent
|
|
261
|
+
if isinstance(x, Module):
|
|
262
|
+
cls = type(x)
|
|
263
|
+
noms = list(cls._yax_dynamics) + list(cls._yax_statics)
|
|
264
|
+
lignes = [f"{pad}{cls.__name__}("]
|
|
265
|
+
for name in noms:
|
|
266
|
+
valeur = _pformat(x.__dict__.get(name), indent + 1).lstrip()
|
|
267
|
+
lignes.append(f"{pad} {name}={valeur},")
|
|
268
|
+
lignes.append(f"{pad})")
|
|
269
|
+
return "\n".join(lignes)
|
|
270
|
+
if _est_tableau(x) or isinstance(x, np.ndarray):
|
|
271
|
+
dtype = str(x.dtype).replace("float", "f").replace("int", "i").replace("bool", "b")
|
|
272
|
+
return pad + f"{dtype}[{','.join(map(str, x.shape))}]"
|
|
273
|
+
if isinstance(x, (list, tuple)):
|
|
274
|
+
if not x:
|
|
275
|
+
return pad + repr(x)
|
|
276
|
+
ouvre, ferme = ("[", "]") if isinstance(x, list) else ("(", ")")
|
|
277
|
+
interieur = ",\n".join(_pformat(v, indent + 1) for v in x)
|
|
278
|
+
return f"{pad}{ouvre}\n{interieur}\n{pad}{ferme}"
|
|
279
|
+
if isinstance(x, dict):
|
|
280
|
+
if not x:
|
|
281
|
+
return pad + "{}"
|
|
282
|
+
lignes = [pad + "{"]
|
|
283
|
+
for k, v in x.items():
|
|
284
|
+
lignes.append(f"{pad} {k!r}: {_pformat(v, indent + 1).lstrip()},")
|
|
285
|
+
lignes.append(pad + "}")
|
|
286
|
+
return "\n".join(lignes)
|
|
287
|
+
return pad + repr(x)
|
|
288
|
+
|
|
289
|
+
|
|
290
|
+
def tree_pprint(x):
|
|
291
|
+
print(_pformat(x, 0))
|
|
292
|
+
|
|
293
|
+
|
|
294
|
+
def tree_at(where, pytree, replace):
|
|
295
|
+
"""Reconstruit `pytree` en remplaçant une ou plusieurs feuilles.
|
|
296
|
+
|
|
297
|
+
`where` : fonction qui, appliquée au pytree, rend LA feuille à remplacer
|
|
298
|
+
(ou un tuple/liste de feuilles). `replace` : la ou les valeurs neuves.
|
|
299
|
+
|
|
300
|
+
model2 = tree_at(lambda m: m.layers[-1].bias, model, jnp.ones(3))
|
|
301
|
+
|
|
302
|
+
Les feuilles sont repérées par identité (`is`) : `where` doit rendre des
|
|
303
|
+
objets extraits du pytree lui-même. L'original est intact.
|
|
304
|
+
"""
|
|
305
|
+
cibles = where(pytree)
|
|
306
|
+
if not isinstance(cibles, (tuple, list)):
|
|
307
|
+
cibles, replace = (cibles,), (replace,)
|
|
308
|
+
if len(cibles) != len(replace):
|
|
309
|
+
raise ValueError(f"tree_at : {len(cibles)} cible(s) mais "
|
|
310
|
+
f"{len(replace)} remplacement(s).")
|
|
311
|
+
|
|
312
|
+
leaves, treedef = jtu.tree_flatten(pytree)
|
|
313
|
+
nouvelles = list(leaves)
|
|
314
|
+
for cible, valeur in zip(cibles, replace):
|
|
315
|
+
indices = [i for i, l in enumerate(leaves) if l is cible]
|
|
316
|
+
if len(indices) != 1:
|
|
317
|
+
raise ValueError(
|
|
318
|
+
"tree_at : cible introuvable ou ambigue — `where` doit rendre "
|
|
319
|
+
"une feuille (un tableau) extraite du pytree lui-même.")
|
|
320
|
+
nouvelles[indices[0]] = valeur
|
|
321
|
+
return jtu.tree_unflatten(treedef, nouvelles)
|
|
322
|
+
|
|
323
|
+
|
|
324
|
+
if __name__ == "__main__":
|
|
325
|
+
import jax.numpy as jnp
|
|
326
|
+
import jax.random as jr
|
|
327
|
+
|
|
328
|
+
class Petit(Module):
|
|
329
|
+
w: jnp.ndarray
|
|
330
|
+
nom: str = StaticField(default_value="petit")
|
|
331
|
+
|
|
332
|
+
def __init__(self, rkey):
|
|
333
|
+
self.w = jr.normal(rkey, (3,))
|
|
334
|
+
|
|
335
|
+
def apply(self, x, rkey=None):
|
|
336
|
+
return jnp.sum(self.w * x)
|
|
337
|
+
|
|
338
|
+
m = Petit(jr.key(0))
|
|
339
|
+
print(m)
|
|
340
|
+
print("feuilles :", jax.tree.leaves(m))
|
|
341
|
+
g = jax.grad(lambda mm: mm.apply(jnp.ones(3)))(m)
|
|
342
|
+
print("gradient :", type(g).__name__, g.w, "| nom conservé :", g.nom)
|
|
343
|
+
m2 = m.set_inference(True)
|
|
344
|
+
print("inference :", m.inference, "->", m2.inference)
|
|
345
|
+
m3 = tree_at(lambda mm: mm.w, m, jnp.zeros(3))
|
|
346
|
+
print("tree_at :", m3.w, "| original intact :", m.w)
|
yax/image/__init__.py
ADDED
|
File without changes
|
|
@@ -0,0 +1,123 @@
|
|
|
1
|
+
"""Enrichissement des données (augmentation), chapitre 11.2.
|
|
2
|
+
|
|
3
|
+
Des transformations qui PRÉSERVENT l'étiquette, écrites pour UNE image
|
|
4
|
+
channels-first (canaux, H, W). Tout est différentiable et vmap-able :
|
|
5
|
+
l'augmentation se fait sur le device, dans le pas d'entraînement — pas dans un
|
|
6
|
+
chargeur de données à part. L'aléatoire vient d'une clé, comme partout : une
|
|
7
|
+
transformation différente à chaque époque pour un même exemple.
|
|
8
|
+
|
|
9
|
+
Convention : f(rkey, img, ...) — la clé d'abord, comme dans jax.random. Ce sont
|
|
10
|
+
des fonctions, pas des modules : elles n'ont aucun paramètre appris.
|
|
11
|
+
|
|
12
|
+
Piège du cours, à répéter : choisir des transformations qui changent
|
|
13
|
+
l'étiquette. Une symétrie horizontale sur des chiffres manuscrits transforme
|
|
14
|
+
un 2 en quelque chose qui n'est plus un 2. L'invariance qu'on impose doit
|
|
15
|
+
être vraie.
|
|
16
|
+
"""
|
|
17
|
+
|
|
18
|
+
import jax
|
|
19
|
+
import jax.numpy as jnp
|
|
20
|
+
import jax.random as jr
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def random_horizontal_flip(rkey, img, prob=0.5):
|
|
24
|
+
"""Symétrie gauche-droite avec probabilité prob."""
|
|
25
|
+
# jnp.where et non un `if` : le tirage est une valeur tracée (chapitre 5)
|
|
26
|
+
flip = jr.bernoulli(rkey, prob)
|
|
27
|
+
return jnp.where(flip, img[:, :, ::-1], img)
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def random_brightness_contrast(rkey, img, max_shift=0.2, max_factor=0.2):
|
|
31
|
+
"""Luminosité : décalage global ; contraste : dilatation autour de la
|
|
32
|
+
moyenne de l'image."""
|
|
33
|
+
k_shift, k_factor = jr.split(rkey)
|
|
34
|
+
shift = jr.uniform(k_shift, minval=-max_shift, maxval=max_shift)
|
|
35
|
+
factor = 1.0 + jr.uniform(k_factor, minval=-max_factor, maxval=max_factor)
|
|
36
|
+
mean = jnp.mean(img)
|
|
37
|
+
return (img - mean) * factor + mean + shift
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def random_noise(rkey, img, sigma=0.05):
|
|
41
|
+
return img + sigma * jr.normal(rkey, img.shape)
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def random_scale_translate(rkey, img, max_zoom=0.2, max_shift=2.0):
|
|
45
|
+
"""Homothétie (zoom autour du centre) et translation, par
|
|
46
|
+
jax.image.scale_and_translate. max_shift est en pixels."""
|
|
47
|
+
C, H, W = img.shape
|
|
48
|
+
k_zoom, k_shift = jr.split(rkey)
|
|
49
|
+
zoom = 1.0 + jr.uniform(k_zoom, minval=-max_zoom, maxval=max_zoom)
|
|
50
|
+
shift = jr.uniform(k_shift, (2,), minval=-max_shift, maxval=max_shift)
|
|
51
|
+
# scale_and_translate applique sortie(y) = entree((y - translation)/scale) :
|
|
52
|
+
# cette translation-ci recentre le zoom sur le milieu de l'image
|
|
53
|
+
scale = jnp.array([zoom, zoom])
|
|
54
|
+
translation = (1.0 - zoom) * jnp.array([H, W]) / 2.0 + shift
|
|
55
|
+
return jax.image.scale_and_translate(img, img.shape, (1, 2),
|
|
56
|
+
scale, translation, method="linear")
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def random_rotation(rkey, img, max_angle_degrees=15.0):
|
|
60
|
+
"""Rotation d'angle quelconque autour du centre, par map_coordinates :
|
|
61
|
+
on construit la grille de coordonnées TOURNÉE, on interpole (ordre 1).
|
|
62
|
+
C'est la rotation inverse qu'on applique aux coordonnées — pour remplir
|
|
63
|
+
chaque pixel de la sortie, on va chercher d'où il vient dans l'entrée."""
|
|
64
|
+
C, H, W = img.shape
|
|
65
|
+
angle = jnp.deg2rad(jr.uniform(rkey, minval=-max_angle_degrees,
|
|
66
|
+
maxval=max_angle_degrees))
|
|
67
|
+
cos, sin = jnp.cos(angle), jnp.sin(angle)
|
|
68
|
+
rows = jnp.arange(H) - (H - 1) / 2.0
|
|
69
|
+
cols = jnp.arange(W) - (W - 1) / 2.0
|
|
70
|
+
r, c = jnp.meshgrid(rows, cols, indexing="ij")
|
|
71
|
+
src_r = cos * r - sin * c + (H - 1) / 2.0
|
|
72
|
+
src_c = sin * r + cos * c + (W - 1) / 2.0
|
|
73
|
+
|
|
74
|
+
def rotate_channel(channel):
|
|
75
|
+
return jax.scipy.ndimage.map_coordinates(channel, [src_r, src_c],
|
|
76
|
+
order=1, mode="constant", cval=0.0)
|
|
77
|
+
|
|
78
|
+
return jax.vmap(rotate_channel)(img)
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
def random_augmentation(rkey, img):
|
|
82
|
+
"""Pipeline d'exemple : composer, c'est splitter la clé et enchaîner.
|
|
83
|
+
À adapter au jeu de données — le flip n'y est volontairement pas, c'est
|
|
84
|
+
la transformation qu'il faut choisir en connaissance de cause."""
|
|
85
|
+
k1, k2, k3, k4 = jr.split(rkey, 4)
|
|
86
|
+
img = random_scale_translate(k1, img)
|
|
87
|
+
img = random_rotation(k2, img)
|
|
88
|
+
img = random_brightness_contrast(k3, img)
|
|
89
|
+
img = random_noise(k4, img)
|
|
90
|
+
return img
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
if __name__ == "__main__":
|
|
94
|
+
import matplotlib
|
|
95
|
+
|
|
96
|
+
# une image test : un carré excentré avec un coin marqué (l'asymétrie rend
|
|
97
|
+
# les transformations visibles)
|
|
98
|
+
img = jnp.zeros((1, 32, 32))
|
|
99
|
+
img = img.at[0, 8:20, 6:18].set(0.6)
|
|
100
|
+
img = img.at[0, 8:12, 6:10].set(1.0)
|
|
101
|
+
|
|
102
|
+
# formes conservées, jit + vmap passent
|
|
103
|
+
batch = jnp.stack([img] * 8)
|
|
104
|
+
keys = jr.split(jr.key(0), 8)
|
|
105
|
+
augmented = jax.jit(jax.vmap(random_augmentation))(keys, batch)
|
|
106
|
+
assert augmented.shape == batch.shape
|
|
107
|
+
# même clé -> même transformation (reproductible). Tolérance large : jit
|
|
108
|
+
# réordonne les flottants, eager et compilé diffèrent au dernier bit.
|
|
109
|
+
again = jax.vmap(random_augmentation)(keys, batch)
|
|
110
|
+
assert jnp.allclose(augmented, again, atol=1e-5)
|
|
111
|
+
print("jit + vmap + reproductibilité : OK")
|
|
112
|
+
|
|
113
|
+
from matplotlib import pyplot as plt
|
|
114
|
+
|
|
115
|
+
fig, axs = plt.subplots(1, 9, figsize=(16, 2.2))
|
|
116
|
+
axs[0].imshow(img[0], cmap="gray", vmin=-0.2, vmax=1.2)
|
|
117
|
+
axs[0].set_title("originale")
|
|
118
|
+
for i in range(8):
|
|
119
|
+
axs[i + 1].imshow(augmented[i, 0], cmap="gray", vmin=-0.2, vmax=1.2)
|
|
120
|
+
for ax in axs:
|
|
121
|
+
ax.axis("off")
|
|
122
|
+
plt.tight_layout()
|
|
123
|
+
plt.show()
|
yax/layers/Conv_layer.py
ADDED
|
@@ -0,0 +1,70 @@
|
|
|
1
|
+
"""Couche de convolution 2D.
|
|
2
|
+
|
|
3
|
+
Implémentée directement sur lax.conv_general_dilated, le primitif XLA — c'est
|
|
4
|
+
exactement ce que ferait une bibliothèque, il n'y a aucune optimisation cachée
|
|
5
|
+
au-dessus. L'API du primitif est générale donc verbeuse (les dimension_numbers
|
|
6
|
+
disent qui est batch, canaux, hauteur, largeur) ; cette couche la fige dans la
|
|
7
|
+
convention du cours.
|
|
8
|
+
|
|
9
|
+
Convention yax2 : écrit pour UN échantillon, en channels-first (canaux, H, W).
|
|
10
|
+
Le batch vient de vmap. Initialisation de Glorot, avec le fan_in d'une
|
|
11
|
+
convolution : canaux d'entrée x hauteur x largeur du noyau.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
from yax.core import Module, StaticField, tree_at
|
|
15
|
+
import jax
|
|
16
|
+
import jax.numpy as jnp
|
|
17
|
+
import jax.random as jr
|
|
18
|
+
from jax import lax
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class Conv_layer(Module):
|
|
22
|
+
weight: jnp.ndarray # (dim_out, dim_in, k, k)
|
|
23
|
+
bias: jnp.ndarray # (dim_out,)
|
|
24
|
+
|
|
25
|
+
stride: int = StaticField()
|
|
26
|
+
padding: str = StaticField()
|
|
27
|
+
|
|
28
|
+
def __init__(self, dim_in, dim_out, kernel_size, rkey, *, stride=1,
|
|
29
|
+
padding="SAME"):
|
|
30
|
+
# padding="SAME" : la sortie garde la taille spatiale de l'entrée
|
|
31
|
+
# (à stride 1) ; "VALID" : pas de remplissage, la taille fond de
|
|
32
|
+
# kernel_size-1. stride>1 : l'alternative au pooling.
|
|
33
|
+
assert padding in ("SAME", "VALID")
|
|
34
|
+
fan_in = dim_in * kernel_size ** 2
|
|
35
|
+
fan_out = dim_out * kernel_size ** 2
|
|
36
|
+
lim = jnp.sqrt(6.0 / (fan_in + fan_out))
|
|
37
|
+
self.weight = jr.uniform(rkey, (dim_out, dim_in, kernel_size, kernel_size),
|
|
38
|
+
minval=-lim, maxval=lim)
|
|
39
|
+
self.bias = jnp.zeros((dim_out,))
|
|
40
|
+
self.stride = stride
|
|
41
|
+
self.padding = padding
|
|
42
|
+
|
|
43
|
+
def apply(self, x, rkey=None):
|
|
44
|
+
# x : (canaux, H, W) — rkey ignoré : couche déterministe.
|
|
45
|
+
# Le primitif attend un batch : on en fabrique un de taille 1.
|
|
46
|
+
y = lax.conv_general_dilated(
|
|
47
|
+
x[None], # (1, dim_in, H, W)
|
|
48
|
+
self.weight, # (dim_out, dim_in, k, k)
|
|
49
|
+
window_strides=(self.stride, self.stride),
|
|
50
|
+
padding=self.padding,
|
|
51
|
+
dimension_numbers=("NCHW", "OIHW", "NCHW"))
|
|
52
|
+
return y[0] + self.bias[:, None, None]
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
if __name__ == "__main__":
|
|
56
|
+
import equinox as eqx # reference de verification uniquement
|
|
57
|
+
# formes, puis vérification numérique contre eqx.nn.Conv2d — le schéma
|
|
58
|
+
# constant : réimplémenter, vérifier contre la référence, adopter LA NÔTRE
|
|
59
|
+
x = jr.normal(jr.key(0), (3, 16, 16))
|
|
60
|
+
print("SAME stride 1 :", Conv_layer(3, 8, 3, jr.key(1)).apply(x).shape)
|
|
61
|
+
print("SAME stride 2 :", Conv_layer(3, 8, 3, jr.key(1), stride=2).apply(x).shape)
|
|
62
|
+
print("VALID stride 1 :", Conv_layer(3, 8, 3, jr.key(1), padding="VALID").apply(x).shape)
|
|
63
|
+
|
|
64
|
+
for stride, padding in [(1, "SAME"), (2, "SAME"), (1, "VALID")]:
|
|
65
|
+
ref = eqx.nn.Conv2d(3, 8, 3, stride=stride, padding=padding, key=jr.key(2))
|
|
66
|
+
ours = Conv_layer(3, 8, 3, jr.key(3), stride=stride, padding=padding)
|
|
67
|
+
ours = tree_at(lambda m: (m.weight, m.bias), ours,
|
|
68
|
+
(ref.weight, ref.bias.reshape(-1)))
|
|
69
|
+
assert jnp.allclose(ours.apply(x), ref(x), atol=1e-5), (stride, padding)
|
|
70
|
+
print("Conv_layer == eqx.nn.Conv2d : OK (SAME/VALID, stride 1/2)")
|
yax/layers/Dropout.py
ADDED
|
@@ -0,0 +1,63 @@
|
|
|
1
|
+
"""Dropout.
|
|
2
|
+
|
|
3
|
+
Deux choses distinctes, et découplées :
|
|
4
|
+
|
|
5
|
+
- le MODE : le drapeau statique `inference`, porté par tout yax.Module et
|
|
6
|
+
basculé par `model = model.set_inference(True/False)` — le Trainer s'en
|
|
7
|
+
charge (entraînement en False, validation et modèle rendu en True) ;
|
|
8
|
+
- la SOURCE D'ALÉATOIRE : l'argument `rkey` de `apply`.
|
|
9
|
+
|
|
10
|
+
En mode inférence, le dropout est l'identité, clé ou pas. En mode
|
|
11
|
+
entraînement, la clé est obligatoire — l'oublier est une erreur explicite,
|
|
12
|
+
pas un comportement silencieux.
|
|
13
|
+
|
|
14
|
+
Le découplage a un bonus : le MC-dropout (échantillonner des prédictions avec
|
|
15
|
+
le dropout actif) s'obtient en repassant simplement le modèle en
|
|
16
|
+
`set_inference(False)` et en fournissant des clés à l'évaluation.
|
|
17
|
+
"""
|
|
18
|
+
|
|
19
|
+
import jax.numpy as jnp
|
|
20
|
+
import jax.random as jr
|
|
21
|
+
|
|
22
|
+
from yax.core import Module, StaticField
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class Dropout(Module):
|
|
26
|
+
rate: float = StaticField(default_value=0.0)
|
|
27
|
+
|
|
28
|
+
def __init__(self, rate):
|
|
29
|
+
assert 0.0 <= rate < 1.0, f"rate:{rate} doit etre dans [0,1)"
|
|
30
|
+
self.rate = rate
|
|
31
|
+
|
|
32
|
+
def apply(self, x, rkey=None):
|
|
33
|
+
if self.inference or self.rate == 0.0:
|
|
34
|
+
return x
|
|
35
|
+
if rkey is None:
|
|
36
|
+
raise ValueError(
|
|
37
|
+
"Dropout en mode entrainement (inference=False) : une cle "
|
|
38
|
+
"est requise — apply(x, rkey). Pour evaluer, basculer le "
|
|
39
|
+
"modele avec model.set_inference(True).")
|
|
40
|
+
keep_prob = 1.0 - self.rate
|
|
41
|
+
mask = jr.bernoulli(rkey, keep_prob, x.shape)
|
|
42
|
+
# division par keep_prob : l'esperance de la sortie est inchangee,
|
|
43
|
+
# rien a corriger a l'evaluation ("inverted dropout").
|
|
44
|
+
return jnp.where(mask, x / keep_prob, 0.0)
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
if __name__ == "__main__":
|
|
48
|
+
drop = Dropout(0.4)
|
|
49
|
+
x = jnp.ones((10_000,))
|
|
50
|
+
|
|
51
|
+
y_train = drop.apply(x, jr.key(0))
|
|
52
|
+
print(f"fraction annulée : {float(jnp.mean(y_train == 0.0)):.3f} (attendu ~0.4)")
|
|
53
|
+
print(f"moyenne : {float(jnp.mean(y_train)):.3f} (attendu ~1.0)")
|
|
54
|
+
|
|
55
|
+
try:
|
|
56
|
+
drop.apply(x) # entrainement sans cle : erreur explicite
|
|
57
|
+
except ValueError as e:
|
|
58
|
+
print("ValueError :", str(e)[:60], "...")
|
|
59
|
+
|
|
60
|
+
drop_eval = drop.set_inference(True)
|
|
61
|
+
assert (drop_eval.apply(x) == x).all()
|
|
62
|
+
assert (drop_eval.apply(x, jr.key(0)) == x).all() # cle ignoree en inference
|
|
63
|
+
print("mode inference : identité, avec ou sans clé : OK")
|
yax/layers/Embedding.py
ADDED
|
@@ -0,0 +1,29 @@
|
|
|
1
|
+
"""Embedding : une table apprise, indexée par des entiers.
|
|
2
|
+
|
|
3
|
+
Le one-hot suivi d'un Linear EST un embedding, en moins efficace (chapitre 8) :
|
|
4
|
+
ici on lit directement la ligne de la table.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from yax.core import Module
|
|
8
|
+
import jax.numpy as jnp
|
|
9
|
+
import jax.random as jr
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class Embedding(Module):
|
|
13
|
+
weight: jnp.ndarray
|
|
14
|
+
|
|
15
|
+
def __init__(self, num_embeddings, dim, rkey):
|
|
16
|
+
# initialisation normale d'écart-type 0.02, l'usage des transformers
|
|
17
|
+
self.weight = 0.02 * jr.normal(rkey, (num_embeddings, dim))
|
|
18
|
+
|
|
19
|
+
def apply(self, ids, rkey=None):
|
|
20
|
+
# ids : entiers, forme quelconque -> sortie ids.shape + (dim,).
|
|
21
|
+
# Indexation par entiers : valide sous jit (contrairement au masque
|
|
22
|
+
# booléen entre crochets, chapitre 1). rkey ignoré : couche déterministe.
|
|
23
|
+
return self.weight[ids]
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
if __name__ == "__main__":
|
|
27
|
+
emb = Embedding(100, 32, jr.key(0))
|
|
28
|
+
ids = jnp.array([1, 5, 10])
|
|
29
|
+
print("sortie :", emb.apply(ids).shape) # (3, 32)
|