yax — le cœur¶
La classe de base des modèles et les fonctions qui opèrent sur un modèle. Ces noms s'écrivent directement sous yax.
Module
¶
Classe de base de tous les modèles et de toutes les couches.
Un Module est un pytree jax. On déclare ses champs par annotation de
classe, on les remplit dans __init__, et on écrit la méthode
apply(x, rkey=None), qui calcule la sortie pour un exemple.
Les champs sont de deux sortes :
- dynamiques (annotation seule) : les paramètres. Ils contiennent des
tableaux jax, des sous-modules, ou des listes, tuples et dictionnaires de
ceux-ci.
jax.grad,jax.jitet les optimiseurs d'optax les voient ; - statiques (
= yax.StaticField()) : tout le reste, tailles, options, fonctions d'activation. Ils décrivent la forme du modèle et ne sont pas entraînés.
Une valeur non autorisée dans un champ dynamique (un flottant, une chaîne…)
déclenche une erreur dès la construction. Un module est immuable : pour
changer un champ, utiliser yax.tree_at.
class Classifieur(yax.Module):
mlp: yax.layers.MLP
dim_in: int = yax.StaticField()
def __init__(self, dim_in, rkey):
self.mlp = yax.layers.MLP((dim_in, 32, 1), "relu", rkey)
self.dim_in = dim_in
def apply(self, x, rkey=None):
return self.mlp.apply(x)
Attributs :
| Nom | Type | Description |
|---|---|---|
inference |
bool
|
mode inférence, |
Code source dans yax/core.py
128 | |
166 167 168 | |
set_inference
¶
set_inference(value)
Renvoie une copie du modèle en mode inférence (True) ou entraînement (False).
Le changement s'applique à tous les sous-modules ; le modèle d'origine
n'est pas modifié. yax.training.train gère ce mode automatiquement.
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
value
|
bool
|
|
obligatoire |
Code source dans yax/core.py
288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 | |
StaticField
¶
yax.StaticField(*, default_value=None)
Déclare un champ statique d'un yax.Module.
Un champ statique contient ce qui n'est pas un paramètre : une taille, une option, une fonction. Il fait partie de la structure du modèle et n'est pas entraîné. Il peut aussi contenir un tableau constant : un masque, un encodage de position, le calendrier de bruit d'une diffusion.
class Couche(yax.Module):
taille: int = yax.StaticField()
mode: str = yax.StaticField(default_value="somme")
Tableaux statiques : quelques Mo au plus. Un tableau statique n'est pas
une entrée du programme compilé par jax.jit, mais une constante écrite
dedans. Sa taille alourdit donc la compilation : rien de visible en
dessous du Mo, une compilation plusieurs fois plus longue au-delà de la
dizaine de Mo. Pour un gros tableau constant, des poids pré-entraînés par
exemple, mieux vaut un champ dynamique, figé dans apply par
jax.lax.stop_gradient (attention : la décroissance des poids d'adamw
le modifierait quand même). Une fois compilé, en revanche, un tableau
statique ne coûte rien de plus à chaque pas, au contraire : on ne calcule
ni son gradient ni sa mise à jour.
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
default_value
|
Any
|
valeur prise par le champ si |
None
|
Code source dans yax/core.py
25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 | |
tree_at
¶
yax.tree_at(where, pytree, replace)
Renvoie une copie d'un modèle dont certaines feuilles sont remplacées.
model2 = yax.tree_at(lambda m: m.layers[-1].bias, model, jnp.ones(3))
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
where
|
Callable
|
fonction qui, appliquée au modèle, renvoie la feuille à remplacer, ou un tuple de feuilles. |
obligatoire |
pytree
|
T
|
le modèle d'origine, qui n'est pas modifié. |
obligatoire |
replace
|
la nouvelle valeur, ou un tuple de valeurs. |
obligatoire |
Renvoie :
| Type | Description |
|---|---|
T
|
Le modèle modifié. |
Lève :
| Type | Description |
|---|---|
ValueError
|
si |
Code source dans yax/core.py
508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 | |
batch_apply
¶
yax.batch_apply(model, x, rkey=None)
Applique un modèle à un lot d'exemples.
Un modèle yax traite un exemple à la fois ; batch_apply le vectorise
avec jax.vmap sur la première dimension de x.
y = yax.batch_apply(model, x) # évaluation, sans aléa
y = yax.batch_apply(model, x, rkey) # avec aléa (dropout...)
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
model
|
Module
|
le modèle. |
obligatoire |
x
|
ArrayLike
|
le lot, de forme |
obligatoire |
rkey
|
Array | None
|
clé aléatoire, répartie entre les exemples ; |
None
|
Renvoie :
| Type | Description |
|---|---|
Array
|
Les sorties du modèle pour chaque exemple, de forme |
Code source dans yax/core.py
362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 | |
pprint
¶
yax.pprint(x)
Affiche l'arborescence d'un modèle, en texte.
Les tableaux sont résumés par leur type et leur forme (f32[2,16]) ; les
champs statiques et les fonctions sont affichés en clair. Voir aussi
ipprint, sa version dépliable pour les notebooks.
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
x
|
Any
|
un modèle ou un pytree quelconque. |
obligatoire |
Code source dans yax/core.py
387 388 389 390 391 392 393 394 395 396 397 | |
ipprint
¶
yax.ipprint(x)
Affiche l'arborescence d'un modèle sous forme dépliable, dans un notebook.
Chaque sous-module se déplie d'un clic et indique son nombre de
paramètres ; un clic sur un tableau affiche ses valeurs. Hors notebook,
ipprint se comporte comme pprint. Un modèle placé en dernière ligne
d'une cellule s'affiche de la même façon.
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
x
|
Any
|
un modèle ou un pytree quelconque. |
obligatoire |
Code source dans yax/core.py
488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 | |