Aller au contenu

yax.obj

Objectifs : les fonctions que yax.training.train minimise.

Un objectif a la signature objective_fn(model, x, y, rkey=None) : il reçoit le modèle et un lot (x, y), et renvoie un scalaire.

Il est très facile de construire des objectifs personnalisés, notamment en utilisant les fonctions de perte d'optax. Par exemple, yax.obj.softmax_cross_entropy_from_integers est l'exact équivalent de :

def mon_objectif(model, x, y, rkey=None):
    logits = yax.batch_apply(model, x, rkey)
    return jnp.mean(optax.softmax_cross_entropy_with_integer_labels(logits, y))

Les objectifs de classification prennent des logits. Le modèle renvoie des scores réels : sa dernière couche ne porte ni sigmoïde ni softmax. Le calcul en est plus stable, même pour de très grands scores. Pour lire des probabilités à l'évaluation, appliquer soi-même jax.nn.sigmoid (deux classes, ou multi-label) ou jax.nn.softmax (classes exclusives) aux sorties du modèle. Cela vaut aussi pour dice_loss, focal_loss et intersection_over_union_loss, qui appliquent la sigmoïde elles-mêmes.

Options. Un objectif s'utilise tel quel, avec ses options par défaut, ou se configure en l'appelant avec ses options nommées :

yax.obj.huber_loss                 # delta=1.0
yax.obj.huber_loss(delta=2.0)      # un nouvel objectif, delta=2.0

Le décorateur yax.obj.objective rend ainsi configurables les options de vos propres objectifs ; sans option, une simple fonction suffit.

Noms. Chaque objectif porte un nom complet, et un alias court quand il est long : yax.obj.mean_squared_error s'écrit aussi yax.obj.mse. Le suffixe _loss n'est gardé que lorsque le nom nu désigne autre chose qu'une perte : dice_loss vaut 1 - dice, et dice est un score que l'on maximise.

Régression

mean_squared_error

yax.obj.mean_squared_error(model, x, y, rkey=None)

Erreur quadratique moyenne, pour la régression. Alias : yax.obj.mse.

Paramètres :

Nom Type Description Défaut
model Module

le modèle.

obligatoire
x ArrayLike

les entrées du lot.

obligatoire
y ArrayLike

les cibles, de même forme que les sorties du modèle.

obligatoire
rkey Array | None

clé aléatoire (dropout), ou None.

None
Code source dans yax/obj.py
111
112
113
114
115
116
117
118
119
120
121
122
123
@objective
def mean_squared_error(model: Module, x: ArrayLike, y: ArrayLike,
                       rkey: jax.Array | None = None) -> jax.Array:
    """Erreur quadratique moyenne, pour la régression. Alias : `yax.obj.mse`.

    Args:
        model: le modèle.
        x: les entrées du lot.
        y: les cibles, de même forme que les sorties du modèle.
        rkey: clé aléatoire (dropout), ou `None`.
    """
    y_pred = batch_apply(model, x, rkey)
    return jnp.mean((y_pred - y) ** 2)

mean_absolute_error

yax.obj.mean_absolute_error(model, x, y, rkey=None)

Erreur absolue moyenne, pour la régression. Alias : yax.obj.mae.

Un grand écart y pèse proportionnellement à sa taille, et non à son carré : les mesures aberrantes tirent beaucoup moins le modèle qu'avec mean_squared_error.

Paramètres :

Nom Type Description Défaut
model Module

le modèle.

obligatoire
x ArrayLike

les entrées du lot.

obligatoire
y ArrayLike

les cibles, de même forme que les sorties du modèle.

obligatoire
rkey Array | None

clé aléatoire (dropout), ou None.

