Aller au contenu

Démos

Chaque démo est un script autonome, synthétique, qui converge en quelques secondes de CPU :

python demos/mlp_demo.py

cnn_demo_classif.py

Classification d'images synthétiques par un petit CNN.

Code
"""Classification d'images synthétiques par un petit CNN.

Trois classes de formes — disque, carré, croix — posées à position et taille
aléatoires sur un fond bruité. Le jeu se génère en une ligne de vmap, converge
en quelques secondes, et l'invariance par translation de la convolution y est
réellement mise à l'épreuve (la position varie d'une image à l'autre).

Le sous-échantillonnage est fait par stride=2 (l'alternative au pooling,
chapitre 11) : 16x16 -> 8x8 -> 4x4, puis tête dense sur l'aplati.
"""

import os

import jax
import jax.numpy as jnp
import jax.random as jr

import yax

IMG_SIZE = 16
NB_CLASSES = 3  # 0: disque, 1: carré, 2: croix


def make_image(rkey, label):
    """UNE image (1, 16, 16) : une forme `label` sur fond bruité.

    Les rayons restent assez grands (4.5 à 5.5) pour qu'un disque et un carré
    diffèrent de plusieurs pixels aux coins : avec des formes plus petites,
    seuls les coins les distinguent, la tâche devient réellement ambiguë sous
    le bruit et la confusion disque/carré plafonne la justesse."""
    rkey_pos, rkey_radius, rkey_noise = jr.split(rkey, 3)
    center = jr.uniform(rkey_pos, (2,), minval=6.0, maxval=IMG_SIZE - 6.0)
    radius = jr.uniform(rkey_radius, minval=4.5, maxval=5.5)

    rows = jnp.arange(IMG_SIZE)[:, None]
    cols = jnp.arange(IMG_SIZE)[None, :]
    dr, dc = rows - center[0], cols - center[1]

    disque = dr**2 + dc**2 < radius**2
    carre = jnp.maximum(jnp.abs(dr), jnp.abs(dc)) < radius
    croix = ((jnp.abs(dr) < 1.0) | (jnp.abs(dc) < 1.0)) & (jnp.abs(dr) + jnp.abs(dc) < 2 * radius)

    shape = jnp.stack([disque, carre, croix])[label]
    img = jnp.where(shape, 1.0, 0.0) + 0.25 * jr.normal(rkey_noise, (IMG_SIZE, IMG_SIZE))
    return img[None, :, :]  # channels-first : (1, H, W)


def make_data(rkey, nb):
    rkey_labels, rkey_images = jr.split(rkey)
    labels = jr.randint(rkey_labels, (nb,), 0, NB_CLASSES)
    X = jax.vmap(make_image)(jr.split(rkey_images, nb), labels)
    return X, labels


def present_data():
    from matplotlib import pyplot as plt

    X, Y = make_data(jr.key(0), 12)
    fig, axs = plt.subplots(3, 4, figsize=(8, 6))
    for ax, img, label in zip(axs.flat, X, Y):
        ax.imshow(img[0], cmap="gray")
        ax.set_title(["disque", "carré", "croix"][int(label)])
        ax.axis("off")
    plt.tight_layout()
    plt.show()


class CNN_classifier(yax.Module):
    conv1: yax.layers.Conv_nd
    conv2: yax.layers.Conv_nd
    head: yax.layers.Linear

    def __init__(self, rkey):
        rkey1, rkey2, rkey3 = jr.split(rkey, 3)
        self.conv1 = yax.layers.Conv_nd(1, 8, 3, 2, rkey1, stride=2)   # (1,16,16) -> (8,8,8)
        self.conv2 = yax.layers.Conv_nd(8, 16, 3, 2, rkey2, stride=2)  # -> (16,4,4)
        self.head = yax.layers.Linear(16 * 4 * 4, NB_CLASSES, rkey3)

    def apply(self, x, rkey=None):
        x = jax.nn.relu(self.conv1.apply(x))
        x = jax.nn.relu(self.conv2.apply(x))
        return self.head.apply(x.reshape(-1))  # logits (NB_CLASSES,)


def accuracy(model, x, y):
    logits = yax.batch_apply(model, x, None)
    return float(jnp.mean(jnp.argmax(logits, axis=-1) == y))


def train_classif(verbose=True):
    rkey_train, rkey_val, rkey_model, rkey_shuffle = jr.split(jr.key(0), 4)
    X_train, Y_train = make_data(rkey_train, 1024)
    X_val, Y_val = make_data(rkey_val, 512)

    model = CNN_classifier(rkey_model)
    acc_before = accuracy(model, X_val, Y_val)

    config = yax.configs.TrainConfig(learning_rate=0.003, nb_epochs=100, lr_final_ratio=0.01)
    mother_folder = os.path.join(yax.training.OUT_FOLDER, "cnn_classif")
    run = yax.training.train(
        mother_folder, config, yax.obj.softmax_ce_from_integers, model,
        yax.training.DatasetSampler(X_train, Y_train, 64), (X_val, Y_val),
        rkey=rkey_shuffle, verbose=verbose)

    acc_after = accuracy(run.trained_model, X_val, Y_val)
    print(f"\naccuracy {acc_before:.3f} -> {acc_after:.3f}   (best val CE = {run.loss:.3g})")
    return run, acc_after


if __name__ == "__main__":
    from matplotlib import pyplot as plt

    present_data()
    run, acc = train_classif()
    history = yax.training.load_run(run.folder, "history").history
    history.plot(title=f"CNN formes : CE={run.loss:.2g}, accuracy={acc:.3f}")
    plt.tight_layout()
    plt.show()

diffusion_demo_generation.py

Génération 2D par diffusion (DDPM) sur la spirale.

Code
"""Génération 2D par diffusion (DDPM) sur la spirale.

Le cas le plus rugueux pour le Trainer : la loss tire t et eps au hasard,
donc l'évaluation aussi a besoin d'aléa. La convention rkey=None => clé fixe
rend la validation déterministe et comparable d'époque en époque — le
checkpointing « meilleur modèle » garde alors un sens.
"""

import os

import jax.random as jr


from generative_common import make_spiral, chamfer
import yax

