Aller au contenu

yax.models

Des modèles complets, prêts à entraîner, avec leur fonction de perte lorsqu'elle leur est propre.

Opérateur de Fourier

Opérateur neuronal de Fourier (FNO, Li et al., 2021), en dimension quelconque.

Un opérateur neuronal apprend une transformation entre fonctions : un champ d'entrée vers un champ de sortie. Ses poids agissent sur les fréquences du champ et non sur ses points : un modèle entraîné sur une grille de 64 points s'évalue tel quel sur une grille de 128.

Les entrées sont de forme (canaux, *grille). Le modèle est équivariant par translation : décaler l'entrée décale la sortie. S'il doit connaître la position, ajouter les coordonnées de la grille aux canaux d'entrée.

FNO_nd

yax.models.FNO_nd(dim_in, dim_out, dim_hidden, nb_modes, nb_layers, nb_dims, rkey, *, activation='gelu')

Bases: Module

Opérateur neuronal de Fourier.

Relève les canaux d'entrée à dim_hidden, applique nb_layers blocs activation(spectrale(x) + linéaire(x)), puis projette vers dim_out.

fno = yax.models.FNO_nd(dim_in=1, dim_out=1, dim_hidden=32, nb_modes=16, nb_layers=3,
                        nb_dims=1, rkey=rkey)
u = fno.apply(u0)          # u0 (1, 64) -> u (1, 64) ; fonctionne aussi en (1, 128)

Paramètres :

Nom Type Description Défaut
dim_in int

canaux d'entrée.

obligatoire
dim_out int

canaux de sortie.

obligatoire
dim_hidden int

nombre de canaux internes.

obligatoire
nb_modes int | tuple[int, ...]

nombre de fréquences gardées par axe.

obligatoire
nb_layers int

nombre de blocs spectraux.

obligatoire
nb_dims int

nombre de dimensions de la grille.

obligatoire
rkey Array

clé aléatoire pour l'initialisation.

obligatoire
activation str | Callable

activation entre les blocs.

'gelu'

Attributs :

Nom Type Description
lift Conv_nd

convolution 1×1 vers dim_hidden canaux.

spectrals list[SpectralConv_nd]

les couches spectrales.

bypasses list[Conv_nd]

les convolutions 1×1 parallèles aux couches spectrales.

proj Conv_nd

convolution 1×1 vers dim_out canaux.

activation_fn Callable

la fonction d'activation.

Code source dans yax/models/FNO_nd.py
92
class FNO_nd(Module):
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
    lift: Conv_nd                  # 1x1 : canaux d'entrée -> dim_hidden
    spectrals: list[SpectralConv_nd]
    bypasses: list[Conv_nd]        # 1x1, en parallèle de chaque couche spectrale
    proj: Conv_nd                  # 1x1 : dim_hidden -> canaux de sortie
    activation_fn: Callable = StaticField()

    def __init__(self, dim_in: int, dim_out: int, dim_hidden: int,
                 nb_modes: int | tuple[int, ...], nb_layers: int, nb_dims: int,
                 rkey: jax.Array, *, activation: str | Callable = "gelu"):
        rkeys = iter(jr.split(rkey, 2 * nb_layers + 2))
        self.lift = Conv_nd(dim_in, dim_hidden, 1, nb_dims, next(rkeys))
        self.spectrals = [SpectralConv_nd(dim_hidden, dim_hidden, nb_modes, nb_dims,
                                          next(rkeys))
                          for _ in range(nb_layers)]
        self.bypasses = [Conv_nd(dim_hidden, dim_hidden, 1, nb_dims, next(rkeys))
                         for _ in range(nb_layers)]
        self.proj = Conv_nd(dim_hidden, dim_out, 1, nb_dims, next(rkeys))
        self.activation_fn = resolve_activation(activation)

apply

apply(x, rkey=None)

Calcule le champ de sortie (dim_out, *grille) pour un champ x de forme (dim_in, *grille).

Code source dans yax/models/FNO_nd.py
140
141
142
143
144
145
146
147
def apply(self, x: jax.Array, rkey: jax.Array | None = None) -> jax.Array:
    """Calcule le champ de sortie `(dim_out, *grille)` pour un champ `x` de forme `(dim_in, *grille)`."""
    # x : (dim_in, *spatial) -> (dim_out, *spatial), à N'IMPORTE quelle
    # résolution compatible avec nb_modes
    x = self.lift.apply(x)
    for spectral, bypass in zip(self.spectrals, self.bypasses):
        x = self.activation_fn(spectral.apply(x) + bypass.apply(x))
    return self.proj.apply(x)

SpectralConv_nd

yax.models.SpectralConv_nd(dim_in, dim_out, nb_modes, nb_dims, rkey)

Bases: Module

Couche spectrale : un filtre appris dans l'espace des fréquences.

Calcule la transformée de Fourier du champ, garde les nb_modes plus basses fréquences de chaque axe, les multiplie par des poids complexes appris qui mélangent aussi les canaux, puis revient dans l'espace des points.

