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
|
Code source dans yax/obj.py
111 112 113 114 115 116 117 118 119 120 121 122 123 | |
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
|
Code source dans yax/obj.py
126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 | |
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
|
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 | |
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
|
Code source dans yax/obj.py
166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 | |
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 |
obligatoire |
rkey
|
Array | None
|
clé aléatoire (dropout), ou |
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 | |
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
|
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 | |
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
|
Code source dans yax/obj.py
226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 | |
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
|
gamma
|
float
|
force de la focalisation ; |
2.0
|
alpha
|
float | None
|
poids de la classe 1, entre 0 et 1 (souvent |
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 | |
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 |
obligatoire |
rkey
|
Array | None
|
clé aléatoire (dropout), ou |
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 | |
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 |
obligatoire |
rkey
|
Array | None
|
clé aléatoire (dropout), ou |
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 | |
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 |
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 | |
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 |
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 | |