if __name__ == "__main__":
    rkey_train, rkey_val, rkey_model, rkey_fit, rkey_gen = jr.split(jr.key(0), 5)
    X_train = make_spiral(rkey_train, 2000)
    X_val = make_spiral(rkey_val, 500)

    model = yax.models.Diffusion(dim=2, dim_hidden=64, nb_steps=100, rkey=rkey_model)
    avant = chamfer(model.sample(rkey_gen, 500), X_val)

    config = yax.configs.TrainConfig(learning_rate=2e-3, nb_epochs=200, lr_final_ratio=0.01)
    run = yax.training.train(
        os.path.join(yax.training.OUT_FOLDER, "diffusion_spirale"), config, yax.models.diffusion_loss, model,
        yax.training.DatasetSampler(X_train, X_train, 128), (X_val, X_val),   # non supervisé : Y = X
        rkey=rkey_fit, verbose=True)

    apres = chamfer(run.trained_model.sample(rkey_gen, 500), X_val)
    print(f"chamfer générés/données : {avant:.3f} avant, {apres:.3f} après")

fno_demo_operator.py

Apprendre un OPÉRATEUR : l'équation de la chaleur, par FNO 1D.

Code
"""Apprendre un OPÉRATEUR : l'équation de la chaleur, par FNO 1D.

L'application à apprendre envoie une fonction sur une fonction :
u0 -> u(t), la solution de la chaleur périodique du/dt = d²u/dx² au temps t.
En Fourier elle est exacte et diagonale — le mode k est multiplié par
exp(-(2*pi*k)² t) — ce qui donne des données parfaites en une ligne, et un
opérateur que le FNO peut représenter exactement.

Le clou du spectacle : le modèle est entraîné sur des grilles de 64 points,
puis évalué sur des grilles de 128 points SANS réentraînement — les poids du
FNO vivent dans l'espace des fréquences, pas sur la grille. Un CNN ordinaire
ne sait même pas ingérer la nouvelle résolution sans bricolage.
"""

import os

import jax.numpy as jnp
import jax.random as jr

import yax

TEMPS = 2e-3       # l'horizon de diffusion : les hauts modes fondent, sans disparaître
NB_COEFFS = 12     # bande des données : en dessous des modes du FNO (16)


def make_data(rkey, nb, resolution):
    """(X, Y) de formes (nb, 1, resolution) : u0 et u(TEMPS).

    On tire des coefficients de Fourier aléatoires à décroissance douce
    (des u0 lisses), et la cible est EXACTE : mode k multiplié par
    exp(-(2 pi k)^2 t). norm="forward" : les valeurs ne dépendent pas de la
    résolution d'échantillonnage — la même fonction à 64 ou 128 points.
    """
    rkey_re, rkey_im = jr.split(rkey)
    k = jnp.arange(NB_COEFFS)
    decroissance = 1.0 / (1.0 + k) ** 1.5
    coeffs = (jr.normal(rkey_re, (nb, NB_COEFFS)) +
              1j * jr.normal(rkey_im, (nb, NB_COEFFS))) * decroissance
    chaleur = jnp.exp(-((2.0 * jnp.pi * k) ** 2) * TEMPS)

    def champs(c):
        return jnp.fft.irfft(c, n=resolution, norm="forward")

    X = champs(coeffs)[:, None, :]              # (nb, 1, resolution)
    Y = champs(coeffs * chaleur)[:, None, :]
    return X, Y


def erreur_relative(model, X, Y):
    import jax
    pred = jax.vmap(model.apply)(X)
    return float(jnp.linalg.norm(pred - Y) / jnp.linalg.norm(Y))


if __name__ == "__main__":
    rkey_train, rkey_val, rkey_model, rkey_fit, rkey_sr = jr.split(jr.key(0), 5)
    X_train, Y_train = make_data(rkey_train, 2000, resolution=64)
    X_val, Y_val = make_data(rkey_val, 400, resolution=64)

    model = yax.models.FNO_nd(dim_in=1, dim_out=1, dim_hidden=32, nb_modes=16, nb_layers=3,
                   rkey=rkey_model, nb_dims=1)

    config = yax.configs.TrainConfig(learning_rate=3e-3, nb_epochs=60, lr_final_ratio=0.01)
    run = yax.training.train(os.path.join(yax.training.OUT_FOLDER, "fno_chaleur"), config, yax.obj.mse, model,
                yax.training.DatasetSampler(X_train, Y_train, 64), (X_val, Y_val), rkey=rkey_fit, verbose=True)

    print(f"erreur L2 relative à 64 points  : "
          f"{erreur_relative(run.trained_model, X_val, Y_val):.4f}")

    # SUPER-RESOLUTION : d'autres fonctions, une grille deux fois plus fine,
    # le meme modele — zéro réentraînement
    X_sr, Y_sr = make_data(rkey_sr, 400, resolution=128)
    print(f"erreur L2 relative à 128 points : "
          f"{erreur_relative(run.trained_model, X_sr, Y_sr):.4f}  (sans réentraînement)")

gnn_demo_nodes.py

Classification de nœuds sur UN graphe fixe, par message passing.

Code
"""Classification de nœuds sur UN graphe fixe, par message passing.

Un graphe à deux communautés (modèle à blocs stochastiques) : les liens sont
denses à l'intérieur d'une communauté, rares entre les deux. Chaque nœud porte
une caractéristique très bruitée de sa communauté — si bruitée qu'un nœud seul
ne dit presque rien : l'optimum théorique nœud par nœud est ~69 %, et le MLP,
qui sur-apprend ses nœuds d'entraînement dans les 7 dimensions de bruit pur,
fait pire encore. Le message passing moyenne l'information sur le voisinage,
qui est surtout de la même communauté : le bruit s'écrase, le GNN dépasse
90 %. C'est l'argument du chapitre en une expérience.

La moitié des nœuds est étiquetée, la justesse est mesurée sur les autres.
Boucle d'entraînement full-batch écrite à la main (comme mlp_demo) : un seul
graphe, pas de mini-batchs — le Trainer n'a rien à mélanger ici.
"""

import jax
import jax.numpy as jnp
import jax.random as jr
import numpy as np
import optax

import yax

NB_PER_BLOCK = 60
DIM_FEATURES = 8
SIGMA = 2.0          # écart-type du bruit sur les caractéristiques
NB_LABELED = 30      # nœuds étiquetés par communauté (la moitié)


