PyTorch en pratique
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()vsmodel.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.
Que déclenche loss.backward() ?