Paramètres :

Nom Type Description Défaut
dim_in int

canaux d'entrée.

obligatoire
dim_out int

canaux de sortie.

obligatoire
nb_modes int | tuple[int, ...]

nombre de fréquences gardées par axe, un entier ou un tuple ; il faut 2 * nb_modes au plus égal à la taille de la grille.

obligatoire
nb_dims int

nombre de dimensions de la grille.

obligatoire
rkey Array

clé aléatoire pour l'initialisation.

obligatoire

Attributs :

Nom Type Description
weights list[Array]

les poids spectraux, parties réelle et imaginaire séparées.

nb_modes tuple[int, ...]

fréquences gardées par axe.

Code source dans yax/models/FNO_nd.py
25
class SpectralConv_nd(Module):
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
    weights: list[jax.Array]       # un bloc (dim_in, dim_out, *nb_modes, 2) par coin
    nb_modes: tuple[int, ...] = StaticField()

    def __init__(self, dim_in: int, dim_out: int, nb_modes: int | tuple[int, ...],
                 nb_dims: int, rkey: jax.Array):
        # nb_modes : nombre de fréquences conservées par dimension (int, ou tuple
        # de longueur nb_dims), à comparer à la taille de grille N :
        # il faut 2*nb_modes <= N (et nb_modes <= N//2+1 sur la dernière dimension).
        if isinstance(nb_modes, int):
            nb_modes = (nb_modes,) * nb_dims
        assert len(nb_modes) == nb_dims, (
            f"nb_modes:{nb_modes} incompatible avec nb_dims:{nb_dims}")
        self.nb_modes = tuple(nb_modes)
        nb_coins = 2 ** (len(self.nb_modes) - 1)
        echelle = 1.0 / (dim_in * dim_out)
        self.weights = [
            echelle * jr.normal(k, (dim_in, dim_out) + self.nb_modes + (2,))
            for k in jr.split(rkey, nb_coins)]

U-Net

U-Net en dimension quelconque, pour la segmentation.

Une branche descendante réduit la résolution et gagne en contexte ; une branche montante retrouve la résolution ; des connexions de saut transmettent à chaque niveau les détails de la descente à la remontée. Le modèle traite des signaux (1D), des images (2D) ou des volumes (3D).

UNet_nd

yax.models.UNet_nd(dim_in, dim_base, nb_levels, nb_dims, rkey, *, dim_out=1)

Bases: Module

U-Net : segmentation, avec une prédiction par pixel (ou voxel).

Paramètres :

Nom Type Description Défaut
dim_in int

canaux d'entrée.

obligatoire
dim_base int

canaux au premier niveau, doublés à chaque niveau.

obligatoire
nb_levels int

nombre de niveaux de réduction.

obligatoire
nb_dims int

nombre de dimensions spatiales.

obligatoire
rkey Array

clé aléatoire pour l'initialisation.

obligatoire
dim_out int

canaux de sortie, un par classe ; 1 pour la segmentation binaire.

1

Attributs :

Nom Type Description
encoders list[ConvBlock]

les blocs de la branche descendante.

downs list[Conv_nd]

les convolutions de réduction, de pas 2.

bottleneck ConvBlock

le bloc du niveau le plus bas.

up_convs list[Conv_nd]

les convolutions qui suivent chaque agrandissement.

decoders list[ConvBlock]

les blocs de la branche montante.

head Conv_nd

la convolution finale 1×1.

nb_levels : nombre de descentes. Chaque taille spatiale doit être divisible par 2*nb_levels. Canaux doublés à chaque niveau : base, 2base...

Code source dans yax/models/UNet_nd.py
49
class UNet_nd(Module):
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
    encoders: list[ConvBlock]  # un ConvBlock par niveau de la descente
    downs: list[Conv_nd]       # convolutions stride 2 (l'alternative au pooling)
    bottleneck: ConvBlock
    up_convs: list[Conv_nd]    # convolution après chaque interpolation
    decoders: list[ConvBlock]  # un ConvBlock par niveau de la remontée (après concat)
    head: Conv_nd     # convolution 1x1 -> logits

    def __init__(self, dim_in: int, dim_base: int, nb_levels: int, nb_dims: int,
                 rkey: jax.Array, *, dim_out: int = 1):
        """nb_levels : nombre de descentes. Chaque taille spatiale doit être
        divisible par 2**nb_levels. Canaux doublés à chaque niveau : base, 2*base..."""
        chs = [dim_base * 2**i for i in range(nb_levels + 1)]  # ex 8, 16, 32
        # 2 clés par niveau à la descente (encoder + down), 2 à la remontée
        # (up_conv + decoder), plus bottleneck et head
        rkeys = iter(jr.split(rkey, 4 * nb_levels + 2))

        encoders, downs = [], []
        d_in = dim_in
        for i in range(nb_levels):
            encoders.append(ConvBlock(d_in, chs[i], nb_dims, next(rkeys)))
            downs.append(Conv_nd(chs[i], chs[i + 1], 3, nb_dims, next(rkeys),
                                 stride=2))
            d_in = chs[i + 1]
        self.encoders = encoders
        self.downs = downs

        self.bottleneck = ConvBlock(chs[nb_levels], chs[nb_levels], nb_dims,
                                    next(rkeys))

        up_convs, decoders = [], []
        for i in reversed(range(nb_levels)):
            up_convs.append(Conv_nd(chs[i + 1], chs[i], 3, nb_dims,
                                    next(rkeys)))
            # après concaténation avec le saut : chs[i] (saut) + chs[i] (remontée)
            decoders.append(ConvBlock(2 * chs[i], chs[i], nb_dims,
                                      next(rkeys)))
        self.up_convs = up_convs
        self.decoders = decoders

        self.head = Conv_nd(chs[0], dim_out, 1, nb_dims, next(rkeys))

