Une boucle d'entraînement PyTorch professionnelle

2 min

Objectif : séparer clairement entraînement, validation, sauvegarde et métriques.

Une époque d'entraînement

python
def train_epoch(model, loader, criterion, optimizer, device):
    model.train()
    total_loss = 0.0
    for x, y in loader:
        x, y = x.to(device), y.to(device)
        optimizer.zero_grad(set_to_none=True)
        logits = model(x)
        loss = criterion(logits, y)
        loss.backward()
        optimizer.step()
        total_loss += loss.item() * x.size(0)
    return total_loss / len(loader.dataset)

Multiplier par x.size(0) reconstruit la somme par exemple avant la moyenne globale. Une moyenne naïve des moyennes de batches biaise le dernier batch s'il est plus petit.

Aperçu — la suite de la leçon est réservée aux inscrits.

Déjà inscrit avec un code ?

Votre accès est rattaché à votre compte, pas à ce lien. Connectez-vous avec le même courriel qu'en classe : votre cours vous attend, inutile de ré-entrer le code.

Se connecterPas encore de compte ? En créer un
Cette leçon fait partie du module « Faire apprendre le réseau »

Les premiers modules du cours sont en accès libre. Pour la suite, trois possibilités : acheter ce cours une fois pour toutes, vous abonner pour tout ouvrir, ou entrer le code si vous suivez le cours en classe.

M’abonner — 19 $ US/moisRetour au plan
Vous êtes étudiant du cours ?

Le code se rattache à votre compte : connectez-vous ou créez un compte, il sera validé automatiquement au retour.

Pas encore de compte ? En créer un