Aller au contenu

yax.preprocessing

Des transformations de données, sans paramètre appris.

Séries temporelles

Fenêtrage de séries temporelles, pour la prévision.

sliding_windows

yax.preprocessing.sliding_windows(series, window, horizon, *, stride=1)

Découpe une série temporelle en couples (fenêtre, suite), pour la prévision.

Chaque exemple est une fenêtre de window pas consécutifs ; sa cible est formée des horizon pas qui la suivent. Deux fenêtres successives sont décalées de stride pas.

série     : s0 s1 s2 s3 s4 s5 s6        window=3, horizon=2, stride=1
exemple 0 : [s0 s1 s2] -> [s3 s4]
exemple 1 :    [s1 s2 s3] -> [s4 s5]
exemple 2 :       [s2 s3 s4] -> [s5 s6]

Découper la série en tranches d'entraînement, de validation et de test avant le fenêtrage, puis fenêtrer chaque tranche séparément.

Paramètres :

Nom Type Description Défaut
series ArrayLike

la série, de forme (T,) ou (T, nb_variables).

obligatoire
window int

longueur d'une fenêtre.

obligatoire
horizon int

nombre de pas à prévoir ; 0 est permis (fenêtres seules).

obligatoire
stride int

décalage entre deux fenêtres.

1

Renvoie :

Type Description
tuple[Array, Array]

Le couple (X, Y), avec X de forme (nb, window, ...) et Y de forme (nb, horizon, ...), où nb = (T - window - horizon) // stride + 1.

Code source dans yax/preprocessing/windows.py
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
def sliding_windows(series: ArrayLike, window: int, horizon: int, *,
                    stride: int = 1) -> tuple[jax.Array, jax.Array]:
    """Découpe une série temporelle en couples (fenêtre, suite), pour la prévision.

    Chaque exemple est une fenêtre de `window` pas consécutifs ; sa cible est
    formée des `horizon` pas qui la suivent. Deux fenêtres successives sont
    décalées de `stride` pas.

    ```
    série     : s0 s1 s2 s3 s4 s5 s6        window=3, horizon=2, stride=1
    exemple 0 : [s0 s1 s2] -> [s3 s4]
    exemple 1 :    [s1 s2 s3] -> [s4 s5]
    exemple 2 :       [s2 s3 s4] -> [s5 s6]
    ```

    Découper la série en tranches d'entraînement, de validation et de test
    **avant** le fenêtrage, puis fenêtrer chaque tranche séparément.

    Args:
        series: la série, de forme `(T,)` ou `(T, nb_variables)`.
        window: longueur d'une fenêtre.
        horizon: nombre de pas à prévoir ; 0 est permis (fenêtres seules).
        stride: décalage entre deux fenêtres.

    Returns:
        Le couple `(X, Y)`, avec `X` de forme `(nb, window, ...)` et `Y` de forme `(nb, horizon, ...)`, où `nb = (T - window - horizon) // stride + 1`.
    """
    series = jnp.asarray(series)
    T = series.shape[0]
    assert window >= 1 and horizon >= 0 and stride >= 1, (window, horizon, stride)
    nb = (T - window - horizon) // stride + 1
    assert nb >= 1, (f"série de longueur {T} : trop courte pour window={window} "
                     f"et horizon={horizon}")
    # (nb, window + horizon) : la position de chaque pas de chaque exemple
    indices = jnp.arange(nb)[:, None] * stride + jnp.arange(window + horizon)[None, :]
    blocs = series[indices]                     # (nb, window + horizon, ...)
    return blocs[:, :window], blocs[:, window:]

Séquences et attention

sinusoidal_positional_encoding

yax.preprocessing.sinusoidal_positional_encoding(seq_len, dim)

Encodage positionnel sinusoïdal, à ajouter aux vecteurs d'une séquence.

L'attention ne tient pas compte de l'ordre des éléments ; cet encodage le réintroduit. Chaque position est représentée par des sinus et cosinus à plusieurs fréquences : les hautes distinguent les positions voisines, les basses situent dans la séquence entière. Rien n'est appris ; la version apprise est un yax.layers.Embedding(seq_len, dim) indexé par la position.