None
Code source dans yax/obj.py
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
@objective
def mean_absolute_error(model: Module, x: ArrayLike, y: ArrayLike,
                        rkey: jax.Array | None = None) -> jax.Array:
    """Erreur absolue moyenne, pour la régression. Alias : `yax.obj.mae`.

    Un grand écart y pèse proportionnellement à sa taille, et non à son carré :
    les mesures aberrantes tirent beaucoup moins le modèle qu'avec
    `mean_squared_error`.

    Args:
        model: le modèle.
        x: les entrées du lot.
        y: les cibles, de même forme que les sorties du modèle.
        rkey: clé aléatoire (dropout), ou `None`.
    """
    y_pred = batch_apply(model, x, rkey)
    return jnp.mean(jnp.abs(y_pred - y))

huber_loss

yax.obj.huber_loss(model, x, y, rkey=None, *, delta=1.0)

Perte de Huber, pour la régression : un compromis entre les deux précédentes.

Quadratique pour les écarts plus petits que delta, linéaire au-delà : elle garde la douceur de mean_squared_error près de la cible, sans se laisser dominer par les valeurs aberrantes. Se configure par yax.obj.huber_loss(delta=2.0).

Paramètres :

Nom Type Description Défaut
model Module

le modèle.

obligatoire
x ArrayLike

les entrées du lot.

obligatoire
y ArrayLike

les cibles, de même forme que les sorties du modèle.

obligatoire
rkey Array | None

clé aléatoire (dropout), ou None.

None
delta float

l'écart où l'on passe du quadratique au linéaire.

1.0
Code source dans yax/obj.py
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
@objective
def huber_loss(model: Module, x: ArrayLike, y: ArrayLike, rkey: jax.Array | None = None, *,
               delta: float = 1.0) -> jax.Array:
    """Perte de Huber, pour la régression : un compromis entre les deux précédentes.

    Quadratique pour les écarts plus petits que `delta`, linéaire au-delà :
    elle garde la douceur de `mean_squared_error` près de la cible, sans se
    laisser dominer par les valeurs aberrantes. Se configure par
    `yax.obj.huber_loss(delta=2.0)`.

    Args:
        model: le modèle.
        x: les entrées du lot.
        y: les cibles, de même forme que les sorties du modèle.
        rkey: clé aléatoire (dropout), ou `None`.
        delta: l'écart où l'on passe du quadratique au linéaire.
    """
    y_pred = batch_apply(model, x, rkey)
    return jnp.mean(optax.huber_loss(y_pred, jnp.asarray(y), delta=delta))

Classification

binary_cross_entropy

yax.obj.binary_cross_entropy(model, x, y, rkey=None)

Entropie croisée binaire, sur les logits. Alias : yax.obj.bce.

Pour la classification à deux classes, la segmentation binaire (une prédiction par pixel), et la classification multi-label, où plusieurs classes sont vraies à la fois : un logit et une sigmoïde par classe.

Paramètres :

Nom Type Description Défaut
model Module

le modèle, dont la sortie ne passe PAS par une sigmoïde.

obligatoire
x ArrayLike

les entrées du lot.

obligatoire
y ArrayLike

les étiquettes 0 ou 1, de même forme que les sorties.

obligatoire
rkey Array | None

clé aléatoire (dropout), ou None.

None
Code source dans yax/obj.py
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
@objective
def binary_cross_entropy(model: Module, x: ArrayLike, y: ArrayLike,
                         rkey: jax.Array | None = None) -> jax.Array:
    """Entropie croisée binaire, sur les **logits**. Alias : `yax.obj.bce`.

    Pour la classification à deux classes, la segmentation binaire (une
    prédiction par pixel), et la classification multi-label, où plusieurs
    classes sont vraies à la fois : un logit et une sigmoïde par classe.

    Args:
        model: le modèle, dont la sortie ne passe PAS par une sigmoïde.
        x: les entrées du lot.
        y: les étiquettes 0 ou 1, de même forme que les sorties.
        rkey: clé aléatoire (dropout), ou `None`.
    """
    logits = batch_apply(model, x, rkey)
    return jnp.mean(optax.sigmoid_binary_cross_entropy(logits, y))

softmax_cross_entropy_from_integers

yax.obj.softmax_cross_entropy_from_integers(model, x, y, rkey=None)

