💻 Modern ArchitecturesLEVEL 2 · INTERMEDIATEChapter 6
Modern Optimizers & Normalization Mechanics
AdamW decoupled weight decay, exponential moving averages, and RMSNorm from scratch.
💡 1. Core Intuition & Concepts
Modern deep learning training relies on adaptive gradient algorithms like AdamW and scale-invariant normalization layers like RMSNorm to accelerate convergence and stabilize optimization.
📐 2. Mathematical Formulations & Derivations
AdamW Update Equations
m_t: 1st moment (momentum), v_t: 2nd raw moment, λ: decoupled weight decay coefficient.
⚙️ 3. Step-by-Step Computational Mechanism
1
Compute Gradients
Obtain parameter gradient vector g_t = ∇_θ L.
2
Update Moving Averages
Accumulate 1st moment m_t (direction) and 2nd moment v_t (uncentered variance).
3
Apply Bias Correction
Correct zero-initialization bias: m_hat = m_t / (1 - beta1^t), v_hat = v_t / (1 - beta2^t).
4
Decoupled Weight Decay Step
Subtract scaled adaptive gradient plus direct weight decay -alpha * lambda * theta.
💻 4. Code from Scratch (python)
optim.py
import numpy as np
class AdamWScratch:
"""AdamW optimizer with decoupled weight decay implemented in pure Python."""
def __init__(self, params, lr=1e-3, beta1=0.9, beta2=0.999, eps=1e-8, weight_decay=0.01):
self.params = params
self.lr = lr
self.beta1 = beta1
self.beta2 = beta2
self.eps = eps
self.weight_decay = weight_decay
self.m = {p: np.zeros_like(v) for p, v in params.items()}
self.v = {p: np.zeros_like(v) for p, v in params.items()}
self.t = 0
def step(self, grads):
self.t += 1
for p in self.params:
g = grads[p]
self.m[p] = self.beta1 * self.m[p] + (1 - self.beta1) * g
self.v[p] = self.beta2 * self.v[p] + (1 - self.beta2) * (g ** 2)
m_hat = self.m[p] / (1 - self.beta1 ** self.t)
v_hat = self.v[p] / (1 - self.beta2 ** self.t)
self.params[p] -= self.lr * (m_hat / (np.sqrt(v_hat) + self.eps) + self.weight_decay * self.params[p])🧠 5. Comprehension Checkpoint
Answer all 1 questions correctly to complete the chapter · 0 / 1 done
Q1/1 Why does AdamW decouple weight decay from the gradient update instead of adding L2 regularization to the loss?
In the catalog
Finished this chapter?