Aller au contenu

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 (dim_in, dim_out), initialisée selon Glorot.

bias Array

vecteur (dim_out,), initialisé à zéro.

Code source dans yax/layers/Linear.py
6
class Linear(Module):
18
19
20
21
22
23
24
25
    weight: jax.Array
    bias: jax.Array

    def __init__(self, dim_in: int, dim_out: int, rkey: jax.Array):
        rkey_w, rkey_b = jr.split(rkey)
        lim = jnp.sqrt(6.0 / (dim_in + dim_out))
        self.weight = jr.uniform(rkey_w, (dim_in, dim_out), minval=-lim, maxval=lim)
        self.bias = jnp.zeros((dim_out,))

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
def apply(self, x: jax.Array, rkey: jax.Array | None = None) -> jax.Array:
    """Calcule `x @ weight + bias` pour un exemple `x` de forme `(dim_in,)`."""
    # rkey ignoré : couche déterministe. La signature uniforme apply(x, rkey)
    # évite aux modules composites de savoir lesquels de leurs enfants
    # consomment de l'aléatoire.
    return x@self.weight + self.bias

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 yax.activations ("relu", "tanh"…) ou fonction.

obligatoire
rkey Array

clé aléatoire pour l'initialisation.

obligatoire

Attributs :

Nom Type Description
layers list[Linear]

les couches Linear.

activation_fn Callable

la fonction d'activation.

layer_sizes tuple[int, ...]

les tailles des couches.

Code source dans yax/layers/MLP.py
10
class MLP(Module):
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
    layers:list[Linear]
    activation_fn:Callable=StaticField()
    layer_sizes:tuple[int, ...]=StaticField()

    def __init__(self,layer_sizes:Sequence[int],activation:str | Callable,rkey:jax.Array):
        # tuple() garantit l'invariant meme si l'appelant passe une liste :
        # un champ statique mutable pourrait changer l'identite du modele
        # pour le cache de jit, par un append accidentel.
        self.layer_sizes=tuple(layer_sizes)
        # activation : une chaine ("tanh") ou une fonction (jax.nn.celu) —
        # voir yax/activations.py
        self.activation_fn = resolve_activation(activation)

        rkeys = jr.split(rkey, len(layer_sizes)-1)
        self.layers = [Linear(dim_in, dim_out, rkey)
                       for dim_in,dim_out,rkey in zip(layer_sizes[:-1],layer_sizes[1:],rkeys)]

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
def apply(self, x: jax.Array, rkey: jax.Array | None = None) -> jax.Array:
    """Calcule la sortie pour un exemple `x` de forme `(layer_sizes[0],)`."""
    # rkey ignoré : le MLP est déterministe (signature uniforme apply(x, rkey))
    for layer in self.layers[:-1]:
        x = layer.apply(x)
        x = self.activation_fn(x)

    return self.layers[-1].apply(x)

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 [0, 1).

obligatoire

Attributs :

Nom Type Description
rate float

la probabilité d'annulation.

Code source dans yax/layers/Dropout.py
8
class Dropout(Module):
22
23
24
25
26
    rate: float = StaticField(default_value=0.0)

    def __init__(self, rate: float):
        assert 0.0 <= rate < 1.0, f"rate:{rate} doit etre dans [0,1)"
        self.rate = rate

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
def apply(self, x: jax.Array, rkey: jax.Array | None = None) -> jax.Array:
    """Applique le dropout à `x`.

    Args:
        x: un tableau de forme quelconque.
        rkey: clé aléatoire, obligatoire en mode entraînement.
    """
    if self.inference or self.rate == 0.0:
        return x
    if rkey is None:
        raise ValueError(
            "Dropout en mode entrainement (inference=False) : une cle "
            "est requise — apply(x, rkey). Pour evaluer, basculer le "
            "modele avec model.set_inference(True).")
    keep_prob = 1.0 - self.rate
    mask = jr.bernoulli(rkey, keep_prob, x.shape)
    # division par keep_prob : l'esperance de la sortie est inchangee,
    # rien a corriger a l'evaluation ("inverted dropout").
    return jnp.where(mask, x / keep_prob, 0.0)

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, C pour une pente par canal (entrée de forme (C, ...)).

1
init float

valeur initiale de la pente.

0.25

Attributs :

Nom Type Description
a Array

les pentes, de forme (dim,).

