Aller au contenu

yax.training

L'entraînement : la fonction train, le rechargement des résultats, et les objets qui fournissent les données.

train

yax.training.train(mother_folder, config, objective_fn, model, training, validation, *, rkey=None, optimizer_state=None, title=None, verbose=False)

Entraîne un modèle et enregistre le meilleur dans un nouveau run.

L'ordre des arguments rappelle celui de l'objectif, objective_fn(model, x, y, rkey) : l'objectif, puis le modèle, puis les données, et la clé en dernier.

À chaque époque, le modèle est optimisé sur les lots d'entraînement puis évalué sur le jeu de validation. Le modèle de plus faible perte de validation est conservé et enregistré, avec l'état de l'optimiseur, la configuration et l'historique, dans un sous-dossier numéroté de mother_folder (0, 1, 2…).

Les données d'entraînement et de validation se donnent sous l'une de ces formes :

  • (X, Y) : toutes les données en un seul lot ;
  • (X, Y, batch_size) : mélangées à chaque époque, puis découpées en lots ;
  • un sampler (yax.training.DatasetSampler, yax.training.FunctionSampler).

La validation utilise toujours les mêmes lots : les pertes des runs d'un même mother_folder sont donc comparables. Le modèle est entraîné en mode entraînement, évalué et renvoyé en mode inférence.

run = yax.training.train("out/sinus", config, yax.obj.mse, model,
                         (X_train, Y_train, 32), (X_val, Y_val), title="MLP tanh")

Paramètres :

Nom Type Description Défaut
mother_folder str

dossier de l'expérience ; un dossier par jeu de validation.

obligatoire
config TrainConfig

la configuration de l'optimisation, un yax.configs.TrainConfig.

obligatoire
objective_fn Callable

l'objectif à minimiser, objective_fn(model, x, y, rkey) ; ceux de yax.obj (yax.obj.mse, yax.obj.bce…) ou le vôtre.

obligatoire
model Module

le modèle à entraîner.

obligatoire
training

les données d'entraînement.

obligatoire
validation

les données de validation. Avec un batch_size, celui-ci doit diviser la taille du jeu.

obligatoire
rkey Array | None

clé aléatoire (mélange des données, dropout) ; None vaut jr.key(0), et l'entraînement est alors reproductible.

None
optimizer_state Any

état de l'optimiseur à reprendre (par exemple run.opt_state), ou None pour repartir de zéro.

None
title str | None

un titre pour ce run, enregistré avec lui ; on le retrouve dans run.title pour légender les comparaisons.

None
verbose bool

affiche la progression.

False

Renvoie :

Type Description
Run

Le run créé, un yax.training.Run qui contient le meilleur modèle.

