Um loop de treinamento PyTorch profissional

2 min

Objetivo: separar claramente treinamento, validação, salvamento e métricas.

Uma época de treinamento

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)

Multiplicar por x.size(0) reconstrói a soma por exemplo antes da média global. Uma média ingênua das médias dos batches enviesa o último batch se ele for menor.

Pré-visualização — o resto da lição está reservado aos inscritos.

Já se inscreveu com um código?

O seu acesso está ligado à sua conta, não a esta ligação. Inicie sessão com o mesmo e-mail usado na aula: o seu curso está à espera, não precisa de introduzir o código de novo.

Iniciar sessãoAinda não tem conta? Criar uma
Esta lição faz parte do módulo «Fazer a rede aprender»

Os primeiros módulos do curso são de acesso livre. Para o resto há três possibilidades: comprar este curso de uma vez por todas, subscrever, ou introduzir o código entregue na aula.

Subscrever — 19 $ US/mêsVoltar ao plano
É estudante do curso?

O código fica ligado à sua conta: inicie sessão ou crie uma conta e ele será aplicado automaticamente ao regressar.

Ainda não tem conta? Criar uma