c
This commit is contained in:
@@ -0,0 +1,529 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "6caa1908",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"# Becoming a backprop ninja\n",
|
||||
"\n",
|
||||
"Compute every gradient **by hand** — **concept -> code -> Your turn** each step.\n",
|
||||
"\n",
|
||||
"This is notebook **04**. So far we have leaned on `loss.backward()` to do the calculus for us.\n",
|
||||
"Here we turn it off and backpropagate **manually** through the whole MLP + BatchNorm +\n",
|
||||
"cross-entropy, checking each result against PyTorch. After this, autograd is never a mystery.\n",
|
||||
"Follows Karpathy's *\"makemore part 4: becoming a backprop ninja\"*."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "9ed68a60",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Prologue — in plain English\n",
|
||||
"\n",
|
||||
"`loss.backward()` answers one question for every weight: *if I nudge you a little, how does\n",
|
||||
"the loss change?* That answer is the gradient. We have trusted PyTorch to compute it.\n",
|
||||
"\n",
|
||||
"Now we earn it. We will:\n",
|
||||
"\n",
|
||||
"1. Run the forward pass but **save every intermediate value**.\n",
|
||||
"2. Walk backward, one operation at a time, applying the chain rule by hand (exactly the\n",
|
||||
" `00_micrograd.ipynb` idea, now on whole tensors).\n",
|
||||
"3. Compare each hand-derived gradient to PyTorch's — they should match.\n",
|
||||
"\n",
|
||||
"Real-life picture: notebook `00` showed one gear; here we backprop through the whole gearbox,\n",
|
||||
"and check our work against the manufacturer's spec (PyTorch)."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "95e96905",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### 0.1 Imports and data"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "3b29d3f0",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import os\n",
|
||||
"os.environ.setdefault(\"KMP_DUPLICATE_LIB_OK\", \"TRUE\")\n",
|
||||
"\n",
|
||||
"import random\n",
|
||||
"import torch\n",
|
||||
"import torch.nn.functional as F\n",
|
||||
"\n",
|
||||
"words = open(\"names.txt\").read().splitlines()\n",
|
||||
"chars = sorted(list(set(\"\".join(words))))\n",
|
||||
"stoi = {ch: i + 1 for i, ch in enumerate(chars)}\n",
|
||||
"stoi[\".\"] = 0\n",
|
||||
"itos = {i: ch for ch, i in stoi.items()}\n",
|
||||
"vocab_size = len(itos)\n",
|
||||
"block_size = 3\n",
|
||||
"\n",
|
||||
"def build_dataset(word_list):\n",
|
||||
" X, Y = [], []\n",
|
||||
" for w in word_list:\n",
|
||||
" context = [0] * block_size\n",
|
||||
" for ch in w + \".\":\n",
|
||||
" ix = stoi[ch]\n",
|
||||
" X.append(context); Y.append(ix)\n",
|
||||
" context = context[1:] + [ix]\n",
|
||||
" return torch.tensor(X), torch.tensor(Y)\n",
|
||||
"\n",
|
||||
"random.seed(42)\n",
|
||||
"random.shuffle(words)\n",
|
||||
"n1, n2 = int(0.8 * len(words)), int(0.9 * len(words))\n",
|
||||
"Xtr, Ytr = build_dataset(words[:n1])\n",
|
||||
"Xdev, Ydev = build_dataset(words[n1:n2])\n",
|
||||
"print(\"train:\", tuple(Xtr.shape), \" vocab:\", vocab_size)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "19612a96",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### 0.2 The gradient-check helper\n",
|
||||
"\n",
|
||||
"`cmp` compares our hand-computed gradient `dt` to PyTorch's `t.grad`:\n",
|
||||
"\n",
|
||||
"- **exact** = bit-for-bit identical\n",
|
||||
"- **approx** = equal within tiny floating-point tolerance (this is the one that matters)\n",
|
||||
"- **maxdiff** = largest single difference"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "2c6c9ccf",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def cmp(name, dt, t):\n",
|
||||
" ex = torch.all(dt == t.grad).item()\n",
|
||||
" app = torch.allclose(dt, t.grad, rtol=1e-4, atol=1e-7)\n",
|
||||
" maxdiff = (dt - t.grad).abs().max().item()\n",
|
||||
" print(f\"{name:12s} | exact: {str(ex):5s} | approx: {str(app):5s} | maxdiff: {maxdiff:.2e}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "ed7b6dec",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### 1.1 Parameters\n",
|
||||
"\n",
|
||||
"We initialise the same MLP + BatchNorm as notebook `03`. Biases are made small but non-zero\n",
|
||||
"so every gradient path is actually exercised (a zero would hide bugs)."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "70e88d72",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"n_embd, n_hidden = 10, 64\n",
|
||||
"g = torch.Generator().manual_seed(2147483647)\n",
|
||||
"\n",
|
||||
"C = torch.randn((vocab_size, n_embd), generator=g)\n",
|
||||
"W1 = torch.randn((n_embd * block_size, n_hidden), generator=g) * (5/3) / ((n_embd * block_size) ** 0.5)\n",
|
||||
"b1 = torch.randn(n_hidden, generator=g) * 0.1\n",
|
||||
"W2 = torch.randn((n_hidden, vocab_size), generator=g) * 0.1\n",
|
||||
"b2 = torch.randn(vocab_size, generator=g) * 0.1\n",
|
||||
"bngain = torch.randn((1, n_hidden), generator=g) * 0.1 + 1.0\n",
|
||||
"bnbias = torch.randn((1, n_hidden), generator=g) * 0.1\n",
|
||||
"\n",
|
||||
"parameters = [C, W1, b1, W2, b2, bngain, bnbias]\n",
|
||||
"for p in parameters:\n",
|
||||
" p.requires_grad = True\n",
|
||||
"print(\"total parameters:\", sum(p.nelement() for p in parameters))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "04fa6a1e",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### 1.2 Forward pass, broken into tiny steps\n",
|
||||
"\n",
|
||||
"Cross-entropy is written out by hand (max-subtract, exp, sum, log) so we can backprop through\n",
|
||||
"each piece. Every intermediate calls `retain_grad()` so PyTorch will store its gradient for us\n",
|
||||
"to compare against."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "6ad34b3b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"n = 32 # fixed batch size for the manual math\n",
|
||||
"ix = torch.randint(0, Xtr.shape[0], (n,), generator=g)\n",
|
||||
"Xb, Yb = Xtr[ix], Ytr[ix]\n",
|
||||
"\n",
|
||||
"emb = C[Xb] # embed the characters\n",
|
||||
"embcat = emb.view(emb.shape[0], -1) # flatten\n",
|
||||
"# linear 1\n",
|
||||
"hprebn = embcat @ W1 + b1\n",
|
||||
"# batchnorm\n",
|
||||
"bnmeani = 1 / n * hprebn.sum(0, keepdim=True)\n",
|
||||
"bndiff = hprebn - bnmeani\n",
|
||||
"bndiff2 = bndiff ** 2\n",
|
||||
"bnvar = 1 / (n - 1) * bndiff2.sum(0, keepdim=True) # Bessel's correction (n-1)\n",
|
||||
"bnvar_inv = (bnvar + 1e-5) ** -0.5\n",
|
||||
"bnraw = bndiff * bnvar_inv\n",
|
||||
"hpreact = bngain * bnraw + bnbias\n",
|
||||
"# nonlinearity\n",
|
||||
"h = torch.tanh(hpreact)\n",
|
||||
"# linear 2\n",
|
||||
"logits = h @ W2 + b2\n",
|
||||
"# cross-entropy (manual)\n",
|
||||
"logit_maxes = logits.max(1, keepdim=True).values\n",
|
||||
"norm_logits = logits - logit_maxes\n",
|
||||
"counts = norm_logits.exp()\n",
|
||||
"counts_sum = counts.sum(1, keepdim=True)\n",
|
||||
"counts_sum_inv = counts_sum ** -1\n",
|
||||
"probs = counts * counts_sum_inv\n",
|
||||
"logprobs = probs.log()\n",
|
||||
"loss = -logprobs[range(n), Yb].mean()\n",
|
||||
"\n",
|
||||
"for t in [logprobs, probs, counts, counts_sum, counts_sum_inv, norm_logits, logit_maxes,\n",
|
||||
" logits, h, hpreact, bnraw, bnvar_inv, bnvar, bndiff2, bndiff, bnmeani, hprebn,\n",
|
||||
" embcat, emb]:\n",
|
||||
" t.retain_grad()\n",
|
||||
"for p in parameters:\n",
|
||||
" p.retain_grad()\n",
|
||||
"loss.backward()\n",
|
||||
"print(\"loss:\", loss.item())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "6fcbe6af",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### 2.1 Backprop through cross-entropy\n",
|
||||
"\n",
|
||||
"Walk backward from `loss`. For each line of the forward pass we write the local derivative and\n",
|
||||
"multiply by the gradient coming from above (the chain rule). Read each line as: *how does this\n",
|
||||
"intermediate affect the loss?*"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "c7bc68e2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"dlogprobs = torch.zeros_like(logprobs)\n",
|
||||
"dlogprobs[range(n), Yb] = -1.0 / n # loss = -mean of the chosen logprobs\n",
|
||||
"dprobs = (1.0 / probs) * dlogprobs # logprobs = log(probs)\n",
|
||||
"dcounts_sum_inv = (counts * dprobs).sum(1, keepdim=True)\n",
|
||||
"dcounts = counts_sum_inv * dprobs # probs = counts * counts_sum_inv\n",
|
||||
"dcounts_sum = (-counts_sum ** -2) * dcounts_sum_inv\n",
|
||||
"dcounts += torch.ones_like(counts) * dcounts_sum\n",
|
||||
"dnorm_logits = counts * dcounts # counts = exp(norm_logits)\n",
|
||||
"dlogits = dnorm_logits.clone()\n",
|
||||
"dlogit_maxes = (-dnorm_logits).sum(1, keepdim=True)\n",
|
||||
"dlogits += F.one_hot(logits.max(1).indices, num_classes=logits.shape[1]) * dlogit_maxes\n",
|
||||
"\n",
|
||||
"cmp(\"logprobs\", dlogprobs, logprobs)\n",
|
||||
"cmp(\"probs\", dprobs, probs)\n",
|
||||
"cmp(\"counts_sum_inv\", dcounts_sum_inv, counts_sum_inv)\n",
|
||||
"cmp(\"counts_sum\", dcounts_sum, counts_sum)\n",
|
||||
"cmp(\"counts\", dcounts, counts)\n",
|
||||
"cmp(\"norm_logits\", dnorm_logits, norm_logits)\n",
|
||||
"cmp(\"logit_maxes\", dlogit_maxes, logit_maxes)\n",
|
||||
"cmp(\"logits\", dlogits, logits)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "51a6e8d7",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### 2.2 Backprop through the second linear layer and tanh\n",
|
||||
"\n",
|
||||
"`logits = h @ W2 + b2`. Matrix-multiply gradients follow the standard transpose rules; `tanh`\n",
|
||||
"contributes its local derivative `1 - tanh(x)^2`."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "39f0e645",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"dh = dlogits @ W2.T\n",
|
||||
"dW2 = h.T @ dlogits\n",
|
||||
"db2 = dlogits.sum(0)\n",
|
||||
"dhpreact = (1.0 - h ** 2) * dh # tanh local derivative\n",
|
||||
"\n",
|
||||
"cmp(\"h\", dh, h)\n",
|
||||
"cmp(\"W2\", dW2, W2)\n",
|
||||
"cmp(\"b2\", db2, b2)\n",
|
||||
"cmp(\"hpreact\", dhpreact, hpreact)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "a7ecc492",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### 2.3 Backprop through BatchNorm\n",
|
||||
"\n",
|
||||
"This is the famously fiddly part. We go through each step: the scale/shift, the normalisation,\n",
|
||||
"the variance, and the mean. Note how the **mean and variance both depend on every example**, so\n",
|
||||
"gradients fan out across the batch."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "6f4f5d49",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"dbngain = (bnraw * dhpreact).sum(0, keepdim=True)\n",
|
||||
"dbnbias = dhpreact.sum(0, keepdim=True)\n",
|
||||
"dbnraw = bngain * dhpreact\n",
|
||||
"dbndiff = bnvar_inv * dbnraw\n",
|
||||
"dbnvar_inv = (bndiff * dbnraw).sum(0, keepdim=True)\n",
|
||||
"dbnvar = (-0.5 * (bnvar + 1e-5) ** -1.5) * dbnvar_inv\n",
|
||||
"dbndiff2 = (1.0 / (n - 1)) * torch.ones_like(bndiff2) * dbnvar\n",
|
||||
"dbndiff += (2 * bndiff) * dbndiff2\n",
|
||||
"dhprebn = dbndiff.clone()\n",
|
||||
"dbnmeani = (-dbndiff).sum(0)\n",
|
||||
"dhprebn += 1.0 / n * (torch.ones_like(hprebn) * dbnmeani)\n",
|
||||
"\n",
|
||||
"cmp(\"bngain\", dbngain, bngain)\n",
|
||||
"cmp(\"bnbias\", dbnbias, bnbias)\n",
|
||||
"cmp(\"bnraw\", dbnraw, bnraw)\n",
|
||||
"cmp(\"bnvar_inv\", dbnvar_inv, bnvar_inv)\n",
|
||||
"cmp(\"bnvar\", dbnvar, bnvar)\n",
|
||||
"cmp(\"bndiff2\", dbndiff2, bndiff2)\n",
|
||||
"cmp(\"bndiff\", dbndiff, bndiff)\n",
|
||||
"cmp(\"bnmeani\", dbnmeani, bnmeani)\n",
|
||||
"cmp(\"hprebn\", dhprebn, hprebn)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "3be0bf67",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### 2.4 Backprop through the first linear layer and the embedding\n",
|
||||
"\n",
|
||||
"`hprebn = embcat @ W1 + b1`, then un-flatten back to `emb`, then scatter each row's gradient to\n",
|
||||
"the right row of the embedding table `C` (an embedding lookup is just \"pick rows\", so its\n",
|
||||
"backward is \"add the gradient back into those rows\")."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "b1317637",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"dembcat = dhprebn @ W1.T\n",
|
||||
"dW1 = embcat.T @ dhprebn\n",
|
||||
"db1 = dhprebn.sum(0)\n",
|
||||
"demb = dembcat.view(emb.shape)\n",
|
||||
"\n",
|
||||
"dC = torch.zeros_like(C)\n",
|
||||
"for k in range(Xb.shape[0]):\n",
|
||||
" for j in range(Xb.shape[1]):\n",
|
||||
" dC[Xb[k, j]] += demb[k, j]\n",
|
||||
"\n",
|
||||
"cmp(\"embcat\", dembcat, embcat)\n",
|
||||
"cmp(\"W1\", dW1, W1)\n",
|
||||
"cmp(\"b1\", db1, b1)\n",
|
||||
"cmp(\"emb\", demb, emb)\n",
|
||||
"cmp(\"C\", dC, C)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "3f6d3bff",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### 3.1 The elegant shortcuts\n",
|
||||
"\n",
|
||||
"Doing it step-by-step proves we understand it. But two of those backward passes collapse into\n",
|
||||
"beautiful one-liners once you do the algebra:\n",
|
||||
"\n",
|
||||
"- **Cross-entropy:** `dlogits = softmax(logits)`, then subtract 1 from the true class, divide by n.\n",
|
||||
" Intuition: push down whatever probability the model gave, pull up the correct class.\n",
|
||||
"- **BatchNorm:** one expression for `dhprebn` (no need for all the intermediate d's).\n",
|
||||
"\n",
|
||||
"We verify both against PyTorch."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "91248dc5",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# elegant cross-entropy backward\n",
|
||||
"dlogits_fast = F.softmax(logits, 1)\n",
|
||||
"dlogits_fast[range(n), Yb] -= 1\n",
|
||||
"dlogits_fast /= n\n",
|
||||
"print(\"dlogits | approx:\", torch.allclose(dlogits_fast, logits.grad, rtol=1e-4, atol=1e-7),\n",
|
||||
" \" | maxdiff:\", (dlogits_fast - logits.grad).abs().max().item())\n",
|
||||
"\n",
|
||||
"# elegant batchnorm backward\n",
|
||||
"dhprebn_fast = bngain * bnvar_inv / n * (\n",
|
||||
" n * dhpreact - dhpreact.sum(0) - n / (n - 1) * bnraw * (dhpreact * bnraw).sum(0)\n",
|
||||
")\n",
|
||||
"print(\"dhprebn | approx:\", torch.allclose(dhprebn_fast, hprebn.grad, rtol=1e-4, atol=1e-7),\n",
|
||||
" \" | maxdiff:\", (dhprebn_fast - hprebn.grad).abs().max().item())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "09834b6c",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### 3.2 Train with 100% hand-written gradients\n",
|
||||
"\n",
|
||||
"The real proof: turn off autograd and train using only our manual gradients (the elegant\n",
|
||||
"versions, plus the manual `W2/b2/W1/b1/C/bngain/bnbias`). If the loss falls, our calculus is\n",
|
||||
"correct."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "9bfb8004",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# fresh params\n",
|
||||
"g = torch.Generator().manual_seed(2147483647)\n",
|
||||
"C = torch.randn((vocab_size, n_embd), generator=g)\n",
|
||||
"W1 = torch.randn((n_embd * block_size, n_hidden), generator=g) * (5/3) / ((n_embd * block_size) ** 0.5)\n",
|
||||
"b1 = torch.randn(n_hidden, generator=g) * 0.1\n",
|
||||
"W2 = torch.randn((n_hidden, vocab_size), generator=g) * 0.1\n",
|
||||
"b2 = torch.randn(vocab_size, generator=g) * 0.1\n",
|
||||
"bngain = torch.randn((1, n_hidden), generator=g) * 0.1 + 1.0\n",
|
||||
"bnbias = torch.randn((1, n_hidden), generator=g) * 0.1\n",
|
||||
"parameters = [C, W1, b1, W2, b2, bngain, bnbias]\n",
|
||||
"for p in parameters:\n",
|
||||
" p.requires_grad = True\n",
|
||||
"\n",
|
||||
"n = 32\n",
|
||||
"for i in range(8000):\n",
|
||||
" ix = torch.randint(0, Xtr.shape[0], (n,), generator=g)\n",
|
||||
" Xb, Yb = Xtr[ix], Ytr[ix]\n",
|
||||
"\n",
|
||||
" # forward (keep what the manual backward needs)\n",
|
||||
" emb = C[Xb]\n",
|
||||
" embcat = emb.view(emb.shape[0], -1)\n",
|
||||
" hprebn = embcat @ W1 + b1\n",
|
||||
" bnmeani = hprebn.mean(0, keepdim=True)\n",
|
||||
" bnvar = hprebn.var(0, keepdim=True, unbiased=True)\n",
|
||||
" bnvar_inv = (bnvar + 1e-5) ** -0.5\n",
|
||||
" bnraw = (hprebn - bnmeani) * bnvar_inv\n",
|
||||
" hpreact = bngain * bnraw + bnbias\n",
|
||||
" h = torch.tanh(hpreact)\n",
|
||||
" logits = h @ W2 + b2\n",
|
||||
" loss = F.cross_entropy(logits, Yb)\n",
|
||||
"\n",
|
||||
" # manual backward (no loss.backward())\n",
|
||||
" dlogits = F.softmax(logits, 1)\n",
|
||||
" dlogits[range(n), Yb] -= 1\n",
|
||||
" dlogits /= n\n",
|
||||
" dh = dlogits @ W2.T\n",
|
||||
" dW2 = h.T @ dlogits\n",
|
||||
" db2 = dlogits.sum(0)\n",
|
||||
" dhpreact = (1.0 - h ** 2) * dh\n",
|
||||
" dbngain = (bnraw * dhpreact).sum(0, keepdim=True)\n",
|
||||
" dbnbias = dhpreact.sum(0, keepdim=True)\n",
|
||||
" dhprebn = bngain * bnvar_inv / n * (\n",
|
||||
" n * dhpreact - dhpreact.sum(0) - n / (n - 1) * bnraw * (dhpreact * bnraw).sum(0)\n",
|
||||
" )\n",
|
||||
" dembcat = dhprebn @ W1.T\n",
|
||||
" dW1 = embcat.T @ dhprebn\n",
|
||||
" db1 = dhprebn.sum(0)\n",
|
||||
" demb = dembcat.view(emb.shape)\n",
|
||||
" dC = torch.zeros_like(C)\n",
|
||||
" dC.index_add_(0, Xb.view(-1), demb.view(-1, n_embd))\n",
|
||||
"\n",
|
||||
" grads = [dC, dW1, db1, dW2, db2, dbngain, dbnbias]\n",
|
||||
" lr = 0.1 if i < 6000 else 0.01\n",
|
||||
" with torch.no_grad():\n",
|
||||
" for p, gr in zip(parameters, grads):\n",
|
||||
" p += -lr * gr\n",
|
||||
"\n",
|
||||
" if i % 2000 == 0 or i == 7999:\n",
|
||||
" print(f\"step {i:5d} loss {loss.item():.4f}\")\n",
|
||||
"\n",
|
||||
"print(\"trained using only hand-written gradients.\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "422e7804",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"**Your turn 4** — In section 2.3, deliberately introduce a bug (e.g. drop the\n",
|
||||
"`n/(n-1)` Bessel factor in `dbnvar`) and re-run. Watch `cmp` flip to `approx: False` and the\n",
|
||||
"`maxdiff` jump. Then fix it. This is exactly how you debug a real backward pass."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "36c4bc1d",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## What's next\n",
|
||||
"\n",
|
||||
"You backpropagated by hand through cross-entropy, a linear layer, `tanh`, BatchNorm, and an\n",
|
||||
"embedding, matched PyTorch exactly, and trained a net with zero autograd. Backprop is now\n",
|
||||
"yours.\n",
|
||||
"\n",
|
||||
"Where this leads in the series:\n",
|
||||
"\n",
|
||||
"- `05_wavenet.ipynb` — stack layers into a deeper, hierarchical model using `torch.nn`-style\n",
|
||||
" building blocks.\n",
|
||||
"- `06_build_gpt_attention.ipynb` — the transformer; `loss.backward()` there is doing exactly\n",
|
||||
" what you just did, across far more nodes.\n",
|
||||
"\n",
|
||||
"**Checklist**\n",
|
||||
"- [ ] Forward pass saved every intermediate\n",
|
||||
"- [ ] Manual backward matched PyTorch (cross-entropy, linear, tanh)\n",
|
||||
"- [ ] Manual backward through BatchNorm matched\n",
|
||||
"- [ ] Embedding backward = scatter-add into rows\n",
|
||||
"- [ ] Elegant one-line cross-entropy and BatchNorm gradients\n",
|
||||
"- [ ] Trained a net with only hand-written gradients"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": ".venv (3.12.10)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"name": "python",
|
||||
"version": "3.12.10"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
Reference in New Issue
Block a user