Code source dans yax/training/train.py
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
def train(mother_folder: str, config: TrainConfig, objective_fn: Callable, model: Module,
          training, validation, *, rkey: jax.Array | None = None, optimizer_state: Any = None,
          title: str | None = None, verbose: bool = False) -> Run:
    """Entraîne un modèle et enregistre le meilleur dans un nouveau run.

    L'ordre des arguments rappelle celui de l'objectif,
    `objective_fn(model, x, y, rkey)` : l'objectif, puis le modèle, puis les
    données, et la clé en dernier.

    À chaque époque, le modèle est optimisé sur les lots d'entraînement puis
    évalué sur le jeu de validation. Le modèle de plus faible perte de
    validation est conservé et enregistré, avec l'état de l'optimiseur, la
    configuration et l'historique, dans un sous-dossier numéroté de
    `mother_folder` (`0`, `1`, `2`…).

    Les données d'entraînement et de validation se donnent sous l'une de ces
    formes :

    - `(X, Y)` : toutes les données en un seul lot ;
    - `(X, Y, batch_size)` : mélangées à chaque époque, puis découpées en lots ;
    - un sampler (`yax.training.DatasetSampler`, `yax.training.FunctionSampler`).

    La validation utilise toujours les mêmes lots : les pertes des runs d'un
    même `mother_folder` sont donc comparables. Le modèle est entraîné en mode
    entraînement, évalué et renvoyé en mode inférence.

    ```python
    run = yax.training.train("out/sinus", config, yax.obj.mse, model,
                             (X_train, Y_train, 32), (X_val, Y_val), title="MLP tanh")
    ```

    Args:
        mother_folder: dossier de l'expérience ; un dossier par jeu de validation.
        config: la configuration de l'optimisation, un `yax.configs.TrainConfig`.
        objective_fn: l'objectif à minimiser, `objective_fn(model, x, y, rkey)` ;
            ceux de `yax.obj` (`yax.obj.mse`, `yax.obj.bce`…) ou le vôtre.
        model: le modèle à entraîner.
        training: les données d'entraînement.
        validation: les données de validation. Avec un `batch_size`, celui-ci
            doit diviser la taille du jeu.
        rkey: clé aléatoire (mélange des données, dropout) ; `None` vaut
            `jr.key(0)`, et l'entraînement est alors reproductible.
        optimizer_state: état de l'optimiseur à reprendre (par exemple
            `run.opt_state`), ou `None` pour repartir de zéro.
        title: un titre pour ce run, enregistré avec lui ; on le retrouve
            dans `run.title` pour légender les comparaisons.
        verbose: affiche la progression.

    Returns:
        Le run créé, un `yax.training.Run` qui contient le meilleur modèle.
    """
    sampler = _as_sampler(training, "training")
    if rkey is None:
        rkey = jr.key(0)
    if config.patience is not None and config.patience < 1:
        raise ValueError(
            f"patience : attendu un nombre d'epoques >= 1, ou None pour aller au bout "
            f"des nb_epochs — recu {config.patience}.")
    try:
        # la config est enregistree dans le dossier du run : c'est elle qui dira
        # comment le modele a ete entraine. On le verifie AVANT d'entrainer.
        pickle.dumps(config)
    except (pickle.PicklingError, AttributeError, TypeError) as e:
        raise TypeError(
            f"la config ne se serialise pas ({e}) — cas courant : un optimizer donne "
            f"en lambda. config.optimizer attend un NOM ({', '.join(OPTIMIZERS)}) et "
            f"ses reglages vont dans config.optimizer_options, par exemple "
            f"optimizer='adamw', optimizer_options={{'weight_decay': 0.1}}.") from e
    model = model.set_inference(False)   # mode entrainement pour toute la boucle

    nb_batches = sampler.nb_batches
    schedule = optax.cosine_decay_schedule(
        config.learning_rate, config.nb_epochs * nb_batches, alpha=config.lr_final_ratio)
    # le schedule ne peut etre fabrique qu'ici : nb_batches n'est connu qu'a ce moment.
    # with_extra_args_support : les optimiseurs a recherche lineaire (lbfgs) reclament
    # la valeur de la perte et la fonction qui la calcule ; les autres les ignorent.
    # Un seul chemin de code, et rien a changer cote utilisateur.
    optimizer = optax.with_extra_args_support(
        resolve_optimizer(config.optimizer)(schedule, **(config.optimizer_options or {})))
    if optimizer_state is None:
        # optimizer.init accepte le modele tel quel : tous ses leaves sont des
        # tableaux, les champs non-tableaux etant declares StaticField.
        optimizer_state = optimizer.init(model)

    folder = new_run_folder(mother_folder)
    save_as_pickle(f"{folder}/config", config)
    save_as_pickle(f"{folder}/title", title)

    @jax.jit
    def optim_step(model, x, y, optimizer_state, step_rkey):
        # la perte vue comme fonction du SEUL modele : c'est ce que reclame une
        # recherche lineaire, qui la reevalue en plusieurs points du segment
        def value_fn(model):
            return objective_fn(model, x, y, step_rkey)

        loss, grads = jax.value_and_grad(value_fn)(model)
        # le modele est passe en 3e argument : les optimiseurs a weight decay
        # (adamw, lion...) en ont besoin, les autres l'ignorent — de meme que les
        # arguments suivants, utiles a la seule recherche lineaire
        updates, optimizer_state = optimizer.update(
            grads, optimizer_state, model, value=loss, grad=grads, value_fn=value_fn)
        # apply_updates RETOURNE un nouveau modele, rien n'est modifie en place
        model = optax.apply_updates(model, updates)
        return model, optimizer_state, loss

    # la validation se fait en mode inference (dropout coupe), sans cle
    eval_loss = jax.jit(objective_fn)
    validation_batches = _validation_batches(validation)

    def validate(model_eval):
        losses = [eval_loss(model_eval, x, y, None) for x, y in validation_batches()]
        return float(jnp.mean(jnp.stack(losses)))

    best_model = model.set_inference(True)
    best_opt_state = optimizer_state
    best_loss = float("inf")
    history = History()
    step = 0
    epochs_without_record = 0

    for _ in range(config.nb_epochs):
        rkey, rkey_sampler, rkey_epoch = jr.split(rkey, 3)
        step_rkeys = jr.split(rkey_epoch, nb_batches)  # une clé de dropout par step
        batch_losses = []
        for (x, y), step_rkey in zip(sampler(rkey_sampler), step_rkeys):
            model, optimizer_state, loss = optim_step(model, x, y, optimizer_state, step_rkey)
            batch_losses.append(loss)
            step += 1
        if len(batch_losses) != nb_batches:
            raise ValueError(
                f"le sampler a rendu {len(batch_losses)} batchs alors qu'il en déclare "
                f"nb_batches={nb_batches} : le schedule du learning rate compte sur ce nombre.")

        # un seul transfert device->host par epoque, et non un par step
        history.add_train_losses(jnp.stack(batch_losses).tolist())
        model_eval = model.set_inference(True)
        val_loss = validate(model_eval)
        history.add_val_loss(step, val_loss)

        if val_loss < best_loss:  # strict : un plateau n'est pas un progres
            best_loss = val_loss
            best_model = model_eval
            best_opt_state = optimizer_state
            epochs_without_record = 0
            save_as_pickle(f"{folder}/trained_model", best_model)
            save_as_pickle(f"{folder}/opt_state", best_opt_state)
            save_as_pickle(f"{folder}/loss", best_loss)
            if verbose:
                print(f"⬊{val_loss:.3g}", end="")
        else:
            epochs_without_record += 1
            if verbose:
                print(".", end="")

        save_as_pickle(f"{folder}/history", history)

        # arret anticipe : le modele rendu est de toute facon le meilleur, on cesse
        # seulement de depenser des epoques qui ne l'ameliorent plus
        if config.patience is not None and epochs_without_record >= config.patience:
            if verbose:
                print(f"| {config.patience} epoques sans record", end="")
            break

    if verbose:
        print(f"| end of the optimization loop. best={best_loss:.3g} in {folder}")
    # l'opt_state rendu est celui DU CHECKPOINT : il correspond au trained_model
    # (celui de fin de boucle appartient au dernier modele, pas au meilleur)
    return Run(folder=folder, title=title, trained_model=best_model, opt_state=best_opt_state,
               loss=best_loss, config=config, history=history)