Entropie croisée multi-classes, sur les logits, à cibles entières.

Alias : yax.obj.softmax_ce_from_integers. Pour des classes exclusives : chaque exemple appartient à une seule. Convient à la classification (logits de forme (nb_classes,) par exemple) comme aux modèles de langue (logits (seq_len, vocab) contre y de forme (seq_len,)).

Paramètres :

Nom Type Description Défaut
model Module

le modèle, dont la sortie ne passe PAS par une softmax.

obligatoire
x ArrayLike

les entrées du lot.

obligatoire
y ArrayLike

les étiquettes entières, entre 0 et nb_classes - 1.

obligatoire
rkey Array | None

clé aléatoire (dropout), ou None.

None
Code source dans yax/obj.py
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
@objective
def softmax_cross_entropy_from_integers(model: Module, x: ArrayLike, y: ArrayLike,
                                        rkey: jax.Array | None = None) -> jax.Array:
    """Entropie croisée multi-classes, sur les **logits**, à cibles entières.

    Alias : `yax.obj.softmax_ce_from_integers`. Pour des classes exclusives :
    chaque exemple appartient à une seule. Convient à la classification
    (logits de forme `(nb_classes,)` par exemple) comme aux modèles de langue
    (logits `(seq_len, vocab)` contre `y` de forme `(seq_len,)`).

    Args:
        model: le modèle, dont la sortie ne passe PAS par une softmax.
        x: les entrées du lot.
        y: les étiquettes entières, entre 0 et `nb_classes - 1`.
        rkey: clé aléatoire (dropout), ou `None`.
    """
    logits = batch_apply(model, x, rkey)
    return jnp.mean(optax.softmax_cross_entropy_with_integer_labels(logits, y))

softmax_cross_entropy_from_vectors

yax.obj.softmax_cross_entropy_from_vectors(model, x, y, rkey=None)

Entropie croisée multi-classes, sur les logits, à cibles vectorielles.

Alias : yax.obj.softmax_ce_from_vectors. Les cibles sont des vecteurs de probabilités, de somme 1 : codage one-hot, étiquettes adoucies (label smoothing), ou distribution d'un modèle enseignant (distillation). Pour des cibles où plusieurs classes valent 1 à la fois (multi-label), c'est binary_cross_entropy qu'il faut.

Paramètres :

Nom Type Description Défaut
model Module

le modèle, dont la sortie ne passe PAS par une softmax.

obligatoire
x ArrayLike

les entrées du lot.

obligatoire
y ArrayLike

les vecteurs de probabilités, de même forme que les sorties.

obligatoire
rkey Array | None

clé aléatoire (dropout), ou None.

None
Code source dans yax/obj.py
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
@objective
def softmax_cross_entropy_from_vectors(model: Module, x: ArrayLike, y: ArrayLike,
                                       rkey: jax.Array | None = None) -> jax.Array:
    """Entropie croisée multi-classes, sur les **logits**, à cibles vectorielles.

    Alias : `yax.obj.softmax_ce_from_vectors`. Les cibles sont des vecteurs de
    probabilités, de somme 1 : codage one-hot, étiquettes adoucies (label
    smoothing), ou distribution d'un modèle enseignant (distillation). Pour des
    cibles où plusieurs classes valent 1 à la fois (multi-label), c'est
    `binary_cross_entropy` qu'il faut.

    Args:
        model: le modèle, dont la sortie ne passe PAS par une softmax.
        x: les entrées du lot.
        y: les vecteurs de probabilités, de même forme que les sorties.
        rkey: clé aléatoire (dropout), ou `None`.
    """
    logits = batch_apply(model, x, rkey)
    return jnp.mean(optax.softmax_cross_entropy(logits, y))

hinge

yax.obj.hinge(model, x, y, rkey=None)

Perte à marge (hinge), pour la classification binaire, à la manière des SVM.