Code source dans yax/layers/PReLU.py
7
class PReLU(Module):
18
19
20
21
22
    a: jax.Array

    def __init__(self, *, dim: int = 1, init: float = 0.25):
        # 0.25 : l'initialisation du papier — proche d'un leaky-relu classique
        self.a = jnp.full((dim,), init)

apply

apply(x, rkey=None)

Applique l'activation à x.

Code source dans yax/layers/PReLU.py
24
25
26
27
def apply(self, x: jax.Array, rkey: jax.Array | None = None) -> jax.Array:
    """Applique l'activation à `x`."""
    forme = (-1,) + (1,) * (x.ndim - 1)   # (C,) -> (C,1,1) si x est (C,H,W)
    return jnp.where(x >= 0.0, x, self.a.reshape(forme) * x)

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
class LayerNorm(Module):
22
23
24
25
26
27
28
29
30
31
    gamma: jax.Array
    beta: jax.Array
    epsilon: float = StaticField()

    def __init__(self, dim: int, *, epsilon: float = 1e-5):
        # gamma et beta redonnent au réseau la liberté d'échelle et de décalage
        # que la normalisation vient de lui retirer
        self.gamma = jnp.ones((dim,))
        self.beta = jnp.zeros((dim,))
        self.epsilon = epsilon

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
def apply(self, x: jax.Array, rkey: jax.Array | None = None) -> jax.Array:
    """Normalise `x` sur son dernier axe."""
    # x : (..., dim) — écrit pour un échantillon (dim,), mais fonctionne tel
    # quel sur (seq, dim) par broadcasting. rkey ignoré : couche déterministe.
    mean = jnp.mean(x, axis=-1, keepdims=True)
    var = jnp.var(x, axis=-1, keepdims=True)
    x_hat = (x - mean) / jnp.sqrt(var + self.epsilon)
    return self.gamma * x_hat + self.beta

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 (nb_embeddings, dim).

Code source dans yax/layers/Embedding.py
7
class Embedding(Module):
21
22
23
24
25
    weight: jax.Array

    def __init__(self, nb_embeddings: int, dim: int, rkey: jax.Array):
        # initialisation normale d'écart-type 0.02, l'usage des transformers
        self.weight = 0.02 * jr.normal(rkey, (nb_embeddings, dim))

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
def apply(self, ids: jax.Array, rkey: jax.Array | None = None) -> jax.Array:
    """Renvoie les vecteurs des identifiants entiers `ids`, de forme `ids.shape + (dim,)`."""
    # ids : entiers, forme quelconque -> sortie ids.shape + (dim,).
    # Indexation par entiers : valide sous jit (contrairement au masque
    # booléen entre crochets). rkey ignoré : couche déterministe.
    return self.weight[ids]

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

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" conserve la taille spatiale (divisée par le pas), "VALID" ne complète pas les bords.

'SAME'

Attributs :

Nom Type Description
weight Array

les noyaux, de forme (dim_out, dim_in, *noyau).

bias Array

les biais, de forme (dim_out,).

stride tuple[int, ...]

pas de la convolution.

padding str

mode de complétion des bords.

Code source dans yax/layers/Conv_nd.py
8
class Conv_nd(Module):
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
    weight: jax.Array            # (dim_out, dim_in, *kernel)
    bias: jax.Array              # (dim_out,)

    stride: tuple[int, ...] = StaticField()
    padding: str = StaticField()

    def __init__(self, dim_in: int, dim_out: int,
                 kernel_size: int | tuple[int, ...], nb_dims: int,
                 rkey: jax.Array, *,
                 stride: int | tuple[int, ...] = 1, padding: str = "SAME"):
        # padding="SAME" : la sortie garde la taille spatiale de l'entrée
        # (à stride 1) ; "VALID" : pas de remplissage, la taille fond de
        # kernel_size-1. stride>1 : l'alternative au pooling.
        # kernel_size et stride : un entier (même valeur partout) ou un tuple
        # de longueur nb_dims.
        assert padding in ("SAME", "VALID")
        if isinstance(kernel_size, int):
            kernel_size = (kernel_size,) * nb_dims
        assert len(kernel_size) == nb_dims, (
            f"kernel_size:{kernel_size} incompatible avec nb_dims:{nb_dims}")
        if isinstance(stride, int):
            stride = (stride,) * len(kernel_size)
        assert len(stride) == len(kernel_size)

        taille_noyau = 1
        for k in kernel_size:
            taille_noyau *= k
        lim = jnp.sqrt(6.0 / ((dim_in + dim_out) * taille_noyau))
        self.weight = jr.uniform(rkey, (dim_out, dim_in) + tuple(kernel_size),
                                 minval=-lim, maxval=lim)
        self.bias = jnp.zeros((dim_out,))
        self.stride = tuple(stride)
        self.padding = padding

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
def apply(self, x: jax.Array, rkey: jax.Array | None = None) -> jax.Array:
    """Calcule la convolution d'un exemple `x` de forme `(dim_in, *spatial)`."""
    # x : (canaux, *spatial) — rkey ignoré : couche déterministe.
    # Le primitif attend un batch : on en fabrique un de taille 1.
    nb_dims = self.weight.ndim - 2
    assert x.ndim == nb_dims + 1, (
        f"Conv_nd : entree de rang {x.ndim} pour une convolution a "
        f"{nb_dims} dimension(s) spatiale(s) — attendu (canaux, "
        f"{', '.join(['taille'] * nb_dims)}).")
    y = lax.conv_general_dilated(
        x[None],                        # (1, dim_in, *spatial)
        self.weight,                    # (dim_out, dim_in, *kernel)
        window_strides=self.stride,
        padding=self.padding)           # dimension_numbers par defaut :
                                        # (batch, canaux, *spatial) a tout rang
    return y[0] + self.bias.reshape((-1,) + (1,) * nb_dims)

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" ou "lstm".

