DEEP LEARNING / 8. TRAINING TRICKS

Training Tricks

The difference between a model that trains and one that doesn't


EXPLANATION

Architecture matters, but training tricks often matter more. These are the techniques that separate a model that converges cleanly from one that explodes or collapses.

Essential tricks:
• Batch Normalization → normalizes activations per mini-batch, stabilizes training, allows higher lr
• Layer Normalization → normalizes per sample (used in transformers, independent of batch size)
• Dropout            → randomly zeroes activations during training, prevents co-adaptation (overfitting)
• Gradient Clipping  → caps gradient norm to prevent exploding gradients (critical for RNNs/transformers)
• Weight Init        → Kaiming (ReLU), Xavier (tanh/sigmoid). Bad init → dead neurons or explosions
• Mixed Precision    → train in float16, keep master weights in float32. ~2× speedup on modern GPUs
• Early Stopping     → stop when val loss stops improving, save best checkpoint

DATA FLOW

Training instabilities and fixes:

  Exploding gradients  → gradient clipping (clip_grad_norm_)
  Vanishing gradients  → residual connections, LayerNorm, better init
  Overfitting          → dropout, weight decay, data augmentation
  Slow convergence     → learning rate warmup, better optimizer
  Covariate shift      → BatchNorm (CNNs) or LayerNorm (Transformers)

  Training checklist:
  ✓ Normalize inputs (zero mean, unit variance)
  ✓ Use proper weight init (Kaiming for ReLU)
  ✓ Clip gradients (max_norm=1.0)
  ✓ Warmup lr for first N steps
  ✓ Monitor grad norms — if exploding, something is wrong

CODE

PYTHON
1import torch
2import torch.nn as nn
3
4model = nn.TransformerEncoderLayer(d_model=512, nhead=8, batch_first=True)
5optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=0.01)
6
7# ── Gradient Clipping (critical for transformers) ─────────────────
8loss = torch.tensor(1.0, requires_grad=True)
9loss.backward()
10torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
11optimizer.step()
12
13# ── Learning rate warmup (linear then cosine decay) ───────────────
14def get_lr(step, d_model=512, warmup_steps=4000):
15 """Original transformer lr schedule from 'Attention is All You Need'"""
16 step = max(step, 1)
17 return d_model**(-0.5) * min(step**(-0.5), step * warmup_steps**(-1.5))
18
19lrs = [get_lr(step) for step in range(1, 8000)]
20
21# ── Dropout ───────────────────────────────────────────────────────
22class ModelWithDropout(nn.Module):
23 def __init__(self):
24 super().__init__()
25 self.net = nn.Sequential(
26 nn.Linear(512, 1024),
27 nn.GELU(),
28 nn.Dropout(p=0.1), # zero 10% of activations randomly
29 nn.Linear(1024, 512),
30 nn.Dropout(p=0.1),
31 )
32 def forward(self, x): return self.net(x)
33
34# IMPORTANT: model.train() enables dropout, model.eval() disables it
35m = ModelWithDropout()
36m.train() # dropout ON during training
37m.eval() # dropout OFF during inference
38
39# ── Mixed precision training (2× speedup on GPU) ──────────────────
40from torch.cuda.amp import autocast, GradScaler
41
42scaler = GradScaler()
43
44def train_step(model, x, y, optimizer):
45 optimizer.zero_grad()
46
47 with autocast(): # runs forward in float16
48 out = model(x)
49 loss = nn.MSELoss()(out, y)
50
51 scaler.scale(loss).backward()
52 scaler.unscale_(optimizer)
53 torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
54 scaler.step(optimizer)
55 scaler.update()
56 return loss.item()
57
58# ── Weight initialization ─────────────────────────────────────────
59def init_weights(module):
60 if isinstance(module, nn.Linear):
61 nn.init.kaiming_normal_(module.weight, nonlinearity="relu")
62 if module.bias is not None:
63 nn.init.zeros_(module.bias)
64 elif isinstance(module, nn.Embedding):
65 nn.init.normal_(module.weight, mean=0.0, std=0.02)
66
67model.apply(init_weights)
← PREV7. Transformers