Vaut max(0, 1 - y * score) : elle ne demande pas seulement le bon signe, mais une marge d'au moins 1. Un exemple déjà bien classé, au-delà de la marge, ne coûte plus rien — au contraire de l'entropie croisée, qui pousse toujours un peu.

Paramètres :

Nom Type Description Défaut
model Module

le modèle, qui rend un score réel.

obligatoire
x ArrayLike

les entrées du lot.

obligatoire
y ArrayLike

les étiquettes -1 ou +1 (et non 0 ou 1), de même forme que les sorties.

obligatoire
rkey Array | None

clé aléatoire (dropout), ou None.

None
Code source dans yax/obj.py
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
@objective
def hinge(model: Module, x: ArrayLike, y: ArrayLike, rkey: jax.Array | None = None) -> jax.Array:
    """Perte à marge (hinge), pour la classification binaire, à la manière des SVM.

    Vaut `max(0, 1 - y * score)` : elle ne demande pas seulement le bon signe,
    mais une marge d'au moins 1. Un exemple déjà bien classé, au-delà de la
    marge, ne coûte plus rien — au contraire de l'entropie croisée, qui pousse
    toujours un peu.

    Args:
        model: le modèle, qui rend un score réel.
        x: les entrées du lot.
        y: les étiquettes **-1 ou +1** (et non 0 ou 1), de même forme que les sorties.
        rkey: clé aléatoire (dropout), ou `None`.
    """
    scores = batch_apply(model, x, rkey)
    return jnp.mean(optax.hinge_loss(scores, jnp.asarray(y)))

focal_loss

yax.obj.focal_loss(model, x, y, rkey=None, *, gamma=2.0, alpha=None)

Entropie croisée binaire focalisée (Lin et al., 2017), sur les logits.

Chaque exemple est pesé par (1 - p)**gamma, où p est la probabilité qu'il donne à la bonne classe : les exemples déjà bien classés ne pèsent presque plus, et l'apprentissage se concentre sur les cas difficiles. Utile quand une classe écrase l'autre — en détection, le fond occupe presque toute l'image. Se configure par yax.obj.focal_loss(gamma=1.0, alpha=0.25).

Paramètres :

Nom Type Description Défaut
model Module

le modèle, dont la sortie ne passe PAS par une sigmoïde.

obligatoire
x ArrayLike

les entrées du lot.

obligatoire
y ArrayLike

les étiquettes 0 ou 1, de même forme que les sorties.

obligatoire
rkey Array | None

clé aléatoire (dropout), ou None.

None
gamma float

force de la focalisation ; 0.0 redonne binary_cross_entropy.

2.0
alpha float | None

poids de la classe 1, entre 0 et 1 (souvent 0.25) ; None pour ne pas pondérer les classes.

None
Code source dans yax/obj.py
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
@objective
def focal_loss(model: Module, x: ArrayLike, y: ArrayLike, rkey: jax.Array | None = None, *,
               gamma: float = 2.0, alpha: float | None = None) -> jax.Array:
    """Entropie croisée binaire focalisée (Lin et al., 2017), sur les **logits**.

    Chaque exemple est pesé par `(1 - p)**gamma`, où `p` est la probabilité
    qu'il donne à la bonne classe : les exemples déjà bien classés ne pèsent
    presque plus, et l'apprentissage se concentre sur les cas difficiles.
    Utile quand une classe écrase l'autre — en détection, le fond occupe
    presque toute l'image. Se configure par `yax.obj.focal_loss(gamma=1.0, alpha=0.25)`.

    Args:
        model: le modèle, dont la sortie ne passe PAS par une sigmoïde.
        x: les entrées du lot.
        y: les étiquettes 0 ou 1, de même forme que les sorties.
        rkey: clé aléatoire (dropout), ou `None`.
        gamma: force de la focalisation ; `0.0` redonne `binary_cross_entropy`.
        alpha: poids de la classe 1, entre 0 et 1 (souvent `0.25`) ; `None`
            pour ne pas pondérer les classes.
    """
    logits = batch_apply(model, x, rkey)
    return jnp.mean(optax.sigmoid_focal_loss(logits, jnp.asarray(y), alpha=alpha, gamma=gamma))

