L'atout de yax est la simplicité :
Un module de yax est simplement un pytree dont les feuilles sont exactement les paramètres.
Conséquence :
jax.grad, jax.jit, jax.vmap et optax
s'utilisent exactement comme dans la documentation jax — rien à
désapprendre, rien à envelopper, rien à filtrer (ce qui offre un net gain de performance sur les petits modèles. Sur les gros cet avantage s'estompe, car le cout dominant est celui des calculs).
Né pour un cours de master, yax est dimensionné pour servir au-delà : des modèles de recherche compacts et lisibles. Tout ce que yax ne fournit pas s'écrit en jax ordinaire, sans friction.
Sommaire
- Installation
- Démarrage rapide
- yax.Module : un modèle est un pytree
- Aléatoire et mode inférence
- Fonctions d'activation
- Les couches
- Les modèles
- Modèles génératifs
- L'entraînement : le Trainer
- Conventions
- Les démonstrations
Installation
La distribution s'appelle yaxlib, l'import s'appelle
yax :
pip install yaxlib # dépendances : jax, optax, numpy, matplotlib
pip install yaxlib[dev] # + pytest
Démarrage rapide
import jax
import jax.numpy as jnp
import jax.random as jr
import yax
model = yax.MLP((1, 32, 32, 1), "tanh", jr.key(0))
y = model.apply(jnp.ones(1)) # UN échantillon…
Y = jax.vmap(model.apply)(X) # …le batch vient de jax.vmap
Et le gradient s'obtient comme dans n'importe quel tutoriel jax — le modèle est le pytree de paramètres :
def loss(model):
pred = jax.vmap(model.apply)(X)
return jnp.mean((pred - Y) ** 2)
grads = jax.grad(loss)(model) # un MLP dont les feuilles sont les gradients
optimizer = optax.adam(1e-3)
opt_state = optimizer.init(model) # optax accepte le modèle tel quel
yax.Module : un modèle est un pytree
Les champs d'un module se déclarent par annotations de classe et se rangent en deux familles :
- Les champs dynamiques (annotation seule) : les feuilles du pytree, c'est-à-dire les paramètres. Ils ne peuvent contenir que des tableaux jax, des sous-modules, ou des list/tuple/dict de ceux-ci — tout écart est refusé à la construction, avec un message explicite.
- Les champs statiques (
= yax.StaticField()) : des métadonnées rangées dans la structure du pytree, invisibles pourgradet les optimiseurs. - Attention, comme pour les fameuses Dataclass de python, le typage avec les deux-points est indispensable pour différencier les champs d'instance des champs de classe.
Commençons par le commencement, codons à la main le fameux layer Affine (qui existe déjà sous le nom de yax.Linear)
class Affine(yax.Module):
weight: jnp.ndarray # dynamique : un paramètre
bias: jnp.ndarray
dim_out: int = yax.StaticField() # statique : une métadonnée
def __init__(self, dim_in, dim_out, rkey):
self.weight = jr.normal(rkey, (dim_in, dim_out)) / jnp.sqrt(dim_in)
self.bias = jnp.zeros(dim_out)
self.dim_out = dim_out
def apply(self, x, rkey=None):
return x @ self.weight + self.bias
La composition est
gratuite : Par exemple, empilons notre layer Affine en tête
d'une pile de convolutions :
class PetitCNN(yax.Module):
convs: list
tete: Affine
def __init__(self, rkey):
rkey1, rkey2, rkey3 = jr.split(rkey, 3)
self.convs = [yax.Conv_nd(1, 8, 3, 2, rkey1, stride=2), # (1,28,28) -> (8,14,14)
yax.Conv_nd(8, 16, 3, 2, rkey2, stride=2)] # -> (16,7,7)
self.tete = Affine(16 * 7 * 7, 10, rkey3)
def apply(self, x, rkey=None):
for conv in self.convs:
x = jax.nn.relu(conv.apply(x))
return self.tete.apply(x.reshape(-1)) # 10 logits
(Affine n'était qu'un prétexte : cette couche existe déjà
dans yax, sous le nom Linear.)
yax.tree_pprint(model) affiche l'arbre complet, et les modules
étant immuables, on « modifie » avec
yax.tree_at(lambda m: m.tete.bias, model, jnp.ones(10)).
Aléatoire et mode inférence
Deux choses distinctes, et découplées :
- la source d'aléatoire passe par l'argument
rkeydeapply(dropout, reparamétrisation d'un VAE…). Une couche stochastique en mode entraînement exige sa clé — l'oublier provoque une erreur explicite. - le mode est le drapeau statique
inference(False par défaut), basculé récursivement parmodel = model.set_inference(True). En inférence, le dropout VAE sont déterministes, donc sans clé.
Le Trainer intégré (à usage facultatif) entraîne avec
inference=False alors qu'il valide et retourne le meilleur modèle en
inference=True.
Fonctions d'activation
Partout où yax attend une activation, on donne une chaîne ou une
fonction : la chaîne est piochée dans le dictionnaire
yax.ACTIVATIONS
(identity, relu, leaky_relu, relu6, tanh, sigmoid, gelu,
gelu_approximate, silu/swish, elu, celu, selu, mish, softplus),
une fonction passe telle quelle.
yax.MLP((4, 64, 2), "gelu", jr.key(0)) # une chaîne
yax.MLP((4, 64, 2), jax.nn.silu, jr.key(0)) # une fonction jax
yax.MLP((4, 64, 2), ma_fonction, jr.key(0)) # la vôtre
Pour une pente apprise, voir la couche PReLU.
Les couches
Toutes écrites pour un échantillon (le batch vient de
jax.vmap), toutes avec la signature
apply(x, rkey=None).
| Couche | Rôle |
|---|---|
Linear | affine x @ W + b |
MLP | perceptron multicouches |
Dropout | identité en inférence, random-clé obligatoire en entraînement |
PReLU | leaky-relu à pente apprise (paramètre partagée ou par canal) |
LayerNorm | normalisation classique des Transformers |
Embedding | table indices → vecteurs |
sinusoidal_positional_encoding | encodage positionnel statique |
Conv_nd | convolution 1D/2D/3D… (nb_dims requis), channels-first (C, *spatial), sur lax.conv_general_dilated |
RNN_layer | récurrence GRU ou LSTM (cellules GRUCell/LSTMCell exposées) |
MultiHeadAttention | attention multi-têtes ; causal_mask, padding_mask, cartes d'attention |
TransformerBlock | attention + feed-forward, pré-norm, dropout |
MessagePassing_layer | graphes : message par arête, agrégation sum/mean/max/attention (softmax par récepteur, façon GATv2) |
Les modèles
Des références compactes et lisibles, dans yax.models — chacune
avec sa perte dans le même fichier :
| Modèle | Rôle |
|---|---|
UNet_nd | segmentation en tout rang (signal, image, volume) : descente/remontée avec connexions de saut, interpolation + convolution |
MiniYOLO | détection à une passe : grille 8×8, une boîte par cellule, yolo_loss, detect avec NMS |
VAE | auto-encodeur variationnel — vae_loss (ELBO), generate |
RealNVP | flot normalisant à vraisemblance exacte — realnvp_loss, sample (flot inversé) |
Diffusion | DDPM — diffusion_loss, sample (débruitage ancestral) |
Modèles génératifs
Les trois modèles génératifs s'entraînent avec le Trainer tel quel,
en non-supervisé : on passe Y_train = X_train
(la perte du VAE reconstruit y — donner un x bruité
et un y propre en fait d'ailleurs un débruiteur) :
from yax.models.VAE import VAE, vae_loss
model = VAE(dim_in=2, dim_hidden=64, dim_latent=2, rkey=jr.key(0))
run = train(model, mother_folder, config,
DatasetSampler(X_train, X_train, 64), (X_val, X_val), # non supervisé : Y = X
rkey=jr.key(1), loss_fn=vae_loss)
neuves = run.trained_model.generate(jr.key(2), 500) # des données neuves
Chaque famille illustre un régime d'aléa différent, tous compatibles avec la
validation déterministe du Trainer : le flot est entièrement déterministe,
le VAE échantillonne à l'entraînement mais encode par la moyenne en inférence,
et la diffusion — qui tire t et ε même pour
évaluer — utilise une clé fixe quand rkey=None.
L'entraînement : le Trainer
Facultatif — tout s'entraîne aussi à la main en dix lignes d'optax — mais le yax.Trainer standardise la boucle, la validation et effectue des checkpoints à chaque fois que la loss bas un nouveau record:
from yax.training.Trainer import train, load_run, find_best_run, OUT_FOLDER, DatasetSampler
from yax.training.configs import TrainConfig
from yax.training.losses import mse_loss
config = TrainConfig(learning_rate=2e-2, n_epoch=300, lr_final_ratio=0.01) # cosine decay
run = train(model, os.path.join(OUT_FOLDER, "mon_experience"), config,
DatasetSampler(X_train, Y_train, 32), # une époque = un mélange, batchs de 32
(X_val, Y_val), # la validation : un batch fixe
rkey=jr.key(1), loss_fn=mse_loss)
run.trained_model # le meilleur modèle, en inference=True — prêt à évaluer
run.loss # sa loss de validation
Le Trainer ne connaît pas les données : une époque est un appel au
sampler suivi d'une validation. DatasetSampler(X, Y, batch_size)
mélange un jeu fini sans remise ; FunctionSampler(f, nb_batches) tire
chaque batch par f(rkey) — pour les problèmes sans données
(physique, PINN, Ritz), où y vaut None. La validation est
soit un batch fixe (x_val, y_val), soit elle-même un sampler, appelé avec
une clé constante : le jeu de validation est identique à toutes les
époques et pour tous les runs d'un mother_folder. Ainsi TrainConfig
ne parle plus que d'optimisation ; la taille des batchs est une propriété du sampler.
# Deep Ritz : -u'' = f sur [0,1], sans aucune donnée (demos/ritz_demo_poisson.py)
sampler = FunctionSampler(lambda rkey: (jr.uniform(rkey, (512,)), None), nb_batches=25) # y = None : pas de cible
validation = FunctionSampler(lambda rkey: (jnp.linspace(0.0, 1.0, 1025), None), nb_batches=1)
run = train(model, mother_folder, config, sampler, validation, rkey=jr.key(1), loss_fn=ritz_loss)
- Pertes : signature commune
loss_fn(model, x, y, rkey). Fournies :mse_loss,bce_logits_loss,softmax_ce_loss,dice_loss— toutes sur les logits, jamais sur des probabilités. - Optimiseur : dans la config, et sous forme de
données —
optimizerest un nom du dictionnaireyax.OPTIMIZERS("adam"par défaut,"adamw","lion","sgd"…) etoptimizer_optionsses réglages :TrainConfig(5e-3, 300, optimizer="adamw", optimizer_options={"weight_decay": 0.1}). La config étant enregistrée avec le run, celui-ci dit à lui seul comment il a été entraîné. Le Trainer fabrique le schedule en cosinus — seul à connaître le nombre de steps — et le passe au constructeur. - Quasi-Newton :
optimizer="lbfgs"suffit, sans rien réécrire. Le Trainer fournit à tous les optimiseurs la valeur de la perte et la fonction qui la calcule — ce que réclame une recherche linéaire, et que les autres ignorent ; le coût pour eux est nul (mesuré, et pas à un bit près de différence). Condition d'emploi d'une recherche linéaire : un objectif déterministe, donc un seul batch fixe (DatasetSampler(X, Y, len(X))). - Mode : le Trainer entraîne en
inference=False(les clés de dropout sont dérivées derkey, une par step), valide et rend le meilleur modèle eninference=True. - Checkpoints : un
mother_folder= une expérience = un jeu de validation fixe ; chaque run y est un sous-dossier numéroté (trained_model,opt_state,loss,config,history) ;find_best_rundépartage les runs,load_runrecharge. Le critère est la loss de validation — pas d'accuracy cachée.config.patiencearrête la boucle après N époques sans record : un budget épargné, pas un gain de qualité — le modèle rendu est de toute façon le meilleur. - History : les loss d'entraînement par step, celles
de validation par époque — et une méthode
plot.
run.history.plot() — le train par step en
transparent (le bruit des batchs), sa moyenne par époque en opaque, la
validation en petits points, les records en vert, et le
meilleur modèle — celui que train() rend et sauvegarde — en
gros point vert. Ici un sur-apprentissage différé, provoqué
exprès (MLP surdimensionné, 24 points bruités, learning rate
constant) : le train continue de descendre en mémorisant le bruit
pendant que la validation remonte — et le checkpointing garde le modèle
du creux.Recharger un checkpoint
from yax.training.Trainer import load_run, find_best_run
run = load_run(find_best_run(mother_folder)) # le run de plus basse val loss
model = run.trained_model # déjà en inference=True : prêt à évaluer
run.history.plot()
run = load_run(folder, "trained_model", "config") # ou le strict nécessaire
train et load_run rendent tous deux un objet de type
yax.Run : la complétion automatique propose les attributs
propose trained_model, opt_state,
loss, config, history,
folder, et un nom inconnu est une erreur explicite qui liste
les champs valides. (trained_model et non model :
dans run = train(model, ...), pas de confusion avec le modèle
initial ; loss sans « best » : le « meilleur » se dit
des runs, cf. find_best_run.) Pour poursuivre un entraînement, repasser simplement
le modèle à train() (il le rebascule lui-même en mode
entraînement) — avec optimizer_state=run.opt_state si l'on veut
conserver les moments d'Adam, ou sans, pour repartir d'un optimiseur neuf (le
bon choix quand les données ou le learning rate changent).
Pourquoi pickle ?
Un modèle yax est un objet Python ordinaire — un pytree de tableaux et de
métadonnées. pickle le sauvegarde entier (structure,
poids, hyperparamètres) en une ligne et le rend prêt à l'emploi au
rechargement : pas de format de checkpoint à définir, pas de squelette de
modèle à reconstruire avant d'y verser des poids. La contrepartie est
assumée : recharger demande le même code (même version de yax) — ce qui
est le régime naturel de checkpoints d'expériences, qui vivent avec leur
code.
Conventions
applyet non__call__: on distingue le pytree de paramètres de la fonction qu'il définit.- Écrit pour un échantillon ; le batch vient de
jax.vmap(helperbatch_applydans les pertes). - Images en channels-first
(C, H, W). - Les variables contenant des clés PRNG se nomment avec
rkey(rkey,rkeys,rkey_model…) — « key » seul est trop générique. - Classification : les modèles rendent des logits ; sigmoid et softmax ne servent qu'à l'affichage.
Les démonstrations
Le dossier demos/ contient une démonstration par famille de
modèles, sur des données synthétiques calibrées pour converger en
quelques dizaines de secondes sur CPU : on voit chaque architecture
apprendre de bout en bout, sans rien télécharger. (Dépôt GitHub à venir ;
en attendant, les démos sont livrées dans l'archive source de PyPI.)
| Démo | Ce qu'elle montre |
|---|---|
mlp_demo | régression 1D avec la descente de gradient écrite à la main — la preuve que le modèle est un pytree ordinaire, sans optax ni Trainer |
cnn_demo_classif | classification de formes (disque, carré, croix) posées sur fond bruité, par CNN |
rnn_demo_classif | classification binaire de séquences 2D, par GRU ou LSTM |
transformer_demo_charlm | modèle de langue caractère par caractère (masque causal), avec génération de texte |
unet_demo_segmentation | segmentation binaire par pixel (U-Net), métrique de Dice |
yolo_demo_detection | détection de boîtes (MiniYOLO) : encodage en grille, perte composite, NMS |
gnn_demo_nodes | classification de nœuds sur un graphe à deux communautés — avec un MLP témoin qui ignore le graphe, pour mesurer ce que le message passing apporte — et les quatre agrégations comparées |
vae_demo_generation | génération 2D des « deux
lunes » par VAE (non supervisé : Y = X) |
realnvp_demo_generation | les deux lunes par flot normalisant, avec la vraisemblance exacte en nats |
diffusion_demo_generation | une spirale par diffusion (DDPM), échantillonnage ancestral |
Les trois démos génératives rapportent la même métrique — la distance de Chamfer entre points générés et données de validation, avant et après entraînement — pour comparer les familles à armes égales.
yaxlib — licence MIT. Conçu pour les cours de réseaux de neurones du master (jax, M1–M2) et ouvert aux usages de recherche.