Optimisation and Regularisation
Almost every model is trained by iterative optimisation, and almost every deployed model is regularised. Interviewers probe the mechanics: what gradient descent does, why learning rates matter, how momentum and Adam work, why normalisation helps, and what each regulariser actually does.
1. Gradient descent
Given a loss , repeat . The gradient points uphill; stepping against it reduces the loss locally.
The learning rate
- Too small: painfully slow, may stall in flat regions.
- Too large: overshoots, oscillates or diverges.
- On a quadratic with curvature , gradient descent converges only if . In ill-conditioned problems (very different curvature in different directions) a single learning rate cannot suit all directions, which is why feature scaling and adaptive optimisers matter.
import numpy as np
# minimise f(w) = 0.5 * h * w^2, so grad = h * w; the update multiplies w by (1 - eta*h)
def run(eta, h=10.0, steps=100, w0=1.0):
w = w0
for _ in range(steps):
w -= eta * h * w
return w
assert abs(run(0.05)) < 1e-6 # eta*h = 0.5: converges
assert abs(run(0.19)) < 1e-2 # eta*h = 1.9: converges, oscillating (still below 2/h = 0.2)
assert abs(run(0.21)) > 1e3 # eta*h = 2.1: diverges
Batch, stochastic and mini-batch
Full-batch gradients are exact but expensive. A mini-batch gradient is an unbiased noisy estimate of the true gradient. The noise slows exact convergence but adds exploration and, empirically, often helps generalisation. A common recipe is a mini-batch of 32 to 1024 with a decaying learning rate.
Convex versus non-convex
For convex losses (linear and logistic regression with a convex penalty) any local minimum is global. Neural networks are non-convex, but in practice bad local minima are rarer than saddle points and flat regions, and sensible initialisation plus SGD variants find solutions that generalise.
2. Improving on plain SGD
Momentum
Keep a velocity that accumulates past gradients: , . It damps oscillations across narrow valleys and speeds travel along them, like a heavy ball rolling.
RMSProp and Adam
Adaptive methods scale each parameter's step by a running estimate of its recent gradient magnitude. Adam combines momentum (first moment ) with RMS scaling (second moment ) and bias correction:
Typical defaults: , , . The bias correction matters in the first steps, when and start at zero and would be biased low.
import numpy as np
def adam_minimise(grad, theta0, lr=0.1, steps=500, b1=0.9, b2=0.999, eps=1e-8):
theta = np.array(theta0, dtype=float)
m = np.zeros_like(theta); v = np.zeros_like(theta)
for t in range(1, steps + 1):
g = grad(theta)
m = b1 * m + (1 - b1) * g
v = b2 * v + (1 - b2) * g * g
mh, vh = m / (1 - b1 ** t), v / (1 - b2 ** t)
theta -= lr * mh / (np.sqrt(vh) + eps)
return theta
# a badly scaled bowl: curvature 100 in one direction, 1 in the other
grad = lambda th: np.array([100 * th[0], 1 * th[1]])
theta = adam_minimise(grad, [1.0, 1.0], lr=0.05, steps=800)
assert np.abs(theta).max() < 0.05 # both directions converge despite the 100x scale gap
SGD with momentum often generalises slightly better on some vision tasks; AdamW (Adam with decoupled weight decay) is the standard for transformers.
Learning-rate schedules
Warm up (start small and ramp up, essential for transformers), then decay by step, cosine or inverse-square-root. Cyclical and one-cycle schedules can speed training.
3. Initialisation and normalisation
- Initialisation: small random weights break symmetry; their scale should keep activations and gradients from shrinking or exploding with depth. Xavier/Glorot (variance about to ) suits tanh; He (variance ) suits ReLU.
- Batch normalisation: normalise activations per mini-batch, then scale and shift. Stabilises and speeds training, allows larger learning rates, and gives a mild regularising effect. It behaves differently in training and inference (running statistics).
- Layer normalisation: normalise across features within one example, independent of the batch. The standard in transformers and sequence models.
- Gradient clipping: cap the gradient norm to stop explosions, important for recurrent nets and large-model training.
4. Regularisation
Regularisation is anything that reduces generalisation error without (necessarily) reducing training error.
| Technique | How it works |
|---|---|
| L2 / weight decay | Penalise : shrinks weights; equals a Gaussian prior |
| L1 | Penalise : sparse weights; equals a Laplace prior |
| Dropout | Randomly zero activations during training (with scaling), so units cannot co-adapt; like averaging many thinned networks |
| Early stopping | Stop when validation loss stops improving; limits effective capacity |
| Data augmentation | Create label-preserving variants (flips, crops, noise, back-translation) |
| Smaller model / fewer features | Reduces capacity directly |
| Ensembling | Averages out variance |
| Label smoothing | Softens one-hot targets to prevent overconfidence |
| Batch size and noise | Small batches add implicit regularisation |
L2 versus weight decay
For plain SGD, an L2 penalty and weight decay are equivalent. For adaptive optimisers such as Adam they are not, because the penalty gradient gets rescaled by the adaptive denominator. That is the motivation for AdamW, which applies decay directly to the weights.
Why L1 gives sparsity and L2 does not
The L1 penalty has a constant-magnitude gradient pulling every weight toward zero regardless of size, so small weights reach exactly zero. The L2 gradient is proportional to the weight, so it shrinks big weights a lot and small weights barely, and never reaches zero.
import numpy as np
# proximal / closed-form view for a single weight with data-fit minimiser w0
w0 = 0.2
lam = 0.5
l2 = w0 / (1 + lam) # ridge shrinks but never reaches zero
l1 = np.sign(w0) * max(abs(w0) - lam / 2, 0) # soft-thresholding: exactly zero when |w0| <= lam/2
assert 0 < l2 < w0
assert l1 == 0.0
assert np.isclose(np.sign(0.9) * max(abs(0.9) - lam / 2, 0), 0.65) # large weights survive L1, reduced by a constant
Dropout details
During training each unit is dropped with probability and the kept ones are scaled by ("inverted dropout"), so no change is needed at inference. It is rarely used with batch norm in convolutional nets, and it is common in dense layers and attention.
import numpy as np
rng = np.random.default_rng(0)
x = np.ones(1_000_000)
p = 0.3
mask = rng.uniform(size=x.shape) >= p
y = x * mask / (1 - p)
assert abs(y.mean() - 1.0) < 0.01 # inverted dropout keeps the expected activation unchanged
5. Diagnosing training problems
| Symptom | Likely causes | Fixes |
|---|---|---|
| Loss is NaN or explodes | Learning rate too high, bad initialisation, log(0), exploding gradients | Lower the rate, clip gradients, check data, stabilise softmax |
| Loss does not decrease | Learning rate too low or too high, bug, dead ReLUs, wrong labels | Overfit a tiny batch first; check data pipeline and loss |
| Training good, validation bad | Overfitting | Regularise, augment, more data, early stopping |
| Both poor | Underfitting, features inadequate | Bigger model, better features, longer training |
| Validation noisy | Small validation set, high learning rate | More data, lower rate, average over seeds |
The single most useful debugging step: overfit a single small batch. If a network cannot drive the loss to near zero on 10 examples, there is a bug.
6. Common mistakes
- Not scaling inputs before gradient-based training.
- Comparing optimisers with the learning rate left untuned for each.
- Applying weight decay to biases and normalisation parameters by default.
- Evaluating with dropout and batch norm still in training mode.
- Confusing L2 and weight decay under Adam.
- Reading validation loss from a noisy single run as proof of an improvement.
7. Practice questions
- Explain how the learning rate affects convergence. When does gradient descent diverge on a quadratic?
- Describe momentum and why it helps in narrow valleys.
- Walk through Adam's update. Why the bias correction?
- L1 versus L2: effect on weights, priors and geometry.
- How does dropout regularise, and what is "inverted dropout"?
- Why is layer norm preferred to batch norm in transformers?
- Your loss becomes NaN after a few hundred steps. List causes and checks.
- What is the difference between L2 regularisation and weight decay in Adam?