Aller au contenu

yax.configs

La configuration d'un entraînement, enregistrée avec chaque run.

TrainConfig dataclass

yax.configs.TrainConfig(learning_rate, nb_epochs, lr_final_ratio=1.0, optimizer='adam', optimizer_options=None, patience=None)

Paramètres d'optimisation d'un entraînement.

La configuration est enregistrée avec chaque run : elle suffit à savoir comment un modèle a été entraîné. La taille des lots n'y figure pas : elle se donne avec les données.

config = yax.configs.TrainConfig(learning_rate=5e-3, nb_epochs=300,
                                 optimizer="adamw",
                                 optimizer_options={"weight_decay": 0.1})

Attributs :

Nom Type Description
learning_rate float

pas d'apprentissage initial.

nb_epochs int

nombre d'époques.

lr_final_ratio float

rapport entre le pas final et le pas initial ; le pas décroît selon un cosinus. 1.0 (défaut) garde un pas constant.

optimizer str

nom d'un optimiseur de yax.optimizers.OPTIMIZERS, "adam" par défaut.

optimizer_options dict | None

options de l'optimiseur, par exemple {"weight_decay": 0.1} pour "adamw".

patience int | None

arrête l'entraînement après ce nombre d'époques sans amélioration de la perte de validation ; None pour aller jusqu'au bout.

Code source dans yax/configs.py
4
5
@dataclasses.dataclass
class TrainConfig:
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
    learning_rate: float
    nb_epochs: int
    # alpha de la cosine decay : lr finale = lr_final_ratio * learning_rate.
    # 1.0 redonne exactement le pas constant, (1-a)*cos+a valant alors 1.
    lr_final_ratio: float = 1.0
    # un nom du dictionnaire yax.optimizers.OPTIMIZERS (yax/optimizers.py)
    optimizer: str = "adam"
    # les arguments passes au constructeur, apres le learning rate :
    # {"weight_decay": 0.1} pour adamw, {"momentum": 0.9} pour sgd...
    # None plutot que {} (pas de default_factory) : chaque champ optionnel a alors
    # une valeur de CLASSE, donc une config picklee avant son ajout se relit sans
    # exploser — et "adam" sans option, sans patience, decrit bien ces anciens runs.
    optimizer_options: dict | None = None
    # arret anticipe : nombre d'epoques SANS record de validation au bout duquel on
    # s'arrete. None = aller au bout des nb_epochs. Attention, le schedule du pas est
    # calibre sur nb_epochs : s'arreter avant, c'est s'arreter avant sa descente.
    patience: int | None = None