apply

apply(x, rkey=None)

Renvoie les logits (dim_out, *spatial) pour une entrée x de forme (dim_in, *spatial).

Les tailles spatiales doivent être divisibles par 2 ** nb_levels. Pertes associés : yax.obj.bce, yax.obj.dice_loss, ou leur somme par yax.obj.weighted_sum.

Code source dans yax/models/UNet_nd.py
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
def apply(self, x: jax.Array, rkey: jax.Array | None = None) -> jax.Array:
    """Renvoie les logits `(dim_out, *spatial)` pour une entrée `x` de forme `(dim_in, *spatial)`.

    Les tailles spatiales doivent être divisibles par `2 ** nb_levels`. Pertes
    associés : `yax.obj.bce`, `yax.obj.dice_loss`, ou leur somme
    par `yax.obj.weighted_sum`.
    """
    # x : (dim_in, *spatial) -> logits (dim_out, *spatial)
    skips = []
    for encoder, down in zip(self.encoders, self.downs):
        x = encoder.apply(x)
        skips.append(x)                      # gardé pour la connexion de saut
        x = jax.nn.relu(down.apply(x))
    x = self.bottleneck.apply(x)
    for up_conv, decoder, skip in zip(self.up_convs, self.decoders,
                                      reversed(skips)):
        x = jax.nn.relu(up_conv.apply(upsample(x)))
        x = jnp.concatenate([skip, x], axis=0)   # axe 0 = canaux
        x = decoder.apply(x)
    return self.head.apply(x)

ConvBlock

yax.models.ConvBlock(dim_in, dim_out, nb_dims, rkey)

Bases: Module

Bloc de base du U-Net : deux convolutions de noyau 3, suivies chacune d'une relu.

Paramètres :

Nom Type Description Défaut
dim_in int

canaux d'entrée.

obligatoire
dim_out int

canaux de sortie.

obligatoire
nb_dims int

nombre de dimensions spatiales.

obligatoire
rkey Array

clé aléatoire pour l'initialisation.

obligatoire

Attributs :

Nom Type Description
conv1 Conv_nd

la première convolution.

conv2 Conv_nd

la seconde convolution.

Code source dans yax/models/UNet_nd.py
17
class ConvBlock(Module):
30
31
32
33
34
35
36
    conv1: Conv_nd
    conv2: Conv_nd

    def __init__(self, dim_in: int, dim_out: int, nb_dims: int, rkey: jax.Array):
        rkey1, rkey2 = jr.split(rkey)
        self.conv1 = Conv_nd(dim_in, dim_out, 3, nb_dims, rkey1)
        self.conv2 = Conv_nd(dim_out, dim_out, 3, nb_dims, rkey2)

VAE

Auto-encodeur variationnel (VAE, Kingma et Welling, 2013).

L'encodeur associe à chaque donnée une loi gaussienne dans un espace latent ; le décodeur reconstruit la donnée à partir d'un point tiré de cette loi. La perte rapproche ces lois de la normale standard, ce qui permet de générer des données nouvelles en décodant des points tirés au hasard.

VAE

yax.models.VAE(dim_in, dim_hidden, dim_latent, rkey, *, activation='gelu')

Bases: Module

Auto-encodeur variationnel.

Apprentissage non supervisé : passer les données comme entrées et comme cibles (Y = X), avec la perte yax.models.vae_loss.

Paramètres :

Nom Type Description Défaut
dim_in int

taille d'une donnée.

obligatoire
dim_hidden int

taille des couches cachées.

obligatoire
dim_latent int

taille de l'espace latent.

obligatoire
rkey Array

clé aléatoire pour l'initialisation.

obligatoire
activation str | Callable

activation des MLP.

'gelu'

Attributs :

Nom Type Description
encoder MLP

MLP qui renvoie la moyenne et la log-variance de la loi latente.

decoder MLP

MLP qui reconstruit une donnée à partir d'un point latent.