load_run

yax.training.load_run(folder, *names)

Relit un run enregistré par train.

run = yax.training.load_run("out/anneau/0")
model = run.trained_model

La classe du modèle doit être définie avant le chargement.

Paramètres :

Nom Type Description Défaut
folder str

le dossier du run.

obligatoire
*names str

les champs à charger ("trained_model", "config"…), tous par défaut. Les champs non demandés valent None.

()

Renvoie :

Type Description
Run

Un yax.training.Run.

Code source dans yax/training/train.py
54
55
56
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
83
def load_run(folder: str, *names: str) -> Run:
    """Relit un run enregistré par `train`.

    ```python
    run = yax.training.load_run("out/anneau/0")
    model = run.trained_model
    ```

    La classe du modèle doit être définie avant le chargement.

    Args:
        folder: le dossier du run.
        *names: les champs à charger (`"trained_model"`, `"config"`…), tous
            par défaut. Les champs non demandés valent `None`.

    Returns:
        Un `yax.training.Run`.
    """
    assert os.path.isdir(folder), f"folder:{folder} does not exist"
    if not names:
        names = RUN_ITEMS
    inconnus = [name for name in names if name not in RUN_ITEMS]
    if inconnus:
        raise ValueError(f"load_run : champ(s) inconnu(s) {inconnus} — les "
                         f"champs d'un Run sont : {', '.join(RUN_ITEMS)}.")
    items = {}
    for name in names:
        path = f"{folder}/{name}"
        items[name] = None if name == "title" and not os.path.exists(path) else load_from_pickle(path)
    return Run(folder=folder, **items)

load_runs

yax.training.load_runs(mother_folder, *names)

Relit tous les runs d'une expérience, dans l'ordre de leur création.

for run in yax.training.load_runs("out/sinus"):
    print(run.title, run.loss)

Les runs interrompus avant leur premier enregistrement sont ignorés.

Paramètres :

Nom Type Description Défaut
mother_folder str

le dossier de l'expérience.

obligatoire
*names str

les champs à charger, comme pour load_run ; tous par défaut.

()

Renvoie :

Type Description
list[Run]

La liste des yax.training.Run, du premier créé au dernier.

Code source dans yax/training/train.py
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
def load_runs(mother_folder: str, *names: str) -> list[Run]:
    """Relit tous les runs d'une expérience, dans l'ordre de leur création.

    ```python
    for run in yax.training.load_runs("out/sinus"):
        print(run.title, run.loss)
    ```

    Les runs interrompus avant leur premier enregistrement sont ignorés.

    Args:
        mother_folder: le dossier de l'expérience.
        *names: les champs à charger, comme pour `load_run` ; tous par défaut.

    Returns:
        La liste des `yax.training.Run`, du premier créé au dernier.
    """
    assert os.path.isdir(mother_folder), f"mother_folder:{mother_folder} does not exist"
    # tri numerique : "10" vient apres "9" ; .DS_Store et autres sont ignores
    numbers = sorted(int(name) for name in os.listdir(mother_folder) if name.isdigit())
    folders = [os.path.join(mother_folder, str(number)) for number in numbers]
    return [load_run(folder, *names) for folder in folders if os.path.exists(f"{folder}/loss")]

