learn.chetana.fr

PyTorch en pratique

13 min de lectureEssentiel

PyTorch = numpy + gradients + GPU

Tu viens d'écrire la backprop à la main ; PyTorch industrialise exactement ça (PyTorch n'est pas dans le navigateur — cette leçon se lit, le code s'annote) :

import torch

x = torch.randn(32, 2, device="cuda")   # un ndarray… sur GPU
W = torch.randn(2, 8, requires_grad=True) # "suis les gradients de ce tensor"
y = (x @ W).relu().sum()
y.backward()                              # ← TOUTE la backprop de la leçon 6.2
print(W.grad)                             # les gradients, calculés pour toi

requires_grad + backward() : c'est l'autograd — PyTorch enregistre le graphe des opérations au forward et le rejoue à l'envers. Tes 15 lignes de chain rule manuelle tiennent maintenant en une.

Les 4 abstractions à connaître

# 1. nn.Module : le composant réseau (ton "service" avec état)
class MonReseau(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.stack = torch.nn.Sequential(
            torch.nn.Linear(2, 8), torch.nn.ReLU(), torch.nn.Linear(8, 1),
        )
    def forward(self, x):
        return self.stack(x)

# 2. DataLoader : batching + shuffle + chargement parallèle (num_workers, module 1 !)
loader = torch.utils.data.DataLoader(dataset, batch_size=64, shuffle=True, num_workers=4)

# 3. Optimizer : l'update intelligent (Adam = SGD + momentum + lr adaptatif par paramètre)
opt = torch.optim.AdamW(model.parameters(), lr=3e-4)

# 4. La boucle — reconnais les 5 étapes de la leçon 6.2 :
for X_b, y_b in loader:
    pred = model(X_b)                        # forward
    loss = torch.nn.functional.mse_loss(pred, y_b)
    loss.backward()                          # backward (autograd)
    opt.step()                               # update
    opt.zero_grad()                          # zéro

Les réflexes qui évitent 80 % des bugs PyTorch

  • model.train() vs model.eval() : certains composants (dropout, batchnorm) changent de comportement — l'oubli classique qui fausse l'évaluation ;
  • with torch.no_grad(): autour de l'inférence : pas de graphe de gradients = mémoire et vitesse ;
  • device discipline : modèle ET données sur le même device (x.to("cuda")) — sinon l'erreur runtime la plus vue du métier ;
  • shapes, shapes, shapes : print(x.shape) sans honte. (batch, features) attendu partout.
🧩 Quiz1/4

Que déclenche loss.backward() ?

🃏 Flashcards1/5