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 |
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 |
activation_fn |
Callable
|
la fonction d'activation. |
Code source dans yax/models/FNO_nd.py
92 | |
121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 | |
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 | |
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 |
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 | |
44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 | |
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 | |
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 | |
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 | |
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 | |
30 31 32 33 34 35 36 | |
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 | |
40 41 42 43 44 45 46 47 48 49 50 | |
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 | |
decode
¶
decode(z)
Reconstruit une donnée à partir du point latent z.
Code source dans yax/models/VAE.py
57 58 59 | |
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 | |
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 | |
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 | |
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 à |
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 | |
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 | |
87 88 89 90 91 92 93 94 95 96 | |
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 | |
inverse
¶
inverse(z)
Transforme un point gaussien z en donnée.
Code source dans yax/models/RealNVP.py
106 107 108 109 110 | |
sample
¶
sample(rkey, nb)
Génère nb données nouvelles.
Code source dans yax/models/RealNVP.py
112 113 114 115 | |
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 |
swap |
bool
|
rôle des deux moitiés. |
Code source dans yax/models/RealNVP.py
22 | |
38 39 40 41 42 43 44 45 46 | |
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 | |
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 | |
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 |
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 | |
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 | |
40 41 42 43 44 45 46 47 48 49 50 51 52 53 | |
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 | |
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 | |
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 |
obligatoire |
rkey
|
Array | None
|
clé aléatoire, ou |
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 | |
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 | |
53 54 55 56 57 58 59 60 61 62 63 64 65 | |
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 | |
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 |
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 | |
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 |
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 |
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 | |