yax.layers¶
Les couches de réseaux de neurones. Chacune est un yax.Module et traite un exemple à la fois ; yax.batch_apply ou jax.vmap traitent un lot.
Linear
¶
yax.layers.Linear(dim_in, dim_out, rkey)
Bases: Module
Couche linéaire (affine) : x @ weight + bias.
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
dim_in
|
int
|
taille de l'entrée. |
obligatoire |
dim_out
|
int
|
taille de la sortie. |
obligatoire |
rkey
|
Array
|
clé aléatoire pour l'initialisation. |
obligatoire |
Attributs :
| Nom | Type | Description |
|---|---|---|
weight |
Array
|
matrice |
bias |
Array
|
vecteur |
Code source dans yax/layers/Linear.py
6 | |
18 19 20 21 22 23 24 25 | |
apply
¶
apply(x, rkey=None)
Calcule x @ weight + bias pour un exemple x de forme (dim_in,).
Code source dans yax/layers/Linear.py
27 28 29 30 31 32 | |
MLP
¶
yax.layers.MLP(layer_sizes, activation, rkey)
Bases: Module
Perceptron multicouche : une suite de couches Linear séparées par une activation.
Aucune activation n'est appliquée après la dernière couche : la sortie est une valeur réelle ou un logit.
mlp = yax.layers.MLP((2, 32, 32, 1), "tanh", rkey) # deux couches cachées de 32
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
layer_sizes
|
Sequence[int]
|
toutes les tailles, de l'entrée à la sortie. |
obligatoire |
activation
|
str | Callable
|
nom d'une activation de |
obligatoire |
rkey
|
Array
|
clé aléatoire pour l'initialisation. |
obligatoire |
Attributs :
| Nom | Type | Description |
|---|---|---|
layers |
list[Linear]
|
les couches |
activation_fn |
Callable
|
la fonction d'activation. |
layer_sizes |
tuple[int, ...]
|
les tailles des couches. |
Code source dans yax/layers/MLP.py
10 | |
31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 | |
apply
¶
apply(x, rkey=None)
Calcule la sortie pour un exemple x de forme (layer_sizes[0],).
Code source dans yax/layers/MLP.py
48 49 50 51 52 53 54 55 | |
Dropout
¶
yax.layers.Dropout(rate)
Bases: Module
Dropout : met à zéro une fraction des composantes pendant l'entraînement.
En mode entraînement, chaque composante est annulée avec la probabilité
rate et les autres sont multipliées par 1 / (1 - rate) ; une clé
aléatoire est alors obligatoire. En mode inférence, la couche renvoie
l'entrée inchangée.
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
rate
|
float
|
probabilité d'annuler une composante, dans |
obligatoire |
Attributs :
| Nom | Type | Description |
|---|---|---|
rate |
float
|
la probabilité d'annulation. |
Code source dans yax/layers/Dropout.py
8 | |
22 23 24 25 26 | |
apply
¶
apply(x, rkey=None)
Applique le dropout à x.
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
x
|
Array
|
un tableau de forme quelconque. |
obligatoire |
rkey
|
Array | None
|
clé aléatoire, obligatoire en mode entraînement. |
None
|
Code source dans yax/layers/Dropout.py
28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 | |
PReLU
¶
yax.layers.PReLU(*, dim=1, init=0.25)
Bases: Module
Activation PReLU : x si x ≥ 0, a · x sinon, où la pente a est apprise.
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
dim
|
int
|
nombre de pentes ; 1 pour une pente commune, |
1
|
init
|
float
|
valeur initiale de la pente. |
0.25
|
Attributs :
| Nom | Type | Description |
|---|---|---|
a |
Array
|
les pentes, de forme |
Code source dans yax/layers/PReLU.py
7 | |
18 19 20 21 22 | |
apply
¶
apply(x, rkey=None)
Applique l'activation à x.
Code source dans yax/layers/PReLU.py
24 25 26 27 | |
LayerNorm
¶
yax.layers.LayerNorm(dim, *, epsilon=1e-05)
Bases: Module
Normalisation de couche : centre et réduit les composantes de chaque exemple.
Moyenne et variance sont calculées sur le dernier axe, exemple par
exemple ; le résultat est ensuite multiplié par gamma et décalé de
beta.
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
dim
|
int
|
taille du dernier axe. |
obligatoire |
epsilon
|
float
|
petite constante ajoutée à la variance. |
1e-05
|
Attributs :
| Nom | Type | Description |
|---|---|---|
gamma |
Array
|
facteur d'échelle, initialisé à 1. |
beta |
Array
|
décalage, initialisé à 0. |
epsilon |
float
|
la constante de stabilité. |
Code source dans yax/layers/LayerNorm.py
6 | |
22 23 24 25 26 27 28 29 30 31 | |
apply
¶
apply(x, rkey=None)
Normalise x sur son dernier axe.
Code source dans yax/layers/LayerNorm.py
33 34 35 36 37 38 39 40 | |
Embedding
¶
yax.layers.Embedding(nb_embeddings, dim, rkey)
Bases: Module
Table de vecteurs apprise, indexée par des entiers.
Associe un vecteur de taille dim à chaque identifiant (un mot, une
catégorie…).
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
nb_embeddings
|
int
|
nombre d'identifiants possibles. |
obligatoire |
dim
|
int
|
taille de chaque vecteur. |
obligatoire |
rkey
|
Array
|
clé aléatoire pour l'initialisation. |
obligatoire |
Attributs :
| Nom | Type | Description |
|---|---|---|
weight |
Array
|
la table, de forme |
Code source dans yax/layers/Embedding.py
7 | |
21 22 23 24 25 | |
apply
¶
apply(ids, rkey=None)
Renvoie les vecteurs des identifiants entiers ids, de forme ids.shape + (dim,).
Code source dans yax/layers/Embedding.py
27 28 29 30 31 32 | |
Conv_nd
¶
yax.layers.Conv_nd(dim_in, dim_out, kernel_size, nb_dims, rkey, *, stride=1, padding='SAME')
Bases: Module
Convolution en dimension quelconque : signal (1D), image (2D), volume (3D)…
Un exemple est de forme (dim_in, *spatial), les canaux en premier ; la
sortie est de forme (dim_out, *spatial').
conv = yax.layers.Conv_nd(dim_in=3, dim_out=16, kernel_size=3, nb_dims=2, rkey=rkey)
y = conv.apply(image) # image (3, 32, 32) -> y (16, 32, 32)
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
dim_in
|
int
|
nombre de canaux d'entrée. |
obligatoire |
dim_out
|
int
|
nombre de canaux de sortie. |
obligatoire |
kernel_size
|
int | tuple[int, ...]
|
taille du noyau, un entier ou un tuple de longueur |
obligatoire |
nb_dims
|
int
|
nombre de dimensions spatiales. |
obligatoire |
rkey
|
Array
|
clé aléatoire pour l'initialisation. |
obligatoire |
stride
|
int | tuple[int, ...]
|
pas de la convolution, un entier ou un tuple. |
1
|
padding
|
str
|
|
'SAME'
|
Attributs :
| Nom | Type | Description |
|---|---|---|
weight |
Array
|
les noyaux, de forme |
bias |
Array
|
les biais, de forme |
stride |
tuple[int, ...]
|
pas de la convolution. |
padding |
str
|
mode de complétion des bords. |
Code source dans yax/layers/Conv_nd.py
8 | |
35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 | |
apply
¶
apply(x, rkey=None)
Calcule la convolution d'un exemple x de forme (dim_in, *spatial).
Code source dans yax/layers/Conv_nd.py
69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 | |
Couches récurrentes : les cellules GRU et LSTM, et la couche RNN_layer qui les applique à une séquence.
RNN_layer
¶
yax.layers.RNN_layer(dim_in, dim_hidden, rkey, *, cell_type='gru')
Bases: Module
Couche récurrente : applique une cellule GRU ou LSTM à toute une séquence.
La séquence est parcourue du premier au dernier pas, en partant d'un état
nul, et la couche renvoie l'état caché de chaque pas. Le dernier, hs[-1],
résume toute la séquence.
rnn = yax.layers.RNN_layer(dim_in=3, dim_hidden=16, rkey=rkey)
hs = rnn.apply(xs) # xs (T, 3) -> hs (T, 16) : un état par pas
h = hs[-1] # l'état final
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
dim_in
|
int
|
taille de l'entrée à chaque pas. |
obligatoire |
dim_hidden
|
int
|
taille de l'état caché. |
obligatoire |
rkey
|
Array
|
clé aléatoire pour l'initialisation. |
obligatoire |
cell_type
|
str
|
|
'gru'
|
Attributs :
| Nom | Type | Description |
|---|---|---|
cell |
GRUCell | LSTMCell
|
la cellule récurrente. |
dim_in |
int
|
taille de l'entrée. |
dim_hidden |
int
|
taille de l'état caché. |
cell_type |
str
|
type de cellule. |
Code source dans yax/layers/RNN_layer.py
114 | |
139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 | |
apply
¶
apply(xs, rkey=None)
Parcourt la séquence xs, de forme (T, dim_in), et renvoie les états cachés, de forme (T, dim_hidden).
Code source dans yax/layers/RNN_layer.py
158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 | |
GRUCell
¶
yax.layers.GRUCell(dim_in, dim_hidden, rkey)
Bases: Module
Cellule GRU : un pas de récurrence.
Calcule le nouvel état caché à partir de l'entrée du pas et de l'état
précédent. Pour traiter une séquence entière, utiliser RNN_layer.
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
dim_in
|
int
|
taille de l'entrée d'un pas. |
obligatoire |
dim_hidden
|
int
|
taille de l'état caché. |
obligatoire |
rkey
|
Array
|
clé aléatoire pour l'initialisation. |
obligatoire |
Attributs :
| Nom | Type | Description |
|---|---|---|
weight_ih |
Array
|
poids appliqués à l'entrée, pour les trois portes. |
weight_hh |
Array
|
poids appliqués à l'état caché, pour les trois portes. |
bias |
Array
|
biais des portes. |
bias_n |
Array
|
biais de la porte candidate, côté état caché. |
Code source dans yax/layers/RNN_layer.py
14 | |
31 32 33 34 35 36 37 38 39 40 41 42 | |
apply
¶
apply(x, rkey=None, *, carry)
Calcule un pas de récurrence.
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
x
|
Array
|
entrée du pas, de forme |
obligatoire |
rkey
|
Array | None
|
inutilisé (cellule déterministe). |
None
|
carry
|
Array
|
état caché précédent |
obligatoire |
Renvoie :
| Type | Description |
|---|---|
Array
|
Le nouvel état caché, de forme |
Code source dans yax/layers/RNN_layer.py
44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 | |
LSTMCell
¶
yax.layers.LSTMCell(dim_in, dim_hidden, rkey)
Bases: Module
Cellule LSTM : un pas de récurrence, avec un état double (h, c).
Pour traiter une séquence entière, utiliser RNN_layer.
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
dim_in
|
int
|
taille de l'entrée d'un pas. |
obligatoire |
dim_hidden
|
int
|
taille de l'état caché. |
obligatoire |
rkey
|
Array
|
clé aléatoire pour l'initialisation. |
obligatoire |
Attributs :
| Nom | Type | Description |
|---|---|---|
weight_ih |
Array
|
poids appliqués à l'entrée, pour les quatre portes. |
weight_hh |
Array
|
poids appliqués à l'état caché, pour les quatre portes. |
bias |
Array
|
biais des portes. |
Code source dans yax/layers/RNN_layer.py
65 | |
80 81 82 83 84 85 86 87 88 89 | |
apply
¶
apply(x, rkey=None, *, carry)
Calcule un pas de récurrence.
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
x
|
Array
|
entrée du pas, de forme |
obligatoire |
rkey
|
Array | None
|
inutilisé (cellule déterministe). |
None
|
carry
|
tuple[Array, Array]
|
le couple |
obligatoire |
Renvoie :
| Type | Description |
|---|---|
tuple[Array, Array]
|
Le nouveau couple |
Code source dans yax/layers/RNN_layer.py
91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 | |
MultiHeadAttention
¶
yax.layers.MultiHeadAttention(dim, nb_heads, rkey, *, dropout_rate=0.0)
Bases: Module
Attention multi-têtes, pour une séquence.
Calcule softmax(Q Kᵀ / √d) V dans nb_heads sous-espaces en parallèle.
La même couche couvre trois usages :
- auto-attention :
apply(x); - attention croisée :
apply(x, x_kv=memoire), les clés et valeurs venant d'une autre séquence ; - attention causale :
apply(x, mask=yax.preprocessing.causal_mask(n)).
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
dim
|
int
|
taille des vecteurs de la séquence ; doit être divisible par |
obligatoire |
nb_heads
|
int
|
nombre de têtes. |
obligatoire |
rkey
|
Array
|
clé aléatoire pour l'initialisation. |
obligatoire |
dropout_rate
|
float
|
taux de dropout sur les poids d'attention. |
0.0
|
Attributs :
| Nom | Type | Description |
|---|---|---|
q_proj |
Linear
|
projection des requêtes. |
k_proj |
Linear
|
projection des clés. |
v_proj |
Linear
|
projection des valeurs. |
out_proj |
Linear
|
projection de sortie. |
dropout |
Dropout
|
le dropout des poids d'attention. |
nb_heads |
int
|
nombre de têtes. |
Code source dans yax/layers/MultiHeadAttention.py
12 | |
37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 | |
apply
¶
apply(x, rkey=None, *, x_kv=None, mask=None)
Calcule l'attention pour une séquence.
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
x
|
Array
|
la séquence des requêtes, de forme |
obligatoire |
rkey
|
Array | None
|
clé aléatoire pour le dropout, en mode entraînement. |
None
|
x_kv
|
Array | None
|
la séquence des clés et valeurs, de forme |
None
|
mask
|
Array | None
|
masque additif (0 visible, |
None
|
Renvoie :
| Type | Description |
|---|---|
Array
|
Un tableau de forme |
Code source dans yax/layers/MultiHeadAttention.py
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 | |
attention_maps
¶
attention_maps(x, *, x_kv=None, mask=None)
Renvoie les poids d'attention, pour les visualiser.
Mêmes arguments que apply, sans dropout.
Renvoie :
| Type | Description |
|---|---|
Array
|
Les poids, de forme |
Code source dans yax/layers/MultiHeadAttention.py
105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 | |
TransformerBlock
¶
yax.layers.TransformerBlock(dim, nb_heads, dim_ff, rkey, *, dropout_rate=0.0, norm_position='pre')
Bases: Module
Bloc transformer : attention multi-têtes puis feed-forward, chacun avec connexion résiduelle et normalisation.
Deux placements de la normalisation sont possibles :
"pre"(défaut) :x + sous_couche(norm(x)), la variante la plus stable à entraîner, celle des architectures récentes ;"post":norm(x + sous_couche(x)), la variante de l'article original.
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
dim
|
int
|
taille des vecteurs de la séquence. |
obligatoire |
nb_heads
|
int
|
nombre de têtes d'attention. |
obligatoire |
dim_ff
|
int
|
taille de la couche cachée du feed-forward. |
obligatoire |
rkey
|
Array
|
clé aléatoire pour l'initialisation. |
obligatoire |
dropout_rate
|
float
|
taux de dropout. |
0.0
|
norm_position
|
str
|
|
'pre'
|
Attributs :
| Nom | Type | Description |
|---|---|---|
attention |
MultiHeadAttention
|
la couche d'attention multi-têtes. |
feed_forward |
MLP
|
le MLP |
norm1 |
LayerNorm
|
la normalisation de la partie attention. |
norm2 |
LayerNorm
|
la normalisation de la partie feed-forward. |
dropout |
Dropout
|
le dropout du bloc. |
norm_position |
str
|
placement de la normalisation. |
Code source dans yax/layers/TransformerBlock.py
11 | |
36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 | |
apply
¶
apply(x, rkey=None, *, mask=None)
Transforme une séquence.
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
x
|
Array
|
la séquence, de forme |
obligatoire |
rkey
|
Array | None
|
clé aléatoire pour le dropout, en mode entraînement. |
None
|
mask
|
Array | None
|
masque additif transmis à l'attention. |
None
|
Renvoie :
| Type | Description |
|---|---|
Array
|
Une séquence de même forme. |
Code source dans yax/layers/TransformerBlock.py
57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 | |
MessagePassing_layer
¶
yax.layers.MessagePassing_layer(dim_node, dim_hidden, dim_out, rkey, *, activation='relu', aggregation='sum')
Bases: Module
Couche de passage de messages sur un graphe.
Pour chaque arête, un MLP calcule un message à partir des états de ses deux extrémités. Chaque nœud agrège les messages qu'il reçoit, puis un second MLP met à jour son état.
Le graphe est donné par deux tableaux d'entiers de même longueur : l'arête
e va de senders[e] à receivers[e]. Un lien non orienté compte pour
deux arêtes. Quatre agrégations sont proposées :
"sum": somme des messages ; tient compte du nombre de voisins ;"mean": moyenne ; indépendante du nombre de voisins ;"max": maximum composante par composante ;"attention": moyenne pondérée par des poids appris, à la manière de GATv2.
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
dim_node
|
int
|
taille de l'état d'un nœud en entrée. |
obligatoire |
dim_hidden
|
int
|
taille des messages. |
obligatoire |
dim_out
|
int
|
taille de l'état d'un nœud en sortie. |
obligatoire |
rkey
|
Array
|
clé aléatoire pour l'initialisation. |
obligatoire |
activation
|
str | Callable
|
activation des MLP. |
'relu'
|
aggregation
|
str
|
|
'sum'
|
Attributs :
| Nom | Type | Description |
|---|---|---|
message_mlp |
MLP
|
le MLP qui calcule les messages. |
update_mlp |
MLP
|
le MLP qui met à jour les nœuds. |
score_mlp |
MLP | list
|
le MLP des poids d'attention (vide pour les autres agrégations). |
aggregation |
str
|
le mode d'agrégation. |
Code source dans yax/layers/MessagePassing_layer.py
11 | |
41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 | |
apply
¶
apply(h, rkey=None, *, senders, receivers)
Met à jour l'état de tous les nœuds.
Paramètres :
| Nom | Type | Description | Défaut |
|---|---|---|---|
h
|
Array
|
états des nœuds, de forme |
obligatoire |
rkey
|
Array | None
|
inutilisé (couche déterministe). |
None
|
senders
|
Array
|
indices des nœuds de départ, de forme |
obligatoire |
receivers
|
Array
|
indices des nœuds d'arrivée, de forme |
obligatoire |
Renvoie :
| Type | Description |
|---|---|
Array
|
Les nouveaux états, de forme |
Code source dans yax/layers/MessagePassing_layer.py
60 61 62 63 64 65 66 67 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 | |