find_best_run

yax.training.find_best_run(mother_folder)

Renvoie le dossier du run de plus faible perte de validation.

Paramètres :

Nom Type Description Défaut
mother_folder str

le dossier de l'expérience.

obligatoire

Renvoie :

Type Description
str

Le chemin du meilleur run.

Code source dans yax/training/train.py
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
def find_best_run(mother_folder: str) -> str:
    """Renvoie le dossier du run de plus faible perte de validation.

    Args:
        mother_folder: le dossier de l'expérience.

    Returns:
        Le chemin du meilleur run.
    """
    assert os.path.isdir(mother_folder), f"mother_folder:{mother_folder} does not exist"
    best_loss = float("inf")
    best_folder = None
    for name in sorted(os.listdir(mother_folder)):
        # noinspection bad-argument-type
        folder = os.path.join(mother_folder, name)
        if not os.path.isdir(folder) or not os.path.exists(f"{folder}/loss"):
            continue  # .DS_Store, fichier de synchro, run interrompu...
        loss = load_from_pickle(f"{folder}/loss")
        if loss < best_loss:
            best_loss = loss
            best_folder = folder
    assert best_folder is not None, f"aucun run complet dans {mother_folder}"
    return best_folder

Run dataclass

yax.training.Run(folder=None, title=None, trained_model=None, opt_state=None, loss=None, config=None, history=None)

Résultat d'un entraînement, renvoyé par train et par load_run.

Attributs :

Nom Type Description
folder str | None

le dossier où le run est enregistré.

title str | None

le titre donné à train, pour légender les comparaisons ; None s'il n'y en a pas.

trained_model Any

le meilleur modèle, en mode inférence.

opt_state Any

l'état de l'optimiseur correspondant.

loss float | None

la perte de validation de ce modèle.

config TrainConfig | None

la configuration de l'entraînement.

history History | None

l'historique des pertes, un yax.training.History.

Code source dans yax/training/Run.py
8
9
@dataclass
class Run:
22
23
24
25
26
27
28
    folder: str | None = None          # le sous-dossier où ce run est sauvegardé
    title: str | None = None
    trained_model: Any = None   # le meilleur modèle (inference=True, prêt à évaluer)
    opt_state: Any = None       # l'état de l'optimiseur correspondant
    loss: float | None = None          # sa loss de validation
    config: TrainConfig | None = None
    history: History | None = None

History dataclass

yax.training.History(train_losses=list(), val_losses=list(), val_steps=list())

Historique des pertes d'un entraînement.

Un pas (step) est une mise à jour des paramètres sur un lot. La perte d'entraînement est notée à chaque pas, la perte de validation à la fin de chaque époque.

Attributs :

Nom Type Description
train_losses list[float]

la perte d'entraînement de chaque pas.

val_losses list[float]

les pertes de validation.

val_steps list[int]

le numéro du pas de chaque validation.

Code source dans yax/training/History.py
6
7
@dataclass
class History:
19
20
21
    train_losses: list[float] = field(default_factory=list)  # une valeur par step
    val_losses: list[float] = field(default_factory=list)
    val_steps: list[int] = field(default_factory=list)     # abscisses des val_losses

plot

plot(*, ax=None, log=True, title=None)

Trace les courbes de perte.

La perte d'entraînement apparaît par pas (trait clair) et en moyenne par époque (trait foncé) ; la perte de validation par des points, en vert quand elle bat son record. Le gros point vert marque le modèle conservé.

Paramètres :

Nom Type Description Défaut
ax

axe matplotlib où tracer ; une nouvelle figure par défaut.

None
log bool

échelle logarithmique en ordonnée.

True
title str | None

titre du graphique.

None

Renvoie :

Type Description
Any

L'axe matplotlib.