'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
class RNN_layer(Module):
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
    cell: GRUCell | LSTMCell

    dim_in: int = StaticField()
    dim_hidden: int = StaticField()
    cell_type: str = StaticField()

    def __init__(self, dim_in: int, dim_hidden: int, rkey: jax.Array, *,
                 cell_type: str = "gru"):
        assert cell_type in ["gru", "lstm"]

        self.dim_in = dim_in
        self.dim_hidden = dim_hidden
        self.cell_type = cell_type

        if cell_type == "gru":
            self.cell = GRUCell(dim_in, dim_hidden, rkey)
        else:
            self.cell = LSTMCell(dim_in, dim_hidden, rkey)

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
def apply(self, xs: jax.Array, rkey: jax.Array | None = None) -> jax.Array:
    """Parcourt la séquence `xs`, de forme `(T, dim_in)`, et renvoie les états cachés, de forme `(T, dim_hidden)`."""
    # rkey ignoré : couche déterministe (signature uniforme apply(x, rkey))
    # le carry d'une GRU est h seul, celui d'une LSTM est le couple (h, c)
    zeros = jnp.zeros((self.dim_hidden,))
    carry_init: jax.Array | tuple[jax.Array, jax.Array]
    if self.cell_type == "gru":
        carry_init = zeros
        get_h = lambda carry: carry
    else:
        carry_init = (zeros, zeros)
        get_h = lambda carry: carry[0]

    def f(carry, x):
        carry = self.cell.apply(x, carry=carry)
        # rendu 2 fois : pour le pas suivant, et pour être empilé en sortie
        return carry, get_h(carry)

    _, hs = lax.scan(f, carry_init, xs)
    return hs

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
class GRUCell(Module):
31
32
33
34
35
36
37
38
39
40
41
42
    weight_ih: jax.Array   # (3H, dim_in)  — les trois portes r, z, n empilées
    weight_hh: jax.Array   # (3H, H)
    bias: jax.Array        # (3H,)
    bias_n: jax.Array      # (H,) — le biais propre au candidat n, sous la porte r

    def __init__(self, dim_in: int, dim_hidden: int, rkey: jax.Array):
        lim = 1.0 / jnp.sqrt(dim_hidden)
        rkey1, rkey2, rkey3, rkey4 = jr.split(rkey, 4)
        self.weight_ih = _uniform(rkey1, (3 * dim_hidden, dim_in), lim)
        self.weight_hh = _uniform(rkey2, (3 * dim_hidden, dim_hidden), lim)
        self.bias = _uniform(rkey3, (3 * dim_hidden,), lim)
        self.bias_n = _uniform(rkey4, (dim_hidden,), lim)

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 (dim_in,).

obligatoire
rkey Array | None

inutilisé (cellule déterministe).

None
carry Array

état caché précédent h, de forme (dim_hidden,).

obligatoire

Renvoie :

Type Description
Array

