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 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()
@@ -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")
@@ -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)