dim_latent int

taille de l'espace latent.

Code source dans yax/models/VAE.py
22
class VAE(Module):
40
41
42
43
44
45
46
47
48
49
50
    encoder: MLP
    decoder: MLP
    dim_latent: int = StaticField()

    def __init__(self, dim_in: int, dim_hidden: int, dim_latent: int, rkey: jax.Array, *,
                 activation: str | Callable = "gelu"):
        rkey1, rkey2 = jr.split(rkey)
        # l'encodeur rend mu et log_var concaténés
        self.encoder = MLP((dim_in, dim_hidden, 2 * dim_latent), activation, rkey1)
        self.decoder = MLP((dim_latent, dim_hidden, dim_in), activation, rkey2)
        self.dim_latent = dim_latent

encode

encode(x)

Renvoie (mu, log_var), les paramètres de la loi latente de x.

Code source dans yax/models/VAE.py
52
53
54
55
def encode(self, x: jax.Array) -> tuple[jax.Array, jax.Array]:
    """Renvoie `(mu, log_var)`, les paramètres de la loi latente de `x`."""
    sortie = self.encoder.apply(x)
    return sortie[:self.dim_latent], sortie[self.dim_latent:]   # mu, log_var

decode

decode(z)

Reconstruit une donnée à partir du point latent z.

Code source dans yax/models/VAE.py
57
58
59
def decode(self, z: jax.Array) -> jax.Array:
    """Reconstruit une donnée à partir du point latent `z`."""
    return self.decoder.apply(z)

reparametrise

reparametrise(x, rkey=None)

Tire un point latent pour x et renvoie (mu, log_var, z).

En mode inférence, z = mu et la clé n'est pas nécessaire.

Code source dans yax/models/VAE.py
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
def reparametrise(self, x: jax.Array, rkey: jax.Array | None = None
                  ) -> tuple[jax.Array, jax.Array, jax.Array]:
    """Tire un point latent pour `x` et renvoie `(mu, log_var, z)`.

    En mode inférence, `z = mu` et la clé n'est pas nécessaire.
    """
    mu, log_var = self.encode(x)
    if self.inference:
        return mu, log_var, mu
    if rkey is None:
        raise ValueError(
            "VAE en mode entrainement (inference=False) : une cle est "
            "requise pour la reparametrisation. Pour evaluer, basculer "
            "avec model.set_inference(True).")
    eps = jr.normal(rkey, mu.shape)
    return mu, log_var, mu + jnp.exp(0.5 * log_var) * eps

apply

apply(x, rkey=None)

Reconstruit x : encodage, tirage du point latent, décodage.

Code source dans yax/models/VAE.py
78
79
80
81
82
def apply(self, x: jax.Array, rkey: jax.Array | None = None) -> jax.Array:
    """Reconstruit `x` : encodage, tirage du point latent, décodage."""
    # la reconstruction d'UN échantillon
    _, _, z = self.reparametrise(x, rkey)
    return self.decode(z)

generate

generate(rkey, nb)

Génère nb données nouvelles en décodant des points tirés selon la normale standard.

Code source dans yax/models/VAE.py
84
85
86
87
def generate(self, rkey: jax.Array, nb: int) -> jax.Array:
    """Génère `nb` données nouvelles en décodant des points tirés selon la normale standard."""
    z = jr.normal(rkey, (nb, self.dim_latent))
    return jax.vmap(self.decode)(z)

vae_loss

yax.models.vae_loss(model, x, y, rkey=None, *, beta=1.0)

Perte du VAE : erreur de reconstruction plus divergence de Kullback-Leibler, moyennées sur le lot.

Paramètres :

Nom Type Description Défaut
model VAE

le VAE.

obligatoire
x ArrayLike

les entrées du lot.

obligatoire
y ArrayLike

les cibles, en général égales à x.

obligatoire
rkey Array | None

clé aléatoire pour le tirage latent.

None
beta float

poids de la divergence ; au-delà de 1, l'espace latent est plus structuré, au prix de la reconstruction.

1.0
Code source dans yax/models/VAE.py
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
@objective
def vae_loss(model: VAE, x: ArrayLike, y: ArrayLike, rkey: jax.Array | None = None, *,
             beta: float = 1.0) -> jax.Array:
    """Perte du VAE : erreur de reconstruction plus divergence de Kullback-Leibler, moyennées sur le lot.

    Args:
        model: le VAE.
        x: les entrées du lot.
        y: les cibles, en général égales à `x`.
        rkey: clé aléatoire pour le tirage latent.
        beta: poids de la divergence ; au-delà de 1, l'espace latent est plus
            structuré, au prix de la reconstruction.
    """

    def un_echantillon(x1, y1, rkey1):
        mu, log_var, z = model.reparametrise(x1, rkey1)
        reconstruction = jnp.sum((model.decode(z) - y1) ** 2)
        kl = -0.5 * jnp.sum(1.0 + log_var - mu**2 - jnp.exp(log_var))
        return reconstruction + beta * kl

    if rkey is None:
        return jnp.mean(jax.vmap(un_echantillon, in_axes=(0, 0, None))(x, y, None))
    return jnp.mean(jax.vmap(un_echantillon)(x, y, jr.split(rkey, jnp.shape(x)[0])))