Le nouvel état caché, de forme (dim_hidden,).

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
def apply(self, x: jax.Array, rkey: jax.Array | None = None, *,
          carry: jax.Array) -> jax.Array:
    """Calcule un pas de récurrence.

    Args:
        x: entrée du pas, de forme `(dim_in,)`.
        rkey: inutilisé (cellule déterministe).
        carry: état caché précédent `h`, de forme `(dim_hidden,)`.

    Returns:
        Le nouvel état caché, de forme `(dim_hidden,)`.
    """
    h = carry
    gates_x = jnp.split(self.weight_ih @ x + self.bias, 3)
    gates_h = jnp.split(self.weight_hh @ h, 3)
    r = jax.nn.sigmoid(gates_x[0] + gates_h[0])   # porte de reinitialisation
    z = jax.nn.sigmoid(gates_x[1] + gates_h[1])   # porte de mise a jour
    n = jnp.tanh(gates_x[2] + r * (gates_h[2] + self.bias_n))  # candidat
    return (1.0 - z) * n + z * h

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
class LSTMCell(Module):
80
81
82
83
84
85
86
87
88
89
    weight_ih: jax.Array   # (4H, dim_in) — les quatre portes i, f, g, o empilées
    weight_hh: jax.Array   # (4H, H)
    bias: jax.Array        # (4H,)

    def __init__(self, dim_in: int, dim_hidden: int, rkey: jax.Array):
        lim = 1.0 / jnp.sqrt(dim_hidden)
        rkey1, rkey2, rkey3 = jr.split(rkey, 3)
        self.weight_ih = _uniform(rkey1, (4 * dim_hidden, dim_in), lim)
        self.weight_hh = _uniform(rkey2, (4 * dim_hidden, dim_hidden), lim)
        self.bias = _uniform(rkey3, (4 * dim_hidden,), lim)

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 (dim_in,).

obligatoire
rkey Array | None

inutilisé (cellule déterministe).

None
carry tuple[Array, Array]

le couple (h, c) du pas précédent.

obligatoire

Renvoie :

Type Description
tuple[Array, Array]

Le nouveau couple (h, c).

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
def apply(self, x: jax.Array, rkey: jax.Array | None = None, *,
          carry: tuple[jax.Array, jax.Array]) -> tuple[jax.Array, jax.Array]:
    """Calcule un pas de récurrence.

    Args:
        x: entrée du pas, de forme `(dim_in,)`.
        rkey: inutilisé (cellule déterministe).
        carry: le couple `(h, c)` du pas précédent.

    Returns:
        Le nouveau couple `(h, c)`.
    """
    h, c = carry
    gates = jnp.split(self.weight_ih @ x + self.bias + self.weight_hh @ h, 4)
    i = jax.nn.sigmoid(gates[0])    # porte d'entree
    f = jax.nn.sigmoid(gates[1])    # porte d'oubli
    g = jnp.tanh(gates[2])          # candidat
    o = jax.nn.sigmoid(gates[3])    # porte de sortie
    c = f * c + i * g
    h = o * jnp.tanh(c)
    return h, c

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

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
class MultiHeadAttention(Module):
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
    q_proj: Linear
    k_proj: Linear
    v_proj: Linear
    out_proj: Linear
    dropout: Dropout
    nb_heads: int = StaticField()

    def __init__(self, dim: int, nb_heads: int, rkey: jax.Array, *,
                 dropout_rate: float = 0.0):
        assert dim % nb_heads == 0, f"dim:{dim} doit etre divisible par nb_heads:{nb_heads}"
        rkey_q, rkey_k, rkey_v, rkey_o = jr.split(rkey, 4)
        self.q_proj = Linear(dim, dim, rkey_q)
        self.k_proj = Linear(dim, dim, rkey_k)
        self.v_proj = Linear(dim, dim, rkey_v)
        self.out_proj = Linear(dim, dim, rkey_o)
        self.dropout = Dropout(dropout_rate)
        self.nb_heads = nb_heads

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 (seq_q, dim).

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 (seq_kv, dim) ; None pour l'auto-attention.

None
mask Array | None

masque additif (0 visible, -inf masqué), de forme compatible avec (seq_q, seq_kv) ; voir yax.preprocessing.causal_mask.

None

Renvoie :

Type Description
Array