def make_graph(seed=0):
    """Rend (features, labels, senders, receivers, train_idx, test_idx).
    Génération en numpy : les arêtes sont des tableaux concrets, construits une
    fois pour toutes hors de tout jit — le graphe est FIXE."""
    rng = np.random.default_rng(seed)
    nb_nodes = 2 * NB_PER_BLOCK
    labels = np.repeat([0, 1], NB_PER_BLOCK)

    # liens : p_in dans une communauté, p_out entre communautés
    p = np.where(labels[:, None] == labels[None, :], 0.2, 0.02)
    upper = np.triu(rng.random((nb_nodes, nb_nodes)) < p, k=1)
    i, j = np.nonzero(upper)
    senders = np.concatenate([i, j])      # un lien non orienté = deux arêtes
    receivers = np.concatenate([j, i])

    # caractéristiques : +/-1 sur la première coordonnée, noyé dans le bruit
    mu = np.zeros((2, DIM_FEATURES))
    mu[0, 0], mu[1, 0] = 1.0, -1.0
    features = mu[labels] + SIGMA * rng.normal(size=(nb_nodes, DIM_FEATURES))

    labeled = np.concatenate([rng.choice(np.arange(b * NB_PER_BLOCK, (b + 1) * NB_PER_BLOCK),
                                         NB_LABELED, replace=False) for b in range(2)])
    test_idx = np.setdiff1d(np.arange(nb_nodes), labeled)

    to_j = lambda a, dtype: jnp.asarray(a, dtype=dtype)
    return (to_j(features, jnp.float32), to_j(labels, jnp.int32),
            to_j(senders, jnp.int32), to_j(receivers, jnp.int32),
            to_j(labeled, jnp.int32), to_j(test_idx, jnp.int32))


class GNN_node_classifier(yax.Module):
    mp1: yax.layers.MessagePassing_layer
    mp2: yax.layers.MessagePassing_layer
    head: yax.layers.Linear

    def __init__(self, rkey, *, aggregation="mean"):
        rkey1, rkey2, rkey3 = jr.split(rkey, 3)
        # "mean" est un bon defaut ici : l'échelle de l'agrégat ne dépend pas
        # du degré, qui varie beaucoup d'un nœud à l'autre dans ce graphe —
        # mais run_comparison essaie les quatre agrégations
        self.mp1 = yax.layers.MessagePassing_layer(DIM_FEATURES, 16, 16, rkey1,
                                        aggregation=aggregation)
        self.mp2 = yax.layers.MessagePassing_layer(16, 16, 16, rkey2,
                                        aggregation=aggregation)
        self.head = yax.layers.Linear(16, 2, rkey3)

    def apply(self, h, rkey=None, *, senders, receivers):
        h = self.mp1.apply(h, senders=senders, receivers=receivers)
        h = self.mp2.apply(h, senders=senders, receivers=receivers)
        return self.head.apply(h)  # logits (nb_noeuds, 2)


def train_nodes(logits_fn, model, labels, train_idx, nb_steps=300, lr=0.01):
    """logits_fn(model) -> (nb_noeuds, 2). Perte : CE sur les seuls nœuds
    étiquetés — l'indexation par entiers reste valide sous jit."""
    optimizer = optax.adam(lr)
    opt_state = optimizer.init(model)

    def loss_fn(model):
        logits = logits_fn(model)
        ce = optax.softmax_cross_entropy_with_integer_labels(
            logits[train_idx], labels[train_idx])
        return jnp.mean(ce)

    @jax.jit
    def step(model, opt_state):
        loss, grads = jax.value_and_grad(loss_fn)(model)
        updates, opt_state = optimizer.update(grads, opt_state)
        return optax.apply_updates(model, updates), opt_state, loss

    for _ in range(nb_steps):
        model, opt_state, loss = step(model, opt_state)
    return model, float(loss)


def accuracy_on(logits_fn, model, labels, idx):
    pred = jnp.argmax(logits_fn(model)[idx], axis=-1)
    return float(jnp.mean(pred == labels[idx]))


def run_comparison(seed=0, verbose=True):
    X, labels, senders, receivers, train_idx, test_idx = make_graph(seed)
    rkey_gnn, rkey_mlp = jr.split(jr.key(seed))

    # référence : un MLP nœud par nœud, aveugle au graphe
    mlp = yax.layers.MLP((DIM_FEATURES, 32, 2), "relu", rkey_mlp)
    mlp_logits = lambda m: m.apply(X)
    mlp, _ = train_nodes(mlp_logits, mlp, labels, train_idx)
    acc_mlp = accuracy_on(mlp_logits, mlp, labels, test_idx)

    # le GNN : mêmes caractéristiques, plus la structure du graphe —
    # avec chacune des quatre agrégations
    if verbose:
        print(f"MLP (sans le graphe) : {acc_mlp:.3f}")
    acc_gnn = {}
    gnn_logits = lambda m: m.apply(X, senders=senders, receivers=receivers)
    for aggregation in ("sum", "mean", "max", "attention"):
        gnn = GNN_node_classifier(rkey_gnn, aggregation=aggregation)
        gnn, _ = train_nodes(gnn_logits, gnn, labels, train_idx)
        acc_gnn[aggregation] = accuracy_on(gnn_logits, gnn, labels, test_idx)
        if verbose:
            print(f"GNN, agrégation {aggregation:9s} : {acc_gnn[aggregation]:.3f}")
    return acc_mlp, acc_gnn


if __name__ == "__main__":
    run_comparison()

mlp_demo.py

Approximation de sinus par un MLP.

Code
"""Approximation de sinus par un MLP.

Descente de gradient a la main : pas d'optax, pas de filter_grad, pas de
partition/combine. Tout repose sur le fait qu'un yax.Module est un pytree
dont les feuilles sont exactement les parametres, a condition d'avoir
declare en `static` tout ce qui n'est pas un tableau.
"""

import jax
import jax.numpy as jnp
import jax.random as jr
import matplotlib.pyplot as plt
import yax


def make_data(rkey, nb, x_min, x_max):
    x = jr.uniform(rkey, (nb, 1), minval=x_min, maxval=x_max)
    y = jnp.sin(x)
    return x, y


def loss_fn(model, x, y):
    y_pred = jax.vmap(model.apply)(x)
    return jnp.mean((y_pred - y) ** 2)