Segmentation

dice_loss

yax.obj.dice_loss(model, x, y, rkey=None, *, epsilon=1e-06)

Perte de Dice (1 - Dice), pour la segmentation binaire.

Mesure le recouvrement entre le masque prédit et le masque vrai. Au contraire de l'entropie croisée par pixel, elle n'est pas dominée par le fond quand l'objet est petit. On la combine souvent avec binary_cross_entropy, par weighted_sum.

Paramètres :

Nom Type Description Défaut
model Module

le modèle, dont la sortie ne passe PAS par une sigmoïde.

obligatoire
x ArrayLike

les entrées du lot.

obligatoire
y ArrayLike

les masques vrais, de 0 et de 1, de forme (batch, 1, H, W).

obligatoire
rkey Array | None

clé aléatoire (dropout), ou None.

None
epsilon float

petite constante qui évite une division par zéro.

1e-06
Code source dans yax/obj.py
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
@objective
def dice_loss(model: Module, x: ArrayLike, y: ArrayLike, rkey: jax.Array | None = None, *,
              epsilon: float = 1e-6) -> jax.Array:
    """Perte de Dice (`1 - Dice`), pour la segmentation binaire.

    Mesure le recouvrement entre le masque prédit et le masque vrai. Au
    contraire de l'entropie croisée par pixel, elle n'est pas dominée par le
    fond quand l'objet est petit. On la combine souvent avec
    `binary_cross_entropy`, par `weighted_sum`.

    Args:
        model: le modèle, dont la sortie ne passe PAS par une sigmoïde.
        x: les entrées du lot.
        y: les masques vrais, de 0 et de 1, de forme `(batch, 1, H, W)`.
        rkey: clé aléatoire (dropout), ou `None`.
        epsilon: petite constante qui évite une division par zéro.
    """
    logits = batch_apply(model, x, rkey)
    p = jax.nn.sigmoid(logits)
    axes = tuple(range(1, p.ndim))  # tout sauf le batch : un Dice par échantillon
    intersection = jnp.sum(p * y, axis=axes)
    denom = jnp.sum(p, axis=axes) + jnp.sum(y, axis=axes)
    dice = (2.0 * intersection + epsilon) / (denom + epsilon)
    return 1.0 - jnp.mean(dice)

intersection_over_union_loss

yax.obj.intersection_over_union_loss(model, x, y, rkey=None, *, epsilon=1e-06)

Perte de Jaccard (1 - IoU), pour la segmentation binaire. Alias : yax.obj.iou_loss.

L'intersection sur l'union est la mesure de recouvrement dont on juge d'ordinaire une segmentation. Elle est proche du Dice, mais punit plus durement les masques qui débordent.

Paramètres :

Nom Type Description Défaut
model Module

le modèle, dont la sortie ne passe PAS par une sigmoïde.

obligatoire
x ArrayLike

les entrées du lot.

obligatoire
y ArrayLike

les masques vrais, de 0 et de 1, de forme (batch, 1, H, W).

obligatoire
rkey Array | None

clé aléatoire (dropout), ou None.

None
epsilon float

petite constante qui évite une division par zéro.

1e-06
Code source dans yax/obj.py
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
@objective
def intersection_over_union_loss(model: Module, x: ArrayLike, y: ArrayLike,
                                 rkey: jax.Array | None = None, *,
                                 epsilon: float = 1e-6) -> jax.Array:
    """Perte de Jaccard (`1 - IoU`), pour la segmentation binaire. Alias : `yax.obj.iou_loss`.

    L'intersection sur l'union est la mesure de recouvrement dont on juge
    d'ordinaire une segmentation. Elle est proche du Dice, mais punit plus
    durement les masques qui débordent.

    Args:
        model: le modèle, dont la sortie ne passe PAS par une sigmoïde.
        x: les entrées du lot.
        y: les masques vrais, de 0 et de 1, de forme `(batch, 1, H, W)`.
        rkey: clé aléatoire (dropout), ou `None`.
        epsilon: petite constante qui évite une division par zéro.
    """
    logits = batch_apply(model, x, rkey)
    p = jax.nn.sigmoid(logits)
    axes = tuple(range(1, p.ndim))  # tout sauf le batch : un IoU par échantillon
    intersection = jnp.sum(p * y, axis=axes)
    union = jnp.sum(p, axis=axes) + jnp.sum(y, axis=axes) - intersection
    return 1.0 - jnp.mean((intersection + epsilon) / (union + epsilon))