Un tableau de forme (seq_q, dim).

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
def apply(self, x: jax.Array, rkey: jax.Array | None = None, *,
          x_kv: jax.Array | None = None,
          mask: jax.Array | None = None) -> jax.Array:
    """Calcule l'attention pour une séquence.

    Args:
        x: la séquence des requêtes, de forme `(seq_q, dim)`.
        rkey: clé aléatoire pour le dropout, en mode entraînement.
        x_kv: la séquence des clés et valeurs, de forme `(seq_kv, dim)` ;
            `None` pour l'auto-attention.
        mask: masque additif (0 visible, `-inf` masqué), de forme compatible
            avec `(seq_q, seq_kv)` ; voir `yax.preprocessing.causal_mask`.

    Returns:
        Un tableau de forme `(seq_q, dim)`.
    """
    if x_kv is None:
        x_kv = x  # auto-attention
    seq_q, dim = x.shape
    head_dim = dim // self.nb_heads

    Q = self._split_heads(self.q_proj.apply(x))       # (heads, seq_q, head_dim)
    K = self._split_heads(self.k_proj.apply(x_kv))    # (heads, seq_kv, head_dim)
    V = self._split_heads(self.v_proj.apply(x_kv))

    # normalisation par sqrt(head_dim) : contrôle de la variance du produit
    # scalaire, sinon le softmax sature dès que head_dim est grand
    scores = Q @ jnp.transpose(K, (0, 2, 1)) / jnp.sqrt(head_dim)
    if mask is not None:
        scores = scores + mask                        # -inf avant le softmax
    weights = jax.nn.softmax(scores, axis=-1)         # (heads, seq_q, seq_kv)
    weights = self.dropout.apply(weights, rkey)

    context = weights @ V                             # (heads, seq_q, head_dim)
    context = jnp.transpose(context, (1, 0, 2)).reshape(seq_q, dim)
    return self.out_proj.apply(context)

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 (nb_heads, seq_q, seq_kv) ; chaque ligne somme à 1.

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
def attention_maps(self, x: jax.Array, *, x_kv: jax.Array | None = None,
                   mask: jax.Array | None = None) -> jax.Array:
    """Renvoie les poids d'attention, pour les visualiser.

    Mêmes arguments que `apply`, sans dropout.

    Returns:
        Les poids, de forme `(nb_heads, seq_q, seq_kv)` ; chaque ligne somme à 1.
    """
    if x_kv is None:
        x_kv = x
    head_dim = x.shape[-1] // self.nb_heads
    Q = self._split_heads(self.q_proj.apply(x))
    K = self._split_heads(self.k_proj.apply(x_kv))
    scores = Q @ jnp.transpose(K, (0, 2, 1)) / jnp.sqrt(head_dim)
    if mask is not None:
        scores = scores + mask
    return jax.nn.softmax(scores, axis=-1)

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" ou "post".

'pre'

Attributs :

Nom Type Description
attention MultiHeadAttention

la couche d'attention multi-têtes.

feed_forward MLP

le MLP dim → dim_ff → dim, activation gelu.

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
class TransformerBlock(Module):
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
    attention: MultiHeadAttention
    feed_forward: MLP
    norm1: LayerNorm
    norm2: LayerNorm
    dropout: Dropout
    norm_position: str = StaticField()

    def __init__(self, dim: int, nb_heads: int, dim_ff: int, rkey: jax.Array, *,
                 dropout_rate: float = 0.0, norm_position: str = "pre"):
        assert norm_position in ("pre", "post")
        rkey_att, rkey_ff = jr.split(rkey)
        self.attention = MultiHeadAttention(dim, nb_heads, rkey_att,
                                            dropout_rate=dropout_rate)
        # le feed-forward est un simple MLP appliqué position par position :
        # écrit pour (dim,), il broadcast tel quel sur (seq, dim)
        self.feed_forward = MLP((dim, dim_ff, dim), "gelu", rkey_ff)
        self.norm1 = LayerNorm(dim)
        self.norm2 = LayerNorm(dim)
        self.dropout = Dropout(dropout_rate)
        self.norm_position = norm_position

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 (seq, dim).

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
def apply(self, x: jax.Array, rkey: jax.Array | None = None, *,
          mask: jax.Array | None = None) -> jax.Array:
    """Transforme une séquence.

    Args:
        x: la séquence, de forme `(seq, dim)`.
        rkey: clé aléatoire pour le dropout, en mode entraînement.
        mask: masque additif transmis à l'attention.

    Returns:
        Une séquence de même forme.
    """
    # x : (seq, dim)
    rkey_att, rkey_d1, rkey_d2 = (None, None, None) if rkey is None else jr.split(rkey, 3)

    if self.norm_position == "pre":
        x = x + self.dropout.apply(
            self.attention.apply(self.norm1.apply(x), rkey_att, mask=mask), rkey_d1)
        x = x + self.dropout.apply(
            self.feed_forward.apply(self.norm2.apply(x)), rkey_d2)
    else:  # post
        x = self.norm1.apply(
            x + self.dropout.apply(self.attention.apply(x, rkey_att, mask=mask), rkey_d1))
        x = self.norm2.apply(
            x + self.dropout.apply(self.feed_forward.apply(x), rkey_d2))
    return x

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", "mean", "max" ou "attention".

