🎤 Karaokê PyTorch, ou como parar de esquecer optimizer.zero_grad()

Muitos iniciantes em DL quebram não na arquitetura complexa das redes, mas no simples boilerplate do PyTorch. A ordem das chamadas no loop de treinamento é algo que você precisa pesquisar repetidamente até que fique gravado no subconsciente.

Esqueceu de colocar o modelo em train()? Receba pesos incorretos devido ao Dropout/BatchNorm.
Esqueceu o zero_grad()? Parabéns, os gradientes acumulam e o treinamento vai para o lixo.
Colocou step() antes de backward()? Bem, você entendeu.

Encontrei um vídeo que provoca dois sentimentos ao mesmo tempo: muito constrangimento e respeito.
Um cara de cueca e microfone no meio da bagunça simplesmente pegou os 5 passos da Retropropagação e colocou numa melodia que não sai da cabeça.
Parece que funciona melhor do que 10 horas de palestras chatas de especialistas que leem slides.

Para quem está por fora, lembro a única ordem correta de ações que você deve tatuar no subconsciente (ou aprender a música):

1️⃣ model.train() — coloca o modelo em modo de treinamento.
2️⃣ y_pred = model(x) — Forward pass.
3️⃣ loss = loss_fn(y_pred, y) — Calcula o quanto erramos.
4️⃣ optimizer.zero_grad() — Zera os gradientes antes do próximo passo, senão eles acumulam e você vai para o espaço.
5️⃣ loss.backward() — Calcula os gradientes (retropropagação).
6️⃣ optimizer.step() — Dá um passo com o otimizador.

Ainda falta uma versão para TensorFlow, mas aí, receio, teria que escrever uma ópera em três atos só para inicializar as variáveis 🌚

#por_diversão