Vincent Vigon
Logo yax : un yack bleu formant le Y, un A vert incliné, un + violet

yax

Réseaux de neurones modulaires, au plus près du jax de base.

pip install yaxlib import yax PyPI

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

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 :

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 :

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

CoucheRôle
Linearaffine x @ W + b
MLPperceptron multicouches
Dropoutidentité en inférence, random-clé obligatoire en entraînement
PReLUleaky-relu à pente apprise (paramètre partagée ou par canal)
LayerNormnormalisation classique des Transformers
Embeddingtable indices → vecteurs
sinusoidal_positional_encodingencodage positionnel statique
Conv_ndconvolution 1D/2D/3D… (nb_dims requis), channels-first (C, *spatial), sur lax.conv_general_dilated
RNN_layerrécurrence GRU ou LSTM (cellules GRUCell/LSTMCell exposées)
MultiHeadAttentionattention multi-têtes ; causal_mask, padding_mask, cartes d'attention
TransformerBlockattention + feed-forward, pré-norm, dropout
MessagePassing_layergraphes : 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èleRôle
UNet_ndsegmentation en tout rang (signal, image, volume) : descente/remontée avec connexions de saut, interpolation + convolution
MiniYOLOdétection à une passe : grille 8×8, une boîte par cellule, yolo_loss, detect avec NMS
VAEauto-encodeur variationnel — vae_loss (ELBO), generate
RealNVPflot normalisant à vraisemblance exacte — realnvp_loss, sample (flot inversé)
DiffusionDDPM — 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)
Courbe d'entraînement : train par step transparent et moyenne par époque opaque en rouge, validation en petits points bleus, records en vert, meilleur modèle en gros point vert
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

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émoCe qu'elle montre
mlp_demoré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_classifclassification de formes (disque, carré, croix) posées sur fond bruité, par CNN
rnn_demo_classifclassification binaire de séquences 2D, par GRU ou LSTM
transformer_demo_charlmmodèle de langue caractère par caractère (masque causal), avec génération de texte
unet_demo_segmentationsegmentation binaire par pixel (U-Net), métrique de Dice
yolo_demo_detectiondétection de boîtes (MiniYOLO) : encodage en grille, perte composite, NMS
gnn_demo_nodesclassification 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_generationgénération 2D des « deux lunes » par VAE (non supervisé : Y = X)
realnvp_demo_generationles deux lunes par flot normalisant, avec la vraisemblance exacte en nats
diffusion_demo_generationune 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.