'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
class MessagePassing_layer(Module):
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
    message_mlp: MLP
    update_mlp: MLP
    score_mlp: MLP | list    # le réseau de score si aggregation="attention", sinon []
    aggregation: str = StaticField()

    def __init__(self, dim_node: int, dim_hidden: int, dim_out: int,
                 rkey: jax.Array, *,
                 activation: str | Callable = "relu", aggregation: str = "sum"):
        assert aggregation in ("sum", "mean", "max", "attention"), (
            f"aggregation:{aggregation!r} — choix : sum, mean, max, attention")
        rkey_msg, rkey_upd, rkey_score = jr.split(rkey, 3)
        self.message_mlp = MLP((2 * dim_node, dim_hidden, dim_hidden),
                               activation, rkey_msg)
        self.update_mlp = MLP((dim_node + dim_hidden, dim_hidden, dim_out),
                              activation, rkey_upd)
        self.score_mlp = (MLP((2 * dim_node, dim_hidden, 1), activation, rkey_score)
                          if aggregation == "attention" else [])
        self.aggregation = aggregation

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 (nb_noeuds, dim_node).

obligatoire
rkey Array | None

inutilisé (couche déterministe).

None
senders Array

indices des nœuds de départ, de forme (nb_aretes,).

obligatoire
receivers Array

indices des nœuds d'arrivée, de forme (nb_aretes,).

obligatoire

Renvoie :

Type Description
Array

Les nouveaux états, de forme (nb_noeuds, dim_out).

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
def apply(self, h: jax.Array, rkey: jax.Array | None = None, *,
          senders: jax.Array, receivers: jax.Array) -> jax.Array:
    """Met à jour l'état de tous les nœuds.

    Args:
        h: états des nœuds, de forme `(nb_noeuds, dim_node)`.
        rkey: inutilisé (couche déterministe).
        senders: indices des nœuds de départ, de forme `(nb_aretes,)`.
        receivers: indices des nœuds d'arrivée, de forme `(nb_aretes,)`.

    Returns:
        Les nouveaux états, de forme `(nb_noeuds, dim_out)`.
    """
    nb_noeuds = h.shape[0]
    # 1. un message par arête, fonction des deux extrémités
    edge_features = jnp.concatenate([h[senders], h[receivers]], axis=-1)
    messages = self.message_mlp.apply(edge_features)      # (nb_aretes, dim_hidden)

    # 2. agrégation des messages reçus par chaque nœud (dispersion) ;
    # num_segments est statique sous jit : c'est h.shape[0]
    if self.aggregation == "max":
        aggregated = jax.ops.segment_max(messages, receivers,
                                         num_segments=nb_noeuds)
        # un nœud sans arête entrante recevrait -inf : il ne reçoit rien
        aggregated = jnp.where(jnp.isneginf(aggregated), 0.0, aggregated)
    elif self.aggregation == "attention":
        scores = self.score_mlp.apply(edge_features)[:, 0]     # type: ignore[union-attr]  # (nb_aretes,)
        # softmax PAR RÉCEPTEUR, stabilisé par le max du segment ; un nœud
        # isolé n'indexe jamais son dénominateur (nul) : il ne reçoit rien
        scores_max = jax.ops.segment_max(scores, receivers,
                                         num_segments=nb_noeuds)
        expo = jnp.exp(scores - scores_max[receivers])
        denominateur = jax.ops.segment_sum(expo, receivers,
                                           num_segments=nb_noeuds)
        alpha = expo / denominateur[receivers]
        aggregated = jax.ops.segment_sum(alpha[:, None] * messages, receivers,
                                         num_segments=nb_noeuds)
    else:
        aggregated = jax.ops.segment_sum(messages, receivers,
                                         num_segments=nb_noeuds)
        if self.aggregation == "mean":
            degrees = jax.ops.segment_sum(
                jnp.ones_like(receivers, dtype=h.dtype),
                receivers, num_segments=nb_noeuds)
            aggregated = aggregated / jnp.maximum(degrees, 1.0)[:, None]

    # 3. mise à jour de l'état du nœud
    return self.update_mlp.apply(jnp.concatenate([h, aggregated], axis=-1))