RealNVP

Flot normalisant RealNVP (Dinh et al., 2017).

Un flot est une transformation inversible entre les données et une variable gaussienne. Sa vraisemblance se calcule exactement, et l'on génère des données en inversant le flot sur des tirages gaussiens. Chaque couche est un couplage affine : une moitié des composantes passe inchangée et détermine une transformation affine de l'autre moitié.

RealNVP

yax.models.RealNVP(dim, dim_hidden, nb_couplings, rkey, *, activation='gelu')

Bases: Module

Flot normalisant RealNVP.

Apprentissage non supervisé : passer Y = X, avec la perte yax.models.realnvp_loss.

Paramètres :

Nom Type Description Défaut
dim int

taille d'une donnée.

obligatoire
dim_hidden int

taille des couches cachées des couplages.

obligatoire
nb_couplings int

nombre de couplages.

obligatoire
rkey Array

clé aléatoire pour l'initialisation.

obligatoire
activation str | Callable

activation des MLP.

'gelu'

Attributs :

Nom Type Description
couplings list[Coupling]

les couplages affines.

dim int

taille d'une donnée.

Code source dans yax/models/RealNVP.py
70
class RealNVP(Module):
87
88
89
90
91
92
93
94
95
96
    couplings: list[Coupling]
    dim: int = StaticField()

    def __init__(self, dim: int, dim_hidden: int, nb_couplings: int, rkey: jax.Array, *,
                 activation: str | Callable = "gelu"):
        rkeys = jr.split(rkey, nb_couplings)
        self.couplings = [Coupling(dim, dim_hidden, bool(i % 2), rkey,
                                   activation=activation)
                          for i, rkey in enumerate(rkeys)]
        self.dim = dim

apply

apply(x, rkey=None)

Transforme une donnée x et renvoie (z, log_det).

Code source dans yax/models/RealNVP.py
 98
 99
100
101
102
103
104
def apply(self, x: jax.Array, rkey: jax.Array | None = None) -> tuple[jax.Array, jax.Array]:
    """Transforme une donnée `x` et renvoie `(z, log_det)`."""
    log_det = jnp.zeros(())
    for coupling in self.couplings:
        x, ld = coupling.forward(x)
        log_det = log_det + ld
    return x, log_det

inverse

inverse(z)

Transforme un point gaussien z en donnée.

Code source dans yax/models/RealNVP.py
106
107
108
109
110
def inverse(self, z: jax.Array) -> jax.Array:
    """Transforme un point gaussien `z` en donnée."""
    for coupling in reversed(self.couplings):
        z = coupling.inverse(z)
    return z

sample

sample(rkey, nb)

Génère nb données nouvelles.

Code source dans yax/models/RealNVP.py
112
113
114
115
def sample(self, rkey: jax.Array, nb: int) -> jax.Array:
    """Génère `nb` données nouvelles."""
    z = jr.normal(rkey, (nb, self.dim))
    return jax.vmap(self.inverse)(z)

Coupling

yax.models.Coupling(dim, dim_hidden, swap, rkey, *, activation='gelu')

Bases: Module

Couplage affine : une moitié de x reste fixe, l'autre devient x * exp(s) + t.

s et t sont calculés par un MLP à partir de la moitié fixe.

Paramètres :

Nom Type Description Défaut
dim int

taille d'une donnée.

obligatoire
dim_hidden int

taille des couches cachées du MLP.

obligatoire
swap bool

échange le rôle des deux moitiés.

obligatoire
rkey Array

clé aléatoire pour l'initialisation.

obligatoire
activation str | Callable

activation du MLP.

'gelu'

Attributs :

Nom Type Description
net MLP

le MLP qui calcule s et t.

swap bool

rôle des deux moitiés.

