पाठ 9 / 25

Modules and the Training Loop

Forward, loss, backward, step, evaluate.

A loop you will write many times

Models are classes derived from nn.Module that create layers in __init__ and define forward. A DataLoader serves shuffled mini-batches. Each training step: zero_grad, forward pass, loss, backward, optimizer.step(). After each epoch, switch to model.eval() and compute validation metrics inside torch.no_grad(), then switch back with model.train(). Log training loss and validation metrics every epoch so you can see learning, plateaus and overfitting.

A full training loop on handwritten digits, run

I ran this on CPU with Python 3, PyTorch 2.14.1, numpy 2.5.3 and scikit-learn 1.9.1, with fixed seeds and one thread. A small two-layer network learns the bundled 8x8 digits dataset: test accuracy is 0.509 after 1 epoch, 0.853 after 5 and 0.947 after 20, while the batch loss falls from 2.234 to 0.394.

import torch
from sklearn.datasets import load_digits
from sklearn.model_selection import train_test_split
torch.manual_seed(0); torch.set_num_threads(1)
X, y = load_digits(return_X_y=True)
X = torch.tensor(X / 16.0, dtype=torch.float32); y = torch.tensor(y)
X_tr, X_te, y_tr, y_te = train_test_split(X, y, test_size=0.25, random_state=0, stratify=y)

class MLP(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.net = torch.nn.Sequential(torch.nn.Linear(64, 64), torch.nn.ReLU(), torch.nn.Linear(64, 10))
    def forward(self, x):
        return self.net(x)

model = MLP(); opt = torch.optim.Adam(model.parameters(), lr=1e-3)
loader = torch.utils.data.DataLoader(torch.utils.data.TensorDataset(X_tr, y_tr), batch_size=64, shuffle=True)
for epoch in range(1, 21):
    model.train()
    for xb, yb in loader:
        opt.zero_grad()
        loss = torch.nn.functional.cross_entropy(model(xb), yb)
        loss.backward(); opt.step()
    if epoch in (1, 5, 20):
        model.eval()
        with torch.no_grad():
            acc = (model(X_te).argmax(1) == y_te).float().mean().item()
        print(f"epoch {epoch:>2}  last batch loss {loss.item():.3f}  test accuracy {acc:.3f}")

Output:

epoch  1  last batch loss 2.234  test accuracy 0.509
epoch  5  last batch loss 1.356  test accuracy 0.853
epoch 20  last batch loss 0.394  test accuracy 0.947

Keep the loop boring

Use one well-tested loop structure across projects; most training bugs come from small variations in it.

त्वरित जाँच: What is the correct order inside a training step?

  • backward, zero_grad, step, forward
  • step, backward, forward, loss
  • zero_grad, forward, loss, backward, step
  • forward, step, zero_grad, backward
Answer

zero_grad, forward, loss, backward, step — Clear, compute, differentiate, update.