@jax.jit
def step(model, x, y, lr):
    loss, grads = jax.value_and_grad(loss_fn)(model, x, y)
    new_model = jax.tree.map(lambda p, g: p - lr * g, model, grads)
    return new_model, loss


def train(model, x, y, lr, nb_steps):
    losses = []
    for i in range(nb_steps):
        model, loss = step(model, x, y, lr)
        losses.append(loss)
        if i % (nb_steps // 10) == 0:
            print(f"step {i:5d}   loss = {loss:.3e}")
    return model, jnp.array(losses)


def plot(model, x_train, y_train, losses, x_min, x_max):
    x_plot = jnp.linspace(x_min, x_max, 400)[:, None]
    y_plot = jax.vmap(model.apply)(x_plot)
    fig, (ax0, ax1) = plt.subplots(1, 2, figsize=(11, 4))

    ax0.plot(x_train[:, 0], y_train[:, 0], ".", ms=4, alpha=0.35,
             color="tab:gray", label="donnees")
    ax0.plot(x_plot[:, 0], jnp.sin(x_plot[:, 0]), "--", lw=2,
             color="tab:blue", label="sin(x)")
    ax0.plot(x_plot[:, 0], y_plot[:, 0], lw=2,
             color="tab:red", label="MLP")
    ax0.set_xlabel("x")
    ax0.set_title("approximation")
    ax0.legend()

    ax1.semilogy(losses, lw=1, color="tab:red")
    ax1.set_xlabel("step")
    ax1.set_title("loss (MSE)")
    ax1.grid(alpha=0.3)

    fig.tight_layout()
    return fig


if __name__ == "__main__":
    x_min, x_max = -3.0, 3.0
    lr = 0.02
    nb_steps = 10000
    rkey_data, rkey_model = jr.split(jr.PRNGKey(0))
    x, y = make_data(rkey_data, 256, x_min, x_max)
    model = yax.layers.MLP((1, 32, 32, 1), "tanh", rkey_model)

    print("loss initiale :", loss_fn(model, x, y))
    model, losses = train(model, x, y, lr, nb_steps)
    print("loss finale   :", losses[-1])

    fig = plot(model, x, y, losses, x_min, x_max)
    fig.savefig("mpl_demo.png", dpi=120)
    plt.show()

realnvp_demo_generation.py

Génération 2D par flot normalisant (RealNVP) sur les deux lunes.

Code
"""Génération 2D par flot normalisant (RealNVP) sur les deux lunes.

Le cas le plus confortable pour le Trainer : la loss est un
-log-vraisemblance EXACT et déterministe, aucun rkey nulle part. On passe
Y = X (ignoré par la loss) et tout le reste est le flux standard.
"""

import os

import jax.random as jr


from generative_common import make_moons, chamfer
import yax

if __name__ == "__main__":
    rkey_train, rkey_val, rkey_model, rkey_fit, rkey_gen = jr.split(jr.key(0), 5)
    X_train = make_moons(rkey_train, 2000)
    X_val = make_moons(rkey_val, 500)

    model = yax.models.RealNVP(dim=2, dim_hidden=64, nb_couplings=6, rkey=rkey_model)
    avant = chamfer(model.sample(rkey_gen, 500), X_val)

    config = yax.configs.TrainConfig(learning_rate=2e-3, nb_epochs=150, lr_final_ratio=0.01)
    run = yax.training.train(
        os.path.join(yax.training.OUT_FOLDER, "realnvp_moons"), config, yax.models.realnvp_loss, model,
        yax.training.DatasetSampler(X_train, X_train, 128), (X_val, X_val),   # non supervisé : Y = X
        rkey=rkey_fit, verbose=True)

    apres = chamfer(run.trained_model.sample(rkey_gen, 500), X_val)
    print(f"NLL de validation : {run.loss:.3f} nats")
    print(f"chamfer générés/données : {avant:.3f} avant, {apres:.3f} après")

ritz_demo_poisson.py

Deep Ritz : résoudre -u'' = f sur [0, 1], u(0) = u(1) = 0, SANS données.

Code
"""Deep Ritz : résoudre -u'' = f sur [0, 1], u(0) = u(1) = 0, SANS données.

On minimise l'énergie J[u] = ∫ ½ u'² - f u, dont le minimiseur est la
solution du problème. Les points d'intégration sont TIRÉS à chaque batch
(yax.training.FunctionSampler), la validation est une grille fixe (yax.training.FunctionSampler appelé
avec la clé constante du Trainer). y vaut None : la loss ne compare à rien.

Avec f = π² sin(πx) : u = sin(πx) et J[u] = -π²/4.
"""

import os

import jax
import jax.numpy as jnp
import jax.random as jr
import matplotlib.pyplot as plt

import yax


class Ritz1D(yax.Module):
    """u(x) = x(1-x)·net(x) : les conditions aux limites sont dans l'ansatz,
    il n'y a rien à pénaliser."""
    net: yax.layers.MLP

    def __init__(self, rkey):
        self.net = yax.layers.MLP((1, 32, 32, 1), "tanh", rkey)

    def apply(self, x, rkey=None):   # x scalaire -> u(x) scalaire
        return x * (1 - x) * self.net.apply(x[None])[0]


def f_source(x):
    return jnp.pi ** 2 * jnp.sin(jnp.pi * x)


def ritz_loss(model, x, y, rkey=None):
    """x : (n,) points uniformes sur [0,1] ⇒ la moyenne est l'intégrale."""
    u, du = jax.vmap(jax.value_and_grad(model.apply))(x)
    return jnp.mean(0.5 * du ** 2 - f_source(x) * u)


def main():
    sampler = yax.training.FunctionSampler(lambda rkey: (jr.uniform(rkey, (512,)), None), nb_batches=25)
    validation = yax.training.FunctionSampler(lambda rkey: (jnp.linspace(0.0, 1.0, 1025), None), nb_batches=1)
    config = yax.configs.TrainConfig(learning_rate=5e-3, nb_epochs=60, lr_final_ratio=0.02)

    run = yax.training.train(os.path.join(yax.training.OUT_FOLDER, "ritz_poisson"), config, ritz_loss,
                Ritz1D(jr.key(0)), sampler, validation, rkey=jr.key(1), verbose=True)

    x = jnp.linspace(0.0, 1.0, 400)
    u = jax.vmap(run.trained_model.apply)(x)
    exact = jnp.sin(jnp.pi * x)
    print(f"énergie de validation : {run.loss:.5f}   (exact -π²/4 = {-jnp.pi ** 2 / 4:.5f})")
    print(f"erreur L2 : {float(jnp.sqrt(jnp.mean((u - exact) ** 2))):.2e}")

    fig, (ax0, ax1) = plt.subplots(1, 2, figsize=(11, 4))
    ax0.plot(x, exact, "--", lw=2, color="tab:blue", label="sin(πx)")
    ax0.plot(x, u, lw=2, color="tab:red", label="Ritz")
    ax0.set_title("solution")
    ax0.legend()
    run.history.plot(ax=ax1, log=False, title="énergie J[u] (entraînement par pas, validation sur grille)")
    fig.tight_layout()
    fig.savefig("ritz_demo.png", dpi=120)
    plt.show()


if __name__ == "__main__":
    main()

rnn_demo_classif.py

Code
import os
import jax.numpy as jnp
import jax.random as jr
from matplotlib import pyplot as plt
import jax
import optax
import yax


#C'est l'exemple qui est proposé dans la doc d'equinox
def get_data(dataset_size:int, rkey):
    t = jnp.linspace(0, 2 * jnp.pi, 16)
    offset = jr.uniform(rkey, (dataset_size, 1), minval=0, maxval=2 * jnp.pi)
    x1 = jnp.sin(t + offset) / (1 + t)
    x2 = jnp.cos(t + offset) / (1 + t)

    half_dataset_size = dataset_size // 2
    x1 = x1.at[:half_dataset_size].multiply(-1)
    y = jnp.ones((dataset_size, 1))
    y = y.at[:half_dataset_size].set(0)
    #axis=-1 : on veut (dataset_size, 16, 2), soit 16 pas de temps de 2 features.
    #Avec axis=1 on obtiendrait (dataset_size, 2, 16), et le scan parcourrait
    #2 pas de temps de 16 features.
    x = jnp.stack([x1, x2], axis=-1)

    #les 2 classes sont rangees en bloc, mais batchs_for_one_epoch melange a
    #chaque epoque, et train/val proviennent de deux appels distincts, chacun
    #equilibre. Ne PAS decouper un seul get_data en tranches train/val.
    return x, y


def present_data():
    n_data = 10
    x, y = get_data(n_data,jr.key(0))
    print(x.shape, y.shape)

    fig, axs = plt.subplots(n_data, 2, figsize=(8, 1.2 * n_data), sharex="all", sharey="all")
    for i in range(n_data):
        axs[i, 0].plot(x[i, :, 0])
        axs[i, 1].plot(x[i, :, 1])
        axs[i, 0].set_ylabel(f"label={int(y[i, 0])}")

    axs[0, 0].set_title("feature 0")
    axs[0, 1].set_title("feature 1")
    plt.show()


class RNN_bin_classifier(yax.Module):
    rnn_layer: yax.layers.RNN_layer
    linear_layer: yax.layers.Linear

    dim_hidden:int = yax.StaticField()

    def __init__(self,dim_hidden,rkey,*,cell_type="gru"):
        self.dim_hidden=dim_hidden

        rk1,rk2=jr.split(rkey)
        self.rnn_layer = yax.layers.RNN_layer(2, dim_hidden, rk1, cell_type=cell_type)
        self.linear_layer=yax.layers.Linear(dim_hidden,1,rk2)

    def apply(self,xs,rkey=None):
        # rkey ignoré : modèle déterministe (signature uniforme apply(x, rkey))
        y=self.rnn_layer.apply(xs)[-1]   # le dernier état résume la séquence
        y=self.linear_layer.apply(y)
        return y #pas de sigmoid: il est ajouté dans la loss


def accuracy(model, x, y):
    logits = yax.batch_apply(model, x, None)
    return float(jnp.mean((logits > 0) == (y > 0.5)))   #logit>0 <=> sigmoid>0.5


def train_classif(cell_type="gru", dim_hidden=16, verbose=True):
    rkey_train, rkey_val, rkey_model, rkey_shuffle = jr.split(jr.key(0), 4)
    X_train, Y_train = get_data(512, rkey_train)
    X_val, Y_val = get_data(256, rkey_val)

    model = RNN_bin_classifier(dim_hidden, rkey_model, cell_type=cell_type)
    acc_before = accuracy(model, X_val, Y_val)

    config = yax.configs.TrainConfig(learning_rate=0.05, nb_epochs=60, lr_final_ratio=0.01)
    mother_folder = os.path.join(yax.training.OUT_FOLDER, f"rnn_classif_{cell_type}")
    run = yax.training.train(
        mother_folder, config, yax.obj.bce, model, yax.training.DatasetSampler(X_train, Y_train, 32), (X_val, Y_val),
        rkey=rkey_shuffle, verbose=verbose)

    acc_after = accuracy(run.trained_model, X_val, Y_val)
    print(f"{cell_type} : accuracy {acc_before:.3f} -> {acc_after:.3f}"
          f"   (best val BCE = {run.loss:.3g})")
    return run, acc_after


def train_and_plot(cell_type="gru"):
    run, acc = train_classif(cell_type)
    history = yax.training.load_run(run.folder, "history").history
    history.plot(title=f"RNN {cell_type} : BCE={run.loss:.2g}, accuracy={acc:.3f}")
    plt.tight_layout()
    plt.show()


if __name__ == "__main__":
    train_and_plot("gru")

transformer_demo_charlm.py

Petit transformer prédicteur (décodeur seul) : génération de texte

Code
"""Petit transformer prédicteur (décodeur seul) : génération de texte
caractère par caractère sur un corpus jouet.

La modélisation causale : prédire le jeton suivant — c'est la prévision
auto-régressive du chapitre 10, sur du texte. Le corpus tient en quelques
proverbes ; le modèle le mémorise en quelques secondes, ce qui suffit à voir
tourner toute la mécanique (embedding + positions + blocs causaux + Trainer),
et la génération auto-régressive qui réinjecte sa propre sortie.

La « tokenisation » est ici la plus simple possible : un caractère = un jeton.
Le vocabulaire est un choix, pas une donnée (chapitre 12).
"""

import os

import jax.numpy as jnp
import jax.random as jr

import yax

CORPUS = (
    "la nuit porte conseil. qui vole un oeuf vole un boeuf. rien ne sert de "
    "courir, il faut partir a point. petit a petit, l'oiseau fait son nid. "
    "c'est en forgeant qu'on devient forgeron. apres la pluie, le beau temps. "
)
SEQ_LEN = 32

# vocabulaire : les caractères du corpus, dans un ordre fixe
VOCAB = sorted(set(CORPUS))
CHAR_TO_ID = {c: i for i, c in enumerate(VOCAB)}


def encode(text):
    return jnp.array([CHAR_TO_ID[c] for c in text], dtype=jnp.int32)


def decode(ids):
    return "".join(VOCAB[int(i)] for i in ids)


def make_data(rkey):
    """Fenêtres glissantes : X = corpus[i:i+L], Y = la même fenêtre décalée
    d'un caractère. La cible est l'entrée décalée d'un pas (chapitre 10)."""
    codes = encode(CORPUS)
    nb = len(codes) - SEQ_LEN
    X = jnp.stack([codes[i:i + SEQ_LEN] for i in range(nb)])
    Y = jnp.stack([codes[i + 1:i + 1 + SEQ_LEN] for i in range(nb)])

    shuffle = jr.permutation(rkey, nb)
    X, Y = X[shuffle], Y[shuffle]
    nb_val = nb // 5
    return X[nb_val:], Y[nb_val:], X[:nb_val], Y[:nb_val]


class TransformerLM(yax.Module):
    token_emb: yax.layers.Embedding
    blocks: list
    final_norm: yax.layers.LayerNorm
    head: yax.layers.Linear

    dim: int = yax.StaticField()

    def __init__(self, vocab_size, dim, nb_heads, dim_ff, nb_blocks, rkey, *,
                 dropout_rate=0.0):
        self.dim = dim
        rkey_emb, rkey_head, *rkey_blocks = jr.split(rkey, 2 + nb_blocks)
        self.token_emb = yax.layers.Embedding(vocab_size, dim, rkey_emb)
        self.blocks = [yax.layers.TransformerBlock(dim, nb_heads, dim_ff, rkey,
                                        dropout_rate=dropout_rate)
                       for rkey in rkey_blocks]
        self.final_norm = yax.layers.LayerNorm(dim)   # norm finale : le pendant du pré-norme
        self.head = yax.layers.Linear(dim, vocab_size, rkey_head)

    def apply(self, ids, rkey=None):
        """ids : (seq,) entiers -> logits (seq, vocab) : à chaque position, la
        distribution du caractère SUIVANT. Le masque causal garantit que la
        position t ne voit que ids[:t+1]."""
        seq_len = ids.shape[0]
        # les positions sinusoïdales sont recalculées ici : fonction de la seule
        # forme (statique), le compilateur les constant-fold — rien d'appris,
        # donc rien qui doive être une feuille du modèle
        x = self.token_emb.apply(ids) + yax.preprocessing.sinusoidal_positional_encoding(seq_len, self.dim)
        mask = yax.preprocessing.causal_mask(seq_len)
        rkeys = (None,) * len(self.blocks) if rkey is None else jr.split(rkey, len(self.blocks))
        for block, rkey_bloc in zip(self.blocks, rkeys):
            x = block.apply(x, rkey_bloc, mask=mask)
        return self.head.apply(self.final_norm.apply(x))


def generate(model, prompt, nb_chars):
    """Génération auto-régressive : on réinjecte sa propre sortie (chapitre 10),
    en gardant les SEQ_LEN derniers caractères comme contexte. Choix glouton —
    pour un modèle qui a mémorisé son corpus, échantillonner n'apporte rien."""
    ids = [int(i) for i in encode(prompt)]
    for _ in range(nb_chars):
        window = jnp.array(ids[-SEQ_LEN:], dtype=jnp.int32)
        logits = model.apply(window)
        ids.append(int(jnp.argmax(logits[-1])))
    return decode(ids)


def train_charlm(verbose=True):
    rkey_data, rkey_model, rkey_train = jr.split(jr.key(0), 3)
    X_train, Y_train, X_val, Y_val = make_data(rkey_data)

    # le dropout n'est pas décoratif : sans lui le modèle mémorise les fenêtres
    # AVEC leur alignement de position (CE train ~0.07, CE val ~0.3, et la
    # génération bégaie « son n nid ») ; avec lui la structure apprise devient
    # indépendante de l'alignement et la génération est exacte. Bel exemple de
    # sur-apprentissage réglé par la régularisation (chapitre 7).
    model = TransformerLM(vocab_size=len(VOCAB), dim=32, nb_heads=4,
                          dim_ff=64, nb_blocks=2, rkey=rkey_model,
                          dropout_rate=0.1)

    config = yax.configs.TrainConfig(learning_rate=0.003, nb_epochs=300, lr_final_ratio=0.01)
    mother_folder = os.path.join(yax.training.OUT_FOLDER, "charlm")
    run = yax.training.train(
        mother_folder, config, yax.obj.softmax_ce_from_integers, model,
        yax.training.DatasetSampler(X_train, Y_train, 32), (X_val, Y_val),
        rkey=rkey_train, verbose=verbose)

    print(f"\nbest val CE = {run.loss:.3g}")
    print("génération :", generate(run.trained_model, "petit a petit", 60))
    return run


if __name__ == "__main__":
    from matplotlib import pyplot as plt

    run = train_charlm()
    history = yax.training.load_run(run.folder, "history").history
    history.plot(title=f"transformer char-LM : CE={run.loss:.2g}")
    plt.tight_layout()
    plt.show()

unet_demo_segmentation.py

Segmentation binaire synthétique par U-Net.

Code
"""Segmentation binaire synthétique par U-Net.

Deux disques clairs posés au hasard sur un fond bruité ; la cible est le masque
des pixels des disques. C'est le passage de « une étiquette par image » à « une
étiquette par pixel » : la sortie a la taille de l'entrée, la perte est
l'entropie croisée PAR PIXEL (yax.obj.bce, la même formule que la
classification binaire, appliquée à plus d'axes).

La métrique rapportée est le Dice : la justesse par pixel est trompeuse quand
le fond domine — prédire « tout fond » donne déjà ~85 % de pixels justes ici.
"""

import os

import jax
import jax.numpy as jnp
import jax.random as jr

import yax

IMG_SIZE = 32
NB_DISQUES = 2


def make_image(rkey):
    """UNE image (1, 32, 32) et son masque (1, 32, 32)."""
    rkey_pos, rkey_radius, rkey_noise = jr.split(rkey, 3)
    centers = jr.uniform(rkey_pos, (NB_DISQUES, 2), minval=5.0, maxval=IMG_SIZE - 5.0)
    radii = jr.uniform(rkey_radius, (NB_DISQUES,), minval=3.0, maxval=5.0)

    rows = jnp.arange(IMG_SIZE)[:, None]
    cols = jnp.arange(IMG_SIZE)[None, :]
    disques = ((rows - centers[:, 0, None, None]) ** 2
               + (cols - centers[:, 1, None, None]) ** 2) < radii[:, None, None] ** 2
    mask = jnp.any(disques, axis=0)

    img = jnp.where(mask, 1.0, 0.0) + 0.4 * jr.normal(rkey_noise, (IMG_SIZE, IMG_SIZE))
    return img[None], mask[None].astype(jnp.float32)


def make_data(rkey, nb):
    return jax.vmap(make_image)(jr.split(rkey, nb))


def dice_score(model, x, y):
    """Dice « dur » (prédiction seuillée), moyenné sur les échantillons."""
    logits = yax.batch_apply(model, x, None)
    pred = (logits > 0).astype(jnp.float32)  # logit>0 <=> sigmoid>0.5
    axes = tuple(range(1, pred.ndim))
    intersection = jnp.sum(pred * y, axis=axes)
    denom = jnp.sum(pred, axis=axes) + jnp.sum(y, axis=axes)
    return float(jnp.mean(2.0 * intersection / (denom + 1e-6)))


def present_data():
    from matplotlib import pyplot as plt

    X, Y = make_data(jr.key(0), 4)
    fig, axs = plt.subplots(2, 4, figsize=(10, 5))
    for i in range(4):
        axs[0, i].imshow(X[i, 0], cmap="gray")
        axs[1, i].imshow(Y[i, 0], cmap="gray")
        axs[0, i].axis("off")
        axs[1, i].axis("off")
    axs[0, 0].set_title("image")
    axs[1, 0].set_title("masque")
    plt.tight_layout()
    plt.show()


def train_segmentation(verbose=True):
    rkey_train, rkey_val, rkey_model, rkey_shuffle = jr.split(jr.key(0), 4)
    X_train, Y_train = make_data(rkey_train, 512)
    X_val, Y_val = make_data(rkey_val, 128)

    model = yax.models.UNet_nd(1, 8, 2, 2, rkey_model)
    dice_before = dice_score(model, X_val, Y_val)

    config = yax.configs.TrainConfig(learning_rate=0.003, nb_epochs=30, lr_final_ratio=0.01)
    mother_folder = os.path.join(yax.training.OUT_FOLDER, "unet_seg")
    run = yax.training.train(
        mother_folder, config, yax.obj.bce, model, yax.training.DatasetSampler(X_train, Y_train, 32), (X_val, Y_val),
        rkey=rkey_shuffle, verbose=verbose)

    dice_after = dice_score(run.trained_model, X_val, Y_val)
    print(f"\nDice {dice_before:.3f} -> {dice_after:.3f}   (best val BCE = {run.loss:.3g})")
    return run, dice_after


def plot_predictions(folder, nb=4):
    from matplotlib import pyplot as plt
    model = yax.training.load_run(folder, "trained_model").trained_model
    X, Y = make_data(jr.key(2), nb)
    logits = yax.batch_apply(model, X, None)
    pred = jax.nn.sigmoid(logits)

    fig, axs = plt.subplots(3, nb, figsize=(2.5 * nb, 7))
    for i in range(nb):
        axs[0, i].imshow(X[i, 0], cmap="gray")
        axs[1, i].imshow(Y[i, 0], cmap="gray")
        axs[2, i].imshow(pred[i, 0], cmap="gray", vmin=0, vmax=1)
        for row in range(3):
            axs[row, i].axis("off")
    axs[0, 0].set_title("image")
    axs[1, 0].set_title("masque vrai")
    axs[2, 0].set_title("sigmoid(logits)")
    plt.tight_layout()
    plt.show()


if __name__ == "__main__":
    run, dice = train_segmentation()
    plot_predictions(run.folder)

vae_demo_generation.py

Génération 2D par VAE sur les deux lunes.

Code
"""Génération 2D par VAE sur les deux lunes.

Non supervisé avec le Trainer TEL QUEL : on passe Y = X, et vae_loss
reconstruit y. L'aléa de la reparamétrisation suit le même chemin que le
dropout (les step_rkeys du Trainer) ; la validation, en mode inference, encode
par la moyenne — déterministe.

beta < 1 : sur des données 2D à petit bruit, le terme de reconstruction est
minuscule devant le KL — on le rééquilibre, sinon le latent s'effondre et le
VAE génère un nuage informe.
"""

import os

import jax.random as jr


from generative_common import make_moons, chamfer
import yax

if __name__ == "__main__":
    rkey_train, rkey_val, rkey_model, rkey_fit, rkey_gen = jr.split(jr.key(0), 5)
    X_train = make_moons(rkey_train, 2000)
    X_val = make_moons(rkey_val, 500)

    model = yax.models.VAE(dim_in=2, dim_hidden=64, dim_latent=2, rkey=rkey_model)
    objective_fn = yax.models.vae_loss(beta=0.05)

    avant = chamfer(model.set_inference(True).generate(rkey_gen, 500), X_val)

    config = yax.configs.TrainConfig(learning_rate=3e-3, nb_epochs=80, lr_final_ratio=0.01)
    run = yax.training.train(
        os.path.join(yax.training.OUT_FOLDER, "vae_moons"), config, objective_fn, model,
        yax.training.DatasetSampler(X_train, X_train, 64), (X_val, X_val),   # non supervisé : Y = X
        rkey=rkey_fit, verbose=True)

    apres = chamfer(run.trained_model.generate(rkey_gen, 500), X_val)
    print(f"chamfer générés/données : {avant:.3f} avant, {apres:.3f} après")

yolo_demo_detection.py

Détection d'objets par mini-YOLO sur données synthétiques.

Code
"""Détection d'objets par mini-YOLO sur données synthétiques.

Le jeu de données du sujet : des formes (disques et carrés, une seule classe
« objet ») posées aléatoirement sur un fond bruité — les boîtes sont connues
par construction, aucune annotation à faire. 1 à 3 objets par image.

La métrique est précision/rappel à IoU >= 0.5 après suppression des
non-maxima : une détection est un vrai positif si sa boîte recouvre une vraie
boîte non encore appariée.
"""

import os

import jax
import jax.numpy as jnp
import jax.random as jr
import numpy as np

import yax
from yax.models.MiniYOLO import IMG_SIZE, MAX_OBJ   # les constantes du modele


def make_scene(rkey):
    """Les boîtes d'abord, l'image ensuite : la vérité est connue par
    construction. Rend (boxes (MAX_OBJ, 4) en coins, present (MAX_OBJ,))."""
    rkey_nb, rkey_pos, rkey_radius = jr.split(rkey, 3)
    # le premier objet est toujours présent, les suivants avec probabilité 0.6
    present = jnp.concatenate([jnp.array([True]),
                               jr.bernoulli(rkey_nb, 0.6, (MAX_OBJ - 1,))])
    centers = jr.uniform(rkey_pos, (MAX_OBJ, 2), minval=6.0, maxval=IMG_SIZE - 6.0)
    radii = jr.uniform(rkey_radius, (MAX_OBJ,), minval=3.0, maxval=5.0)
    boxes = jnp.stack([centers[:, 0] - radii, centers[:, 1] - radii,
                       centers[:, 0] + radii, centers[:, 1] + radii], axis=-1)
    return boxes, present


def make_image(rkey, boxes, present):
    """L'image (1, 32, 32) : disque ou carré dans chaque boîte présente."""
    rkey_shape, rkey_noise = jr.split(rkey)
    is_disk = jr.bernoulli(rkey_shape, 0.5, (MAX_OBJ,))
    rows = jnp.arange(IMG_SIZE)[:, None]
    cols = jnp.arange(IMG_SIZE)[None, :]

    img = jnp.zeros((IMG_SIZE, IMG_SIZE))
    for i in range(MAX_OBJ):  # MAX_OBJ est statique, la boucle est déroulée
        x1, y1, x2, y2 = boxes[i]
        cx, cy, r = (x1 + x2) / 2, (y1 + y2) / 2, (x2 - x1) / 2
        disk = (rows - cy) ** 2 + (cols - cx) ** 2 < r ** 2
        square = jnp.maximum(jnp.abs(rows - cy), jnp.abs(cols - cx)) < r
        shape = jnp.where(is_disk[i], disk, square) & present[i]
        img = jnp.maximum(img, shape.astype(jnp.float32))

    img = img + 0.3 * jr.normal(rkey_noise, (IMG_SIZE, IMG_SIZE))
    return img[None]


def make_data(rkey, nb):
    """Rend (X, Y_grilles, boxes, present) — les grilles pour le Trainer, les
    boîtes brutes pour l'évaluation."""
    rkey_scenes, rkey_images = jr.split(rkey)
    boxes, present = jax.vmap(make_scene)(jr.split(rkey_scenes, nb))
    X = jax.vmap(make_image)(jr.split(rkey_images, nb), boxes, present)
    Y = jax.vmap(yax.models.encode_targets)(boxes, present)
    return X, Y, boxes, present


def precision_recall(model, X, boxes, present, iou_threshold=0.5):
    tp_total = fp_total = fn_total = 0
    for i in range(len(X)):
        pred_boxes, _ = yax.models.detect(model, X[i])
        true_boxes = np.asarray(boxes[i])[np.asarray(present[i])]
        tp, fp, fn = yax.models.match_detections(pred_boxes, true_boxes,
                                      iou_threshold=iou_threshold)
        tp_total += tp
        fp_total += fp
        fn_total += fn
    precision = tp_total / max(tp_total + fp_total, 1)
    recall = tp_total / max(tp_total + fn_total, 1)
    return precision, recall


def train_detection(verbose=True):
    rkey_train, rkey_val, rkey_model, rkey_shuffle = jr.split(jr.key(0), 4)
    X_train, Y_train, _, _ = make_data(rkey_train, 2000)
    X_val, Y_val, boxes_val, present_val = make_data(rkey_val, 200)

    model = yax.models.MiniYOLO(rkey_model)
    config = yax.configs.TrainConfig(learning_rate=0.003, nb_epochs=100, lr_final_ratio=0.01)
    mother_folder = os.path.join(yax.training.OUT_FOLDER, "yolo_detection")
    run = yax.training.train(
        mother_folder, config, yax.models.yolo_loss, model,
        yax.training.DatasetSampler(X_train, Y_train, 64), (X_val, Y_val),
        rkey=rkey_shuffle, verbose=verbose)

    precision, recall = precision_recall(run.trained_model, X_val, boxes_val, present_val)
    print(f"\nprécision {precision:.3f}, rappel {recall:.3f} (IoU >= 0.5)"
          f"   (best val loss = {run.loss:.3g})")
    return run, precision, recall


def plot_detections(folder, nb=6):
    from matplotlib import pyplot as plt
    from matplotlib.patches import Rectangle

    model = yax.training.load_run(folder, "trained_model").trained_model
    X, _, boxes, present = make_data(jr.key(2), nb)

    fig, axs = plt.subplots(1, nb, figsize=(2.6 * nb, 3))
    for i in range(nb):
        axs[i].imshow(X[i, 0], cmap="gray")
        for box in np.asarray(boxes[i])[np.asarray(present[i])]:
            axs[i].add_patch(Rectangle((box[0], box[1]), box[2] - box[0],
                                       box[3] - box[1], fill=False,
                                       edgecolor="lime", lw=2))
        pred_boxes, scores = yax.models.detect(model, X[i])
        for box, score in zip(pred_boxes, scores):
            axs[i].add_patch(Rectangle((box[0], box[1]), box[2] - box[0],
                                       box[3] - box[1], fill=False,
                                       edgecolor="red", lw=1.5, ls="--"))
            axs[i].text(box[0], box[1] - 1, f"{score:.2f}", color="red", fontsize=8)
        axs[i].axis("off")
    axs[0].set_title("vert : vérité — rouge : détections")
    plt.tight_layout()
    plt.show()


if __name__ == "__main__":
    run, precision, recall = train_detection()
    plot_detections(run.folder)