Combiner, écrire les siens

weighted_sum

yax.obj.weighted_sum(objectives)

Combine plusieurs objectifs en un seul, par somme pondérée.

En segmentation, on additionne couramment l'entropie croisée par pixel et la perte de Dice : la première soigne chaque pixel, la seconde le recouvrement global.

objective_fn = yax.obj.weighted_sum({yax.obj.bce: 1.0, yax.obj.dice_loss: 0.5})
run = yax.training.train("out/segmentation", config, objective_fn, model,
                         training, validation)

Paramètres :

Nom Type Description Défaut
objectives dict[Callable, float]

les objectifs à combiner, et le poids de chacun.

obligatoire

Renvoie :

Type Description
Objective

Un objectif de signature objective_fn(model, x, y, rkey=None).

Code source dans yax/obj.py
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
def weighted_sum(objectives: dict[Callable, float]) -> Objective:
    """Combine plusieurs objectifs en un seul, par somme pondérée.

    En segmentation, on additionne couramment l'entropie croisée par pixel et
    la perte de Dice : la première soigne chaque pixel, la seconde le
    recouvrement global.

    ```python
    objective_fn = yax.obj.weighted_sum({yax.obj.bce: 1.0, yax.obj.dice_loss: 0.5})
    run = yax.training.train("out/segmentation", config, objective_fn, model,
                             training, validation)
    ```

    Args:
        objectives: les objectifs à combiner, et le poids de chacun.

    Returns:
        Un objectif de signature `objective_fn(model, x, y, rkey=None)`.
    """
    def somme_ponderee(model: Module, x: ArrayLike, y: ArrayLike,
                       rkey: jax.Array | None = None) -> jax.Array:
        total = jnp.zeros(())
        for objective_fn, poids in objectives.items():
            total = total + poids * objective_fn(model, x, y, rkey)
        return total

    return Objective(somme_ponderee)

objective

yax.obj.objective(fn)

Décorateur : fait d'une fonction (model, x, y, rkey=None, *, options) un objectif configurable.

@yax.obj.objective
def perte_l2(model, x, y, rkey=None, *, lam=1e-3):
    penalite = sum(jnp.sum(p ** 2) for p in jax.tree.leaves(model))
    return yax.obj.mse(model, x, y, rkey) + lam * penalite

run = yax.training.train("out/l2", config, perte_l2(lam=1e-2), model, training, validation)

Paramètres :

Nom Type Description Défaut
fn Callable

la fonction ; ses options sont ses arguments nommés après l'étoile.

obligatoire

Renvoie :

Type Description
Objective

Un yax.obj.Objective.

Code source dans yax/obj.py
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
def objective(fn: Callable) -> Objective:
    """Décorateur : fait d'une fonction `(model, x, y, rkey=None, *, options)` un objectif configurable.

    ```python
    @yax.obj.objective
    def perte_l2(model, x, y, rkey=None, *, lam=1e-3):
        penalite = sum(jnp.sum(p ** 2) for p in jax.tree.leaves(model))
        return yax.obj.mse(model, x, y, rkey) + lam * penalite

    run = yax.training.train("out/l2", config, perte_l2(lam=1e-2), model, training, validation)
    ```

    Args:
        fn: la fonction ; ses options sont ses arguments nommés après l'étoile.

    Returns:
        Un `yax.obj.Objective`.
    """
    return Objective(fn)