Code source dans yax/training/History.py
34
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
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
def plot(self, *, ax=None, log: bool = True, title: str | None = None) -> Any:
    """Trace les courbes de perte.

    La perte d'entraînement apparaît par pas (trait clair) et en moyenne par
    époque (trait foncé) ; la perte de validation par des points, en vert
    quand elle bat son record. Le gros point vert marque le modèle conservé.

    Args:
        ax: axe matplotlib où tracer ; une nouvelle figure par défaut.
        log: échelle logarithmique en ordonnée.
        title: titre du graphique.

    Returns:
        L'axe matplotlib.
    """
    # import local : garder `import yax` leger (matplotlib est long a charger)
    import matplotlib.pyplot as plt

    if ax is None:
        # constrained : la figure fait de la place a la legende, posee a
        # droite de l'axe (voir plus bas)
        _, ax = plt.subplots(figsize=(8, 4), layout="constrained")

    couleur_train = "lightcoral"
    ax.plot(range(self.nb_steps), self.train_losses,
            "-", lw=1, color=couleur_train, alpha=0.3, label="train (steps)")
    # moyenne par epoque : les epoques sont delimitees par val_steps
    if self.val_steps:
        bords = [0] + list(self.val_steps)
        centres = [(a + b) / 2.0 for a, b in zip(bords[:-1], bords[1:])]
        moyennes = [sum(self.train_losses[a:b]) / max(b - a, 1)
                    for a, b in zip(bords[:-1], bords[1:])]
        ax.plot(centres, moyennes, "-", lw=1.5, color=couleur_train,
                label="train (moyenne/époque)")

    ax.plot(self.val_steps, self.val_losses,
            ".", ms=5, color="tab:blue", label="validation")
    # les records : chaque amelioration du minimum courant
    records_x, records_y, minimum = [], [], float("inf")
    for step, loss in zip(self.val_steps, self.val_losses):
        if loss < minimum:
            minimum = loss
            records_x.append(step)
            records_y.append(loss)
    if records_x:
        ax.plot(records_x, records_y, ".", ms=5, color="tab:green",
                label="records")
        ax.plot(records_x[-1], records_y[-1], "o", ms=9, color="tab:green",
                label="meilleur modèle")

    if log:
        ax.set_yscale("log")
    ax.set_xlabel("step (optimisations)")
    ax.set_ylabel("loss")
    if title is not None:
        ax.set_title(title)
    # legende HORS de l'axe, a droite : elle ne cache jamais les courbes, et
    # un emplacement fixe evite la recherche de loc="best", qui compte les
    # points recouverts par chaque position candidate — lente avec des
    # dizaines de milliers de steps
    ax.legend(loc="upper left", bbox_to_anchor=(1.02, 1.0), borderaxespad=0.0,
              frameon=False)
    ax.grid(alpha=0.3)
    return ax

DatasetSampler

yax.training.DatasetSampler(X, Y, batch_size)

Lots tirés d'un jeu de données fini.

À chaque époque, les données sont mélangées puis découpées en lots de batch_size ; les exemples qui ne remplissent pas un dernier lot sont laissés de côté pour cette époque. Passer (X, Y, batch_size) à train revient à utiliser ce sampler.

Paramètres :

Nom Type Description Défaut
X ndarray | Array

les entrées.

obligatoire
Y ndarray | Array

les cibles ; pour un apprentissage non supervisé, passer Y = X.

obligatoire
batch_size int

la taille d'un lot.

obligatoire

Attributs :

Nom Type Description
nb_batches

nombre de lots par époque, len(X) // batch_size.

Code source dans yax/training/samplers.py
31
class DatasetSampler:
48
49
50
51
52
53
54
    def __init__(self, X: np.ndarray | jax.Array, Y: np.ndarray | jax.Array, batch_size: int):
        assert len(X) == len(Y), f"X ({len(X)}) et Y ({len(Y)}) n'ont pas la même longueur"
        assert len(X) // batch_size > 0, f"batch_size:{batch_size} > nb_data:{len(X)}"
        self.X = X
        self.Y = Y
        self.batch_size = batch_size
        self.nb_batches = len(X) // batch_size

FunctionSampler

yax.training.FunctionSampler(f, nb_batches)

Lots calculés par une fonction, sans jeu de données.

Pour les problèmes où les points se tirent au hasard (équations aux dérivées partielles, méthode de Ritz…).

sampler = yax.training.FunctionSampler(lambda rkey: (jr.uniform(rkey, (256,)), None),
                                       nb_batches=20)

Paramètres :

Nom Type Description Défaut
f Callable

fonction f(rkey) -> (x, y) qui fabrique un lot ; y peut valoir None.

obligatoire
nb_batches int

nombre de lots par époque.

obligatoire

Attributs :

Nom Type Description
nb_batches

nombre de lots par époque.

Code source dans yax/training/samplers.py
64
class FunctionSampler:
83
84
85
86
    def __init__(self, f: Callable, nb_batches: int):
        assert nb_batches > 0, nb_batches
        self.f = f
        self.nb_batches = nb_batches