Code source dans yax/models/RealNVP.py
22
class Coupling(Module):
38
39
40
41
42
43
44
45
46
    net: MLP
    swap: bool = StaticField()

    def __init__(self, dim: int, dim_hidden: int, swap: bool, rkey: jax.Array, *,
                 activation: str | Callable = "gelu"):
        assert dim % 2 == 0, f"dim:{dim} doit etre paire (couplage par moities)"
        # le réseau rend s et t concaténés : dim/2 -> dim
        self.net = MLP((dim // 2, dim_hidden, dim), activation, rkey)
        self.swap = swap

forward

forward(x)

Transforme x et renvoie (z, log_det), avec le logarithme du déterminant jacobien.

Code source dans yax/models/RealNVP.py
52
53
54
55
56
57
58
59
def forward(self, x: jax.Array) -> tuple[jax.Array, jax.Array]:
    """Transforme `x` et renvoie `(z, log_det)`, avec le logarithme du déterminant jacobien."""
    a, b = jnp.split(x, 2)
    fixe, mobile = (b, a) if self.swap else (a, b)
    s, t = self._coeffs(fixe)
    mobile = mobile * jnp.exp(s) + t
    z = jnp.concatenate([mobile, fixe] if self.swap else [fixe, mobile])
    return z, jnp.sum(s)

inverse

inverse(z)

Inverse la transformation : renvoie x à partir de z.

Code source dans yax/models/RealNVP.py
61
62
63
64
65
66
67
def inverse(self, z: jax.Array) -> jax.Array:
    """Inverse la transformation : renvoie `x` à partir de `z`."""
    a, b = jnp.split(z, 2)
    fixe, mobile = (b, a) if self.swap else (a, b)
    s, t = self._coeffs(fixe)
    mobile = (mobile - t) * jnp.exp(-s)
    return jnp.concatenate([mobile, fixe] if self.swap else [fixe, mobile])

realnvp_loss

yax.models.realnvp_loss(model, x, y, rkey=None)

Opposé de la log-vraisemblance moyenne du lot, en nats.

Paramètres :

Nom Type Description Défaut
model RealNVP

le flot.

obligatoire
x ArrayLike

les données du lot.

obligatoire
y ArrayLike

ignoré ; passer Y = X.

obligatoire
rkey Array | None

inutilisé.

None
Code source dans yax/models/RealNVP.py
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
@objective
def realnvp_loss(model: RealNVP, x: ArrayLike, y: ArrayLike,
                 rkey: jax.Array | None = None) -> jax.Array:
    """Opposé de la log-vraisemblance moyenne du lot, en nats.

    Args:
        model: le flot.
        x: les données du lot.
        y: ignoré ; passer `Y = X`.
        rkey: inutilisé.
    """
    z, log_det = jax.vmap(model.apply, in_axes=(0, None))(jnp.asarray(x), None)
    log_prior = (-0.5 * jnp.sum(z**2, axis=-1)
                 - 0.5 * model.dim * jnp.log(2.0 * jnp.pi))
    return -jnp.mean(log_prior + log_det)

Diffusion

Modèle de diffusion (DDPM, Ho et al., 2020), pour des données de petite dimension.

Pendant l'entraînement, les données sont progressivement bruitées et le réseau apprend à prédire le bruit ajouté. Pour générer, on part d'un bruit pur et on le débruite pas à pas.

Diffusion

yax.models.Diffusion(dim, dim_hidden, nb_steps, rkey, *, activation='gelu')

Bases: Module

Modèle de diffusion.

Apprentissage non supervisé : passer Y = X, avec la perte yax.models.diffusion_loss.

Paramètres :

Nom Type Description Défaut
dim int

taille d'une donnée.

obligatoire
dim_hidden int

taille des couches cachées.

obligatoire
nb_steps int

nombre de pas de bruitage.

obligatoire
rkey Array

clé aléatoire pour l'initialisation.

obligatoire
activation str | Callable

activation du réseau.

'gelu'

Attributs :

Nom Type Description
net MLP

le réseau qui prédit le bruit.

dim int

taille d'une donnée.

nb_steps int

nombre de pas.

betas Array

variance du bruit ajouté à chaque pas.

alpha_bars Array

part du signal qui subsiste après chaque pas.

Code source dans yax/models/Diffusion.py
20
class Diffusion(Module):
40
41
42
43
44
45
46
47
48
49
50
51
52
53
    net: MLP
    dim: int = StaticField()
    nb_steps: int = StaticField()
    betas: jax.Array = StaticField()
    alpha_bars: jax.Array = StaticField()

    def __init__(self, dim: int, dim_hidden: int, nb_steps: int, rkey: jax.Array, *,
                 activation: str | Callable = "gelu"):
        # entrée du réseau : (x_t, t/nb_steps) -> eps prédit
        self.net = MLP((dim + 1, dim_hidden, dim_hidden, dim), activation, rkey)
        self.dim = dim
        self.nb_steps = nb_steps
        self.betas = jnp.linspace(1e-4, 0.04, nb_steps)     # planning linéaire
        self.alpha_bars = jnp.cumprod(1.0 - self.betas)

apply

apply(x_t, rkey=None, *, t)

Prédit le bruit contenu dans x_t, une donnée bruitée jusqu'au pas t.

Code source dans yax/models/Diffusion.py
55
56
57
58
59
def apply(self, x_t: jax.Array, rkey: jax.Array | None = None, *,
          t: int | jax.Array) -> jax.Array:
    """Prédit le bruit contenu dans `x_t`, une donnée bruitée jusqu'au pas `t`."""
    temps = jnp.asarray(t, dtype=x_t.dtype) / self.nb_steps
    return self.net.apply(jnp.concatenate([x_t, temps[None]]))

sample

sample(rkey, nb)

Génère nb données nouvelles, en nb_steps pas de débruitage.

Code source dans yax/models/Diffusion.py
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
def sample(self, rkey: jax.Array, nb: int) -> jax.Array:
    """Génère `nb` données nouvelles, en `nb_steps` pas de débruitage."""

    def un_pas(x, entree):
        t, rkey_pas = entree
        eps_pred = jax.vmap(lambda xt: self.apply(xt, t=t))(x)
        alpha = 1.0 - self.betas[t]
        x = (x - (1.0 - alpha) / jnp.sqrt(1.0 - self.alpha_bars[t]) * eps_pred) \
            / jnp.sqrt(alpha)
        # du bruit à chaque pas, sauf au dernier (t = 0)
        x = x + jnp.where(t > 0, jnp.sqrt(self.betas[t]), 0.0) * jr.normal(rkey_pas, x.shape)
        return x, None

    rkey_init, rkey_scan = jr.split(rkey)
    x = jr.normal(rkey_init, (nb, self.dim))
    instants = jnp.arange(self.nb_steps - 1, -1, -1)
    x, _ = jax.lax.scan(un_pas, x, (instants, jr.split(rkey_scan, self.nb_steps)))
    return x

diffusion_loss

yax.models.diffusion_loss(model, x, y, rkey=None)

Perte de diffusion : erreur quadratique entre le bruit ajouté et le bruit prédit.

Les pas et les bruits sont tirés au hasard. Sans clé (rkey=None), ce tirage est fixe, ce qui rend la perte de validation comparable d'une époque à l'autre.

Paramètres :

Nom Type Description Défaut
model Diffusion

le modèle de diffusion.

obligatoire
x ArrayLike

les données du lot.

obligatoire
y ArrayLike

ignoré ; passer Y = X.

obligatoire
rkey Array | None

clé aléatoire, ou None.

None
Code source dans yax/models/Diffusion.py
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
@objective
def diffusion_loss(model: Diffusion, x: ArrayLike, y: ArrayLike,
                   rkey: jax.Array | None = None) -> jax.Array:
    """Perte de diffusion : erreur quadratique entre le bruit ajouté et le bruit prédit.

    Les pas et les bruits sont tirés au hasard. Sans clé (`rkey=None`), ce
    tirage est fixe, ce qui rend la perte de validation comparable d'une
    époque à l'autre.

    Args:
        model: le modèle de diffusion.
        x: les données du lot.
        y: ignoré ; passer `Y = X`.
        rkey: clé aléatoire, ou `None`.
    """
    if rkey is None:
        rkey = jr.key(0)
    rkey_t, rkey_eps = jr.split(rkey)
    t = jr.randint(rkey_t, (jnp.shape(x)[0],), 0, model.nb_steps)
    eps = jr.normal(rkey_eps, jnp.shape(x))
    ab = model.alpha_bars[t][:, None]
    x_t = jnp.sqrt(ab) * x + jnp.sqrt(1.0 - ab) * eps
    eps_pred = jax.vmap(lambda xt, t1: model.apply(xt, t=t1))(x_t, t)
    return jnp.mean((eps_pred - eps) ** 2)

Mini-YOLO

Mini-YOLO : détection d'objets d'une seule classe, sur des images 32×32.

L'image est découpée en une grille de GRID × GRID cellules. Pour chaque cellule, le réseau prédit en un seul passage la présence d'un objet dont le centre tombe dans la cellule, la position de ce centre et la taille de la boîte. Une suppression des non-maxima ne garde ensuite qu'une boîte par objet.

Une cellule est décrite par 5 valeurs : la présence d'un objet, la position (dx, dy) du centre dans la cellule, la largeur et la hauteur de la boîte en fraction de l'image. Si deux centres tombent dans la même cellule, un seul objet est détecté.

MiniYOLO

yax.models.MiniYOLO(rkey)

Bases: Module

Détecteur mini-YOLO.

Entrée (1, 32, 32), sortie (5, GRID, GRID) : cinq logits par cellule. Cibles avec yax.models.encode_targets, perte yax.models.yolo_loss, détection complète avec yax.models.detect.

Paramètres :

Nom Type Description Défaut
rkey Array

clé aléatoire pour l'initialisation.

obligatoire

Attributs :

Nom Type Description
conv1 Conv_nd

convolution, 16 canaux.

conv2 Conv_nd

convolution de pas 2, 32 canaux.

conv3 Conv_nd

convolution de pas 2, 32 canaux.

conv4 Conv_nd

convolution, 32 canaux.

head Conv_nd

convolution 1×1 vers les 5 valeurs de chaque cellule.

Code source dans yax/models/MiniYOLO.py
36
class MiniYOLO(Module):
53
54
55
56
57
58
59
60
61
62
63
64
65
    conv1: Conv_nd
    conv2: Conv_nd
    conv3: Conv_nd
    conv4: Conv_nd
    head: Conv_nd

    def __init__(self, rkey: jax.Array):
        rkey1, rkey2, rkey3, rkey4, rkey5 = jr.split(rkey, 5)
        self.conv1 = Conv_nd(1, 16, 3, 2, rkey1)          # (16, 32, 32)
        self.conv2 = Conv_nd(16, 32, 3, 2, rkey2, stride=2)  # (32, 16, 16)
        self.conv3 = Conv_nd(32, 32, 3, 2, rkey3, stride=2)  # (32, 8, 8)
        self.conv4 = Conv_nd(32, 32, 3, 2, rkey4)         # (32, 8, 8)
        self.head = Conv_nd(32, 5, 1, 2, rkey5)           # (5, GRID, GRID)

apply

apply(x, rkey=None)

Renvoie les logits (5, GRID, GRID) pour une image (1, 32, 32).

Code source dans yax/models/MiniYOLO.py
67
68
69
70
71
72
73
74
def apply(self, x: jax.Array, rkey: jax.Array | None = None) -> jax.Array:
    """Renvoie les logits `(5, GRID, GRID)` pour une image `(1, 32, 32)`."""
    # x : (1, 32, 32) -> logits (5, GRID, GRID)
    x = jax.nn.relu(self.conv1.apply(x))
    x = jax.nn.relu(self.conv2.apply(x))
    x = jax.nn.relu(self.conv3.apply(x))
    x = jax.nn.relu(self.conv4.apply(x))
    return self.head.apply(x)

yolo_loss

yax.models.yolo_loss(model, x, y, rkey=None)

Perte du détecteur.

Entropie croisée sur la présence d'objet dans toutes les cellules, plus erreur quadratique sur les boîtes des cellules qui contiennent un objet.

Paramètres :

Nom Type Description Défaut
model MiniYOLO

le détecteur.

obligatoire
x ArrayLike

les images du lot.

obligatoire
y ArrayLike

les grilles cibles, de forme (batch, 5, GRID, GRID).

obligatoire
rkey Array | None

inutilisé.

None
Code source dans yax/models/MiniYOLO.py
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
@objective
def yolo_loss(model: MiniYOLO, x: ArrayLike, y: ArrayLike,
              rkey: jax.Array | None = None) -> jax.Array:
    """Perte du détecteur.

    Entropie croisée sur la présence d'objet dans toutes les cellules, plus
    erreur quadratique sur les boîtes des cellules qui contiennent un objet.

    Args:
        model: le détecteur.
        x: les images du lot.
        y: les grilles cibles, de forme `(batch, 5, GRID, GRID)`.
        rkey: inutilisé.
    """
    logits = batch_apply(model, x, rkey)          # (batch, 5, G, G)
    obj_logits = logits[:, 0]
    box_pred = jax.nn.sigmoid(logits[:, 1:])      # dans [0,1], comme les cibles

    y = jnp.asarray(y)
    obj_true = y[:, 0]
    box_true = y[:, 1:]

    loss_obj = jnp.mean(optax.sigmoid_binary_cross_entropy(obj_logits, obj_true))
    mask = obj_true[:, None]                      # (batch, 1, G, G)
    loss_box = jnp.sum(mask * (box_pred - box_true) ** 2) \
        / (4.0 * jnp.sum(obj_true) + 1e-6)
    return loss_obj + LAMBDA_BOX * loss_box

detect

yax.models.detect(model, img, *, iou_threshold=0.5, score_threshold=0.3)

Détecte les objets d'une image.

Paramètres :

Nom Type Description Défaut
model MiniYOLO

le détecteur.

obligatoire
img Array

l'image, de forme (1, 32, 32).

obligatoire
iou_threshold float

seuil de recouvrement de la suppression des non-maxima.

0.5
score_threshold float

score minimal d'une détection. Le baisser trouve plus d'objets, au prix de davantage de fausses détections.

0.3

Renvoie :

Type Description
tuple[ndarray, ndarray]

Le couple (boxes, scores), tableaux numpy de formes (n, 4) et (n,).

Code source dans yax/models/MiniYOLO.py
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
def detect(model: MiniYOLO, img: jax.Array, *, iou_threshold: float = 0.5,
           score_threshold: float = 0.3) -> tuple[np.ndarray, np.ndarray]:
    """Détecte les objets d'une image.

    Args:
        model: le détecteur.
        img: l'image, de forme `(1, 32, 32)`.
        iou_threshold: seuil de recouvrement de la suppression des non-maxima.
        score_threshold: score minimal d'une détection. Le baisser trouve plus
            d'objets, au prix de davantage de fausses détections.

    Returns:
        Le couple `(boxes, scores)`, tableaux numpy de formes `(n, 4)` et `(n,)`.
    """
    boxes, scores = decode_predictions(model.apply(img))
    keep = nms(boxes, scores, iou_threshold=iou_threshold,
               score_threshold=score_threshold)
    return np.asarray(boxes)[keep], np.asarray(scores)[keep]