Paramètres :

Nom Type Description Défaut
seq_len int

longueur de la séquence.

obligatoire
dim int

taille des vecteurs ; doit être paire.

obligatoire

Renvoie :

Type Description
Array

Un tableau (seq_len, dim), avec les sinus sur la première moitié des colonnes et les cosinus sur la seconde.

Code source dans yax/preprocessing/positional_encoding.py
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
def sinusoidal_positional_encoding(seq_len: int, dim: int) -> jax.Array:
    """Encodage positionnel sinusoïdal, à ajouter aux vecteurs d'une séquence.

    L'attention ne tient pas compte de l'ordre des éléments ; cet encodage le
    réintroduit. Chaque position est représentée par des sinus et cosinus à
    plusieurs fréquences : les hautes distinguent les positions voisines, les
    basses situent dans la séquence entière. Rien n'est appris ; la version
    apprise est un `yax.layers.Embedding(seq_len, dim)` indexé par la position.

    Args:
        seq_len: longueur de la séquence.
        dim: taille des vecteurs ; doit être paire.

    Returns:
        Un tableau `(seq_len, dim)`, avec les sinus sur la première moitié des colonnes et les cosinus sur la seconde.
    """
    assert dim % 2 == 0, f"dim:{dim} doit etre pair"
    positions = jnp.arange(seq_len)[:, None]                    # (seq_len, 1)
    freqs = jnp.arange(dim // 2)[None, :]                       # (1, dim/2)
    angles = positions / (10000.0 ** (2 * freqs / dim))         # (seq_len, dim/2)
    return jnp.concatenate([jnp.sin(angles), jnp.cos(angles)], axis=-1)

causal_mask

yax.preprocessing.causal_mask(seq_len)

Masque causal : chaque position ne voit que les positions précédentes et elle-même.

Paramètres :

Nom Type Description Défaut
seq_len int

longueur de la séquence.

obligatoire

Renvoie :

Type Description
Array

Un tableau (seq_len, seq_len), nul sur et sous la diagonale, égal à -inf au-dessus.

Code source dans yax/preprocessing/masks.py
11
12
13
14
15
16
17
18
19
20
21
def causal_mask(seq_len: int) -> jax.Array:
    """Masque causal : chaque position ne voit que les positions précédentes et elle-même.

    Args:
        seq_len: longueur de la séquence.

    Returns:
        Un tableau `(seq_len, seq_len)`, nul sur et sous la diagonale, égal à `-inf` au-dessus.
    """
    visible = jnp.tril(jnp.ones((seq_len, seq_len), dtype=bool))
    return jnp.where(visible, 0.0, -jnp.inf)

padding_mask

yax.preprocessing.padding_mask(is_real_token)

Masque de remplissage : les positions de remplissage ne sont vues par aucune position.

Paramètres :

Nom Type Description Défaut
is_real_token ArrayLike

booléens de forme (seq_len,), True pour un vrai élément, False pour du remplissage.

obligatoire

Renvoie :

Type Description
Array

Un tableau (1, seq_len), valable pour toutes les requêtes.

Code source dans yax/preprocessing/masks.py
24
25
26
27
28
29
30
31
32
33
34
def padding_mask(is_real_token: ArrayLike) -> jax.Array:
    """Masque de remplissage : les positions de remplissage ne sont vues par aucune position.

    Args:
        is_real_token: booléens de forme `(seq_len,)`, `True` pour un vrai
            élément, `False` pour du remplissage.

    Returns:
        Un tableau `(1, seq_len)`, valable pour toutes les requêtes.
    """
    return jnp.where(is_real_token, 0.0, -jnp.inf)[None, :]

Images

Enrichissement d'images (augmentation de données).

Des transformations aléatoires qui conservent l'étiquette, pour une image de forme (canaux, H, W). Elles sont compatibles avec jax.jit et jax.vmap : on les applique dans le pas d'entraînement, avec une clé différente à chaque époque. Toutes ont la signature f(rkey, img, ...).

Une transformation doit préserver le sens de l'image : une symétrie horizontale convient à des photos d'animaux, pas à des chiffres manuscrits.

random_augmentation

yax.preprocessing.random_augmentation(rkey, img)

Exemple de chaîne d'augmentations, qui compose plusieurs transformations de ce module.

À adapter à ses données. La symétrie horizontale n'y figure pas, car elle ne convient pas à toutes les images.

Paramètres :

Nom Type Description Défaut
rkey Array

clé aléatoire.

obligatoire
img Array

l'image, de forme (canaux, H, W).

obligatoire
Code source dans yax/preprocessing/augmentation.py
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
def random_augmentation(rkey: jax.Array, img: jax.Array) -> jax.Array:
    """Exemple de chaîne d'augmentations, qui compose plusieurs transformations de ce module.

    À adapter à ses données. La symétrie horizontale n'y figure pas, car elle
    ne convient pas à toutes les images.

    Args:
        rkey: clé aléatoire.
        img: l'image, de forme `(canaux, H, W)`.
    """
    rkey1, rkey2, rkey3, rkey4 = jr.split(rkey, 4)
    img = random_scale_translate(rkey1, img)
    img = random_rotation(rkey2, img)
    img = random_brightness_contrast(rkey3, img)
    img = random_noise(rkey4, img)
    return img

random_horizontal_flip

yax.preprocessing.random_horizontal_flip(rkey, img, *, prob=0.5)

Symétrie gauche-droite, appliquée avec la probabilité prob.

Paramètres :

Nom Type Description Défaut
rkey Array

clé aléatoire.

obligatoire
img Array

l'image, de forme (canaux, H, W).

obligatoire
prob float

probabilité d'appliquer la symétrie.

0.5
Code source dans yax/preprocessing/augmentation.py
17
18
19
20
21
22
23
24
25
26
27
28
def random_horizontal_flip(rkey: jax.Array, img: jax.Array, *,
                           prob: float = 0.5) -> jax.Array:
    """Symétrie gauche-droite, appliquée avec la probabilité `prob`.

        Args:
            rkey: clé aléatoire.
            img: l'image, de forme `(canaux, H, W)`.
            prob: probabilité d'appliquer la symétrie.
    """
    # jnp.where et non un `if` : le tirage est une valeur tracée
    flip = jr.bernoulli(rkey, prob)
    return jnp.where(flip, img[:, :, ::-1], img)

random_rotation

yax.preprocessing.random_rotation(rkey, img, *, max_angle_degrees=15.0)

Rotation d'angle aléatoire autour du centre, avec interpolation linéaire.

Paramètres :

Nom Type Description Défaut
rkey Array

clé aléatoire.

obligatoire
img Array

l'image, de forme (canaux, H, W).

obligatoire
max_angle_degrees float

angle maximal, en degrés.

15.0
Code source dans yax/preprocessing/augmentation.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
def random_rotation(rkey: jax.Array, img: jax.Array, *,
                    max_angle_degrees: float = 15.0) -> jax.Array:
    """Rotation d'angle aléatoire autour du centre, avec interpolation linéaire.

        Args:
            rkey: clé aléatoire.
            img: l'image, de forme `(canaux, H, W)`.
            max_angle_degrees: angle maximal, en degrés.
    """
    C, H, W = img.shape
    angle = jnp.deg2rad(jr.uniform(rkey, minval=-max_angle_degrees,
                                   maxval=max_angle_degrees))
    cos, sin = jnp.cos(angle), jnp.sin(angle)
    rows = jnp.arange(H) - (H - 1) / 2.0
    cols = jnp.arange(W) - (W - 1) / 2.0
    r, c = jnp.meshgrid(rows, cols, indexing="ij")
    src_r = cos * r - sin * c + (H - 1) / 2.0
    src_c = sin * r + cos * c + (W - 1) / 2.0

    def rotate_channel(channel):
        return jax.scipy.ndimage.map_coordinates(channel, [src_r, src_c],
                                                 order=1, mode="constant", cval=0.0)

    return jax.vmap(rotate_channel)(img)

random_scale_translate

yax.preprocessing.random_scale_translate(rkey, img, *, max_zoom=0.2, max_shift=2.0)

Zoom autour du centre et translation, aléatoires.

Paramètres :

Nom Type Description Défaut
rkey Array

clé aléatoire.

obligatoire
img Array

l'image, de forme (canaux, H, W).

obligatoire
max_zoom float

variation relative maximale de l'échelle.

0.2
max_shift float

translation maximale, en pixels.

2.0
Code source dans yax/preprocessing/augmentation.py
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
def random_scale_translate(rkey: jax.Array, img: jax.Array, *, max_zoom: float = 0.2,
                           max_shift: float = 2.0) -> jax.Array:
    """Zoom autour du centre et translation, aléatoires.

        Args:
            rkey: clé aléatoire.
            img: l'image, de forme `(canaux, H, W)`.
            max_zoom: variation relative maximale de l'échelle.
            max_shift: translation maximale, en pixels.
    """
    C, H, W = img.shape
    rkey_zoom, rkey_shift = jr.split(rkey)
    zoom = 1.0 + jr.uniform(rkey_zoom, minval=-max_zoom, maxval=max_zoom)
    shift = jr.uniform(rkey_shift, (2,), minval=-max_shift, maxval=max_shift)
    # scale_and_translate applique sortie(y) = entree((y - translation)/scale) :
    # cette translation-ci recentre le zoom sur le milieu de l'image
    scale = jnp.array([zoom, zoom])
    translation = (1.0 - zoom) * jnp.array([H, W]) / 2.0 + shift
    return jax.image.scale_and_translate(img, img.shape, (1, 2),
                                         scale, translation, method="linear")

random_brightness_contrast

yax.preprocessing.random_brightness_contrast(rkey, img, *, max_shift=0.2, max_factor=0.2)

Modifie aléatoirement la luminosité et le contraste.

Paramètres :

Nom Type Description Défaut
rkey Array

clé aléatoire.

obligatoire
img Array

l'image, de forme (canaux, H, W).

obligatoire
max_shift float

décalage maximal de la luminosité.

0.2
max_factor float

variation relative maximale du contraste.

0.2
Code source dans yax/preprocessing/augmentation.py
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
def random_brightness_contrast(rkey: jax.Array, img: jax.Array, *, max_shift: float = 0.2,
                               max_factor: float = 0.2) -> jax.Array:
    """Modifie aléatoirement la luminosité et le contraste.

        Args:
            rkey: clé aléatoire.
            img: l'image, de forme `(canaux, H, W)`.
            max_shift: décalage maximal de la luminosité.
            max_factor: variation relative maximale du contraste.
    """
    rkey_shift, rkey_factor = jr.split(rkey)
    shift = jr.uniform(rkey_shift, minval=-max_shift, maxval=max_shift)
    factor = 1.0 + jr.uniform(rkey_factor, minval=-max_factor, maxval=max_factor)
    mean = jnp.mean(img)
    return (img - mean) * factor + mean + shift

random_noise

yax.preprocessing.random_noise(rkey, img, *, sigma=0.05)

Ajoute un bruit gaussien.

Paramètres :

Nom Type Description Défaut
rkey Array

clé aléatoire.

obligatoire
img Array

l'image, de forme (canaux, H, W).

obligatoire
sigma float

écart-type du bruit.

0.05
Code source dans yax/preprocessing/augmentation.py
48
49
50
51
52
53
54
55
56
def random_noise(rkey: jax.Array, img: jax.Array, *, sigma: float = 0.05) -> jax.Array:
    """Ajoute un bruit gaussien.

        Args:
            rkey: clé aléatoire.
            img: l'image, de forme `(canaux, H, W)`.
            sigma: écart-type du bruit.
    """
    return img + sigma * jr.normal(rkey, img.shape)