diff --git a/mcp-server/src/data/hints.generated.ts b/mcp-server/src/data/hints.generated.ts index dacfed0..7bac22e 100644 --- a/mcp-server/src/data/hints.generated.ts +++ b/mcp-server/src/data/hints.generated.ts @@ -455,13 +455,15 @@ export const authoredHints: Record = { ], "gradient-checkpointing": [ "In forward: stash fn on ctx, ctx.save_for_backward(*args), and run fn(*args) inside torch.no_grad() so no intermediate activation is kept.", - "In backward: detach the saved inputs, requires_grad_ them, re-run fn inside torch.enable_grad(), then torch.autograd.grad(output, inputs, grad_outputs).", - "backward must return one gradient per forward argument — None for fn itself, then a gradient (or None) for each tensor input." + "In backward: detach the saved inputs, requires_grad_ them, re-run fn inside torch.enable_grad(), then torch.autograd.backward(outputs, grad_outputs) — unlike autograd.grad on the inputs, that also reaches fn's own parameters.", + "backward must return one gradient per forward argument — None for fn itself, then a gradient (or None) for each tensor input (read them off the detached copies).", + "Autograd only calls backward if some tensor argument requires grad. When none does (plain input data), have checkpoint() pass an extra dummy tensor with requires_grad=True, or the wrapped layer's weights never train." ], "v3-16": [ "In forward: stash fn on ctx, ctx.save_for_backward(*args), and run fn(*args) inside torch.no_grad() so no intermediate activation is kept.", - "In backward: detach the saved inputs, requires_grad_ them, re-run fn inside torch.enable_grad(), then torch.autograd.grad(output, inputs, grad_outputs).", - "backward must return one gradient per forward argument — None for fn itself, then a gradient (or None) for each tensor input." + "In backward: detach the saved inputs, requires_grad_ them, re-run fn inside torch.enable_grad(), then torch.autograd.backward(outputs, grad_outputs) — unlike autograd.grad on the inputs, that also reaches fn's own parameters.", + "backward must return one gradient per forward argument — None for fn itself, then a gradient (or None) for each tensor input (read them off the detached copies).", + "Autograd only calls backward if some tensor argument requires grad. When none does (plain input data), have checkpoint() pass an extra dummy tensor with requires_grad=True, or the wrapped layer's weights never train." ], "mixture-of-experts": [ "Route each token to its top-k experts and weight their outputs by the gate probabilities.", diff --git a/problems.json b/problems.json index 79c36ac..5c2d3ce 100644 --- a/problems.json +++ b/problems.json @@ -1712,8 +1712,9 @@ ], "hints": [ "In forward: stash fn on ctx, ctx.save_for_backward(*args), and run fn(*args) inside torch.no_grad() so no intermediate activation is kept.", - "In backward: detach the saved inputs, requires_grad_ them, re-run fn inside torch.enable_grad(), then torch.autograd.grad(output, inputs, grad_outputs).", - "backward must return one gradient per forward argument \u2014 None for fn itself, then a gradient (or None) for each tensor input." + "In backward: detach the saved inputs, requires_grad_ them, re-run fn inside torch.enable_grad(), then torch.autograd.backward(outputs, grad_outputs) \u2014 unlike autograd.grad on the inputs, that also reaches fn's own parameters.", + "backward must return one gradient per forward argument \u2014 None for fn itself, then a gradient (or None) for each tensor input (read them off the detached copies).", + "Autograd only calls backward if some tensor argument requires grad. When none does (plain input data), have checkpoint() pass an extra dummy tensor with requires_grad=True, or the wrapped layer's weights never train." ] }, { diff --git a/python/test_grader.py b/python/test_grader.py index 72abac7..04f1e5c 100644 --- a/python/test_grader.py +++ b/python/test_grader.py @@ -156,6 +156,76 @@ def __init__(self): self.f2 = nn.Linear(128, 100) +# ---------------------------------------------------------------------- rope +def rotate_half(x): + x1, x2 = x.chunk(2, dim=-1) + return torch.cat((-x2, x1), dim=-1) + + +def rope_apply(q, k, cos, sin): + return q * cos + rotate_half(q) * sin, k * cos + rotate_half(k) * sin + + +def rope_noop(q, k, cos, sin): + """Broken: ignores cos/sin and returns the inputs untouched.""" + return q, k + + +def rope_query_only(q, k, cos, sin): + """Broken: rotates the query but forgets the key.""" + return q * cos + rotate_half(q) * sin, k + + +# --------------------------------------------------------- gradient checkpointing +class CkptFn(torch.autograd.Function): + @staticmethod + def forward(ctx, fn, *args): + ctx.fn = fn + ctx.save_for_backward(*args) + with torch.no_grad(): + return fn(*args) + + @staticmethod + def backward(ctx, *grad_outputs): + inputs = [t.detach().requires_grad_(t.requires_grad) for t in ctx.saved_tensors] + with torch.enable_grad(): + out = ctx.fn(*inputs) + outs = out if isinstance(out, tuple) else (out,) + pairs = [(o, g) for o, g in zip(outs, grad_outputs) if o.requires_grad] + if pairs: + torch.autograd.backward([o for o, _ in pairs], [g for _, g in pairs]) + return (None, *(t.grad if t.requires_grad else None for t in inputs)) + + +def ckpt(fn, *args): + if torch.is_grad_enabled() and not any( + isinstance(a, torch.Tensor) and a.requires_grad for a in args): + dummy = torch.empty(0, requires_grad=True) + return CkptFn.apply(lambda _d, *a: fn(*a), dummy, *args) + return CkptFn.apply(fn, *args) + + +def ckpt_no_dummy(fn, *args): + """Broken: with a plain (no-grad) input backward never runs, so fn never trains.""" + return CkptFn.apply(fn, *args) + + +class CkptFnInputsOnly(CkptFn): + """Broken: autograd.grad w.r.t. the inputs, so fn's parameters get nothing.""" + + @staticmethod + def backward(ctx, *grad_outputs): + inputs = [t.detach().requires_grad_(True) for t in ctx.saved_tensors] + with torch.enable_grad(): + out = ctx.fn(*inputs) + grads = torch.autograd.grad(out, inputs, grad_outputs, allow_unused=True) + return (None, *grads) + + +def ckpt_inputs_only(fn, *args): + return CkptFnInputsOnly.apply(fn, *args) + + CASES = [ # (label, expected, problem_id, args, kwargs) ("kv-cache canonical", True, "kv-cache", (), dict(KVCache=KVCache, CachedAttention=Attn)), @@ -166,6 +236,12 @@ def __init__(self): ("cnn canonical", True, "cnn", (Cnn,), {}), ("cnn divergent", True, "cnn", (DivergentCnn,), {}), ("cnn wrong classes", False, "cnn", (WrongClassCount,), {}), + ("rope correct", True, "rotary-positional-embedding", (rotate_half, rope_apply), {}), + ("rope no-op", False, "rotary-positional-embedding", (rotate_half, rope_noop), {}), + ("rope query only", False, "rotary-positional-embedding", (rotate_half, rope_query_only), {}), + ("ckpt correct", True, "gradient-checkpointing", (CkptFn, ckpt), {}), + ("ckpt no dummy input", False, "gradient-checkpointing", (CkptFn, ckpt_no_dummy), {}), + ("ckpt inputs-only grads", False, "gradient-checkpointing", (CkptFnInputsOnly, ckpt_inputs_only), {}), ] diff --git a/python/torchleet/problems/gradient_checkpointing.py b/python/torchleet/problems/gradient_checkpointing.py index a2928e5..4ce6df9 100644 --- a/python/torchleet/problems/gradient_checkpointing.py +++ b/python/torchleet/problems/gradient_checkpointing.py @@ -10,6 +10,11 @@ an activation created inside the checkpointed function and asserts it has been freed by the time forward returns. An implementation that quietly calls fn(*args) passes every gradient check and fails that one, which is the whole exercise. + +The other classic bug is a checkpoint that silently freezes the layer it wraps: +the parameters fn closes over must receive gradients (autograd.grad on the inputs +alone never reaches them), including when the input is plain data that does not +require grad, as the first layer of any real network is. """ import gc import weakref @@ -18,8 +23,6 @@ import torch.nn as nn import torch.nn.functional as F -from torchleet.runner import Skip - ENTRIES = ["CheckpointFunction", "checkpoint"] DEVICE = "cpu" EXTRAS = [] @@ -28,9 +31,13 @@ "In forward: stash fn on ctx, ctx.save_for_backward(*args), and run " "fn(*args) inside torch.no_grad() so no intermediate activation is kept.", "In backward: detach the saved inputs, requires_grad_ them, re-run fn inside " - "torch.enable_grad(), then torch.autograd.grad(output, inputs, grad_outputs).", + "torch.enable_grad(), then torch.autograd.backward(outputs, grad_outputs) — " + "unlike autograd.grad on the inputs, that also reaches fn's own parameters.", "backward must return one gradient per forward argument — None for fn itself, " - "then a gradient (or None) for each tensor input.", + "then a gradient (or None) for each tensor input (read them off the detached copies).", + "Autograd only calls backward if some tensor argument requires grad. When none " + "does (plain input data), have checkpoint() pass an extra dummy tensor with " + "requires_grad=True, or the wrapped layer's weights never train.", ] BATCH, DIM, HIDDEN = 6, 8, 16 @@ -153,10 +160,9 @@ def fn(t, weight): def check_module_parameter_gradients(ns): """Parameters captured inside fn rather than passed as arguments. - torch.utils.checkpoint reaches them; the recomputation-plus-autograd.grad - formulation in this problem's solution only reaches its tensor arguments. If - an implementation does populate them they must be right, otherwise this is - reported as not verified rather than silently passed. + This is how checkpointing is used in practice (checkpoint(block, x)), and + torch.utils.checkpoint reaches them. A backward that calls autograd.grad on + the saved inputs alone never does, so the wrapped layer silently never trains. """ torch.manual_seed(13) layer = nn.Linear(DIM, DIM) @@ -167,11 +173,11 @@ def check_module_parameter_gradients(ns): plain(x).pow(2).sum().backward() ns.checkpoint(layer, x.detach().requires_grad_(True)).pow(2).sum().backward() - if layer.weight.grad is None and layer.bias.grad is None: - raise Skip( - "checkpoint() propagates gradients only to the tensors passed as " - "arguments, not to parameters captured inside fn (this problem's own " - "solution behaves the same way) — not verified") + assert layer.weight.grad is not None or layer.bias.grad is not None, ( + "no gradient reached the parameters of the checkpointed module, so it " + "would never train — autograd.grad(outputs, inputs) only covers the tensor " + "arguments; use torch.autograd.backward(outputs, grad_outputs) in backward " + "so the parameters fn closes over are reached too") for name, p, q in (("weight", plain.weight, layer.weight), ("bias", plain.bias, layer.bias)): assert q.grad is not None, \ @@ -182,6 +188,36 @@ def check_module_parameter_gradients(ns): f"(max diff {diff:.2e})") +def check_plain_input_still_trains_parameters(ns): + """The first layer of a network is fed data that does not require grad. + + If no argument of the autograd.Function requires grad, autograd never records + it and backward never runs, so a checkpointed first layer would stay frozen. + """ + torch.manual_seed(14) + layer = nn.Linear(DIM, DIM) + plain = nn.Linear(DIM, DIM) + plain.load_state_dict(layer.state_dict()) + x = torch.randn(BATCH, DIM) # plain data, requires_grad=False + + plain(x).pow(2).sum().backward() + out = ns.checkpoint(layer, x) + assert out.requires_grad, ( + "checkpoint(layer, x) with an input that does not require grad returned a " + "tensor with no grad_fn, so backward never runs and the layer never trains " + "— when no argument requires grad, pass an extra dummy tensor with " + "requires_grad=True to CheckpointFunction.apply") + out.pow(2).sum().backward() + for name, p, q in (("weight", plain.weight, layer.weight), + ("bias", plain.bias, layer.bias)): + assert q.grad is not None, ( + f"with a plain (no-grad) input the module's {name} received no gradient") + diff = (p.grad - q.grad).abs().max().item() + assert torch.allclose(p.grad, q.grad, atol=1e-6), ( + f"with a plain (no-grad) input the gradient on the module's {name} " + f"differs from the uncheckpointed run (max diff {diff:.2e})") + + CHECKS = [ check_forward_output_matches, check_output_is_differentiable, @@ -189,4 +225,5 @@ def check_module_parameter_gradients(ns): check_chained_checkpoints_match, check_intermediates_are_not_kept, check_module_parameter_gradients, + check_plain_input_still_trains_parameters, ] diff --git a/python/torchleet/problems/rotary_positional_embedding.py b/python/torchleet/problems/rotary_positional_embedding.py index 4a7be6f..13c0418 100644 --- a/python/torchleet/problems/rotary_positional_embedding.py +++ b/python/torchleet/problems/rotary_positional_embedding.py @@ -2,7 +2,9 @@ RoPE is a rotation, which gives two exact properties to test: applying rotate_half twice negates the vector, and rotating with a unit (cos, sin) pair -preserves the norm. +preserves the norm. Neither of those, nor the zero-angle check, can tell a real +rotation from returning q and k untouched, so check_quarter_turn pins the angle: +at cos=0, sin=1 the result must be exactly rotate_half of the input. """ import torch @@ -48,6 +50,27 @@ def check_identity_rotation(ns): assert torch.allclose(rk, k, atol=1e-6), "cos=1, sin=0 must leave k unchanged" +def check_quarter_turn(ns): + """cos=0, sin=1 is a 90-degree turn: q*0 + rotate_half(q)*1 = rotate_half(q). + + Uses the solver's own rotate_half, so any consistent pairing convention passes, + but returning q or k unrotated (or ignoring sin) does not. + """ + torch.manual_seed(1) + q, k = torch.randn(B, H, S, D), torch.randn(B, H, S, D) + cos, sin = torch.zeros(S, D), torch.ones(S, D) + rq, rk = ns.apply_rotary_pos_emb(q, k, cos, sin) + assert not torch.allclose(rq, q, atol=1e-4), ( + "with cos=0, sin=1 the query came back unchanged — apply_rotary_pos_emb " + "must actually use cos and sin (q*cos + rotate_half(q)*sin)") + assert torch.allclose(rq, ns.rotate_half(q), atol=1e-6), ( + "with cos=0, sin=1 the rotated query should equal rotate_half(q) — " + "check the q*cos + rotate_half(q)*sin formula") + assert torch.allclose(rk, ns.rotate_half(k), atol=1e-6), ( + "with cos=0, sin=1 the rotated key should equal rotate_half(k) — the key " + "must be rotated exactly like the query") + + def check_rotation_preserves_norm(ns): """A rotation cannot change a vector's length.""" torch.manual_seed(0) @@ -62,4 +85,4 @@ def check_rotation_preserves_norm(ns): CHECKS = [check_rotate_half_shape, check_rotate_half_twice_negates, check_rotate_half_permutes_values, check_identity_rotation, - check_rotation_preserves_norm] + check_quarter_turn, check_rotation_preserves_norm] diff --git a/space/problems.json b/space/problems.json index 79c36ac..5c2d3ce 100644 --- a/space/problems.json +++ b/space/problems.json @@ -1712,8 +1712,9 @@ ], "hints": [ "In forward: stash fn on ctx, ctx.save_for_backward(*args), and run fn(*args) inside torch.no_grad() so no intermediate activation is kept.", - "In backward: detach the saved inputs, requires_grad_ them, re-run fn inside torch.enable_grad(), then torch.autograd.grad(output, inputs, grad_outputs).", - "backward must return one gradient per forward argument \u2014 None for fn itself, then a gradient (or None) for each tensor input." + "In backward: detach the saved inputs, requires_grad_ them, re-run fn inside torch.enable_grad(), then torch.autograd.backward(outputs, grad_outputs) \u2014 unlike autograd.grad on the inputs, that also reaches fn's own parameters.", + "backward must return one gradient per forward argument \u2014 None for fn itself, then a gradient (or None) for each tensor input (read them off the detached copies).", + "Autograd only calls backward if some tensor argument requires grad. When none does (plain input data), have checkpoint() pass an extra dummy tensor with requires_grad=True, or the wrapped layer's weights never train." ] }, { diff --git a/torch/easy/quantize-lm/quantize-language-model.ipynb b/torch/easy/quantize-lm/quantize-language-model.ipynb index 73993bf..dc49570 100644 --- a/torch/easy/quantize-lm/quantize-language-model.ipynb +++ b/torch/easy/quantize-lm/quantize-language-model.ipynb @@ -41,14 +41,23 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": null, "metadata": {}, "outputs": [], "source": [ "import torch\n", "import torch.nn as nn\n", "import torch.optim as optim\n", - "from torch.quantization import quantize_dynamic" + "from torch.ao.quantization import quantize_dynamic\n", + "\n", + "# Dynamic quantization needs a quantized backend. On Apple Silicon (and some other\n", + "# builds) the default engine is 'none', so pick one this torch build supports.\n", + "if torch.backends.quantized.engine == \"none\":\n", + " engines = [e for e in torch.backends.quantized.supported_engines if e != \"none\"]\n", + " if not engines:\n", + " raise RuntimeError(\"this torch build has no quantized backend (fbgemm/qnnpack)\")\n", + " torch.backends.quantized.engine = engines[0]\n", + "print(f\"Quantized engine: {torch.backends.quantized.engine}\")" ] }, { @@ -69,7 +78,7 @@ }, { "cell_type": "code", - "execution_count": 5, + "execution_count": null, "metadata": {}, "outputs": [], "source": [ @@ -87,27 +96,18 @@ "num_layers = 2\n", "model = LanguageModel(vocab_size, embed_size, hidden_size, num_layers)\n", "\n", - "criterion = nn.CrossEntropyLoss()\n", + "# The model ends in a softmax, so it returns probabilities, not logits.\n", + "# nn.CrossEntropyLoss would apply a second softmax on top and the loss would barely\n", + "# move, so train with NLLLoss on the log-probabilities instead.\n", + "criterion = nn.NLLLoss()\n", "optimizer = optim.Adam(model.parameters(), lr=0.001)" ] }, { "cell_type": "code", - "execution_count": 6, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Epoch [1/5] - Loss: 3.9118\n", - "Epoch [2/5] - Loss: 3.9113\n", - "Epoch [3/5] - Loss: 3.9108\n", - "Epoch [4/5] - Loss: 3.9103\n", - "Epoch [5/5] - Loss: 3.9097\n" - ] - } - ], + "outputs": [], "source": [ "# Training loop\n", "epochs = 5\n", @@ -115,7 +115,7 @@ " model.train()\n", " optimizer.zero_grad()\n", " output = model(X_train)\n", - " loss = criterion(output, y_train)\n", + " loss = criterion(torch.log(output.clamp_min(1e-9)), y_train)\n", " loss.backward()\n", " optimizer.step()\n", "\n", @@ -126,49 +126,30 @@ "# Quantization: Apply dynamic quantization to the language model\n", "quantized_model = quantize_dynamic(model, {nn.Linear, nn.LSTM}, dtype=torch.qint8)\n", "\n", - "# Save the quantized model\n", - "torch.save(quantized_model.state_dict(), \"quantized_language_model.pth\")\n" + "# Save the float weights. A dynamically quantized state_dict stores the LSTM's packed\n", + "# weights as torch.ScriptObject, which torch.load(weights_only=True) (the default\n", + "# since PyTorch 2.6) refuses to unpickle. Dynamic quantization is deterministic, so\n", + "# quantizing the float weights again after loading rebuilds the same model.\n", + "torch.save(model.state_dict(), \"language_model_fp32.pth\")\n" ] }, { "cell_type": "code", - "execution_count": 7, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "data": { - "text/plain": [ - "" - ] - }, - "execution_count": 7, - "metadata": {}, - "output_type": "execute_result" - } - ], + "outputs": [], "source": [ - "# Load the quantized model and test it\n", + "# Load the float weights, then apply the same dynamic quantization again\n", "quantized_model = LanguageModel(vocab_size, embed_size, hidden_size, num_layers)\n", - "\n", - "# Apply dynamic quantization on the model after defining it\n", - "quantized_model = quantize_dynamic(quantized_model, {nn.Linear, nn.LSTM}, dtype=torch.qint8)\n", - "\n", - "quantized_model.load_state_dict(torch.load(\"quantized_language_model.pth\"))" + "quantized_model.load_state_dict(torch.load(\"language_model_fp32.pth\", weights_only=True))\n", + "quantized_model = quantize_dynamic(quantized_model, {nn.Linear, nn.LSTM}, dtype=torch.qint8)" ] }, { "cell_type": "code", - "execution_count": 8, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Prediction for input [[15, 28, 33, 19, 37, 24, 48, 42, 33, 35]]: 49\n" - ] - } - ], + "outputs": [], "source": [ "# Testing the quantized model on a sample input\n", "quantized_model.eval()\n", diff --git a/torch/easy/quantize-lm/quantize-language-model_SOLN.ipynb b/torch/easy/quantize-lm/quantize-language-model_SOLN.ipynb index 35f0ed1..4b0240c 100644 --- a/torch/easy/quantize-lm/quantize-language-model_SOLN.ipynb +++ b/torch/easy/quantize-lm/quantize-language-model_SOLN.ipynb @@ -30,14 +30,23 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": null, "metadata": {}, "outputs": [], "source": [ "import torch\n", "import torch.nn as nn\n", "import torch.optim as optim\n", - "from torch.quantization import quantize_dynamic" + "from torch.ao.quantization import quantize_dynamic\n", + "\n", + "# Dynamic quantization needs a quantized backend. On Apple Silicon (and some other\n", + "# builds) the default engine is 'none', so pick one this torch build supports.\n", + "if torch.backends.quantized.engine == \"none\":\n", + " engines = [e for e in torch.backends.quantized.supported_engines if e != \"none\"]\n", + " if not engines:\n", + " raise RuntimeError(\"this torch build has no quantized backend (fbgemm/qnnpack)\")\n", + " torch.backends.quantized.engine = engines[0]\n", + "print(f\"Quantized engine: {torch.backends.quantized.engine}\")" ] }, { @@ -64,7 +73,7 @@ }, { "cell_type": "code", - "execution_count": 5, + "execution_count": null, "metadata": {}, "outputs": [], "source": [ @@ -82,27 +91,18 @@ "num_layers = 2\n", "model = LanguageModel(vocab_size, embed_size, hidden_size, num_layers)\n", "\n", - "criterion = nn.CrossEntropyLoss()\n", + "# The model ends in a softmax, so it returns probabilities, not logits.\n", + "# nn.CrossEntropyLoss would apply a second softmax on top and the loss would barely\n", + "# move, so train with NLLLoss on the log-probabilities instead.\n", + "criterion = nn.NLLLoss()\n", "optimizer = optim.Adam(model.parameters(), lr=0.001)" ] }, { "cell_type": "code", - "execution_count": 6, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Epoch [1/5] - Loss: 3.9118\n", - "Epoch [2/5] - Loss: 3.9113\n", - "Epoch [3/5] - Loss: 3.9108\n", - "Epoch [4/5] - Loss: 3.9103\n", - "Epoch [5/5] - Loss: 3.9097\n" - ] - } - ], + "outputs": [], "source": [ "# Training loop\n", "epochs = 5\n", @@ -110,7 +110,7 @@ " model.train()\n", " optimizer.zero_grad()\n", " output = model(X_train)\n", - " loss = criterion(output, y_train)\n", + " loss = criterion(torch.log(output.clamp_min(1e-9)), y_train)\n", " loss.backward()\n", " optimizer.step()\n", "\n", @@ -121,49 +121,30 @@ "# Quantization: Apply dynamic quantization to the language model\n", "quantized_model = quantize_dynamic(model, {nn.Linear, nn.LSTM}, dtype=torch.qint8)\n", "\n", - "# Save the quantized model\n", - "torch.save(quantized_model.state_dict(), \"quantized_language_model.pth\")\n" + "# Save the float weights. A dynamically quantized state_dict stores the LSTM's packed\n", + "# weights as torch.ScriptObject, which torch.load(weights_only=True) (the default\n", + "# since PyTorch 2.6) refuses to unpickle. Dynamic quantization is deterministic, so\n", + "# quantizing the float weights again after loading rebuilds the same model.\n", + "torch.save(model.state_dict(), \"language_model_fp32.pth\")\n" ] }, { "cell_type": "code", - "execution_count": 7, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "data": { - "text/plain": [ - "" - ] - }, - "execution_count": 7, - "metadata": {}, - "output_type": "execute_result" - } - ], + "outputs": [], "source": [ - "# Load the quantized model and test it\n", + "# Load the float weights, then apply the same dynamic quantization again\n", "quantized_model = LanguageModel(vocab_size, embed_size, hidden_size, num_layers)\n", - "\n", - "# Apply dynamic quantization on the model after defining it\n", - "quantized_model = quantize_dynamic(quantized_model, {nn.Linear, nn.LSTM}, dtype=torch.qint8)\n", - "\n", - "quantized_model.load_state_dict(torch.load(\"quantized_language_model.pth\"))" + "quantized_model.load_state_dict(torch.load(\"language_model_fp32.pth\", weights_only=True))\n", + "quantized_model = quantize_dynamic(quantized_model, {nn.Linear, nn.LSTM}, dtype=torch.qint8)" ] }, { "cell_type": "code", - "execution_count": 8, + "execution_count": null, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Prediction for input [[15, 28, 33, 19, 37, 24, 48, 42, 33, 35]]: 49\n" - ] - } - ], + "outputs": [], "source": [ "# Testing the quantized model on a sample input\n", "quantized_model.eval()\n", @@ -172,6 +153,33 @@ " prediction = quantized_model(test_input)\n", " print(f\"Prediction for input {test_input.tolist()}: {prediction.argmax(dim=1).item()}\")" ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# Compare the quantized model with the float one: size of the saved weights and\n", + "# whether the two agree on the next-token prediction for the same inputs\n", + "import io\n", + "\n", + "def state_dict_mb(m):\n", + " buf = io.BytesIO()\n", + " torch.save(m.state_dict(), buf)\n", + " return buf.getbuffer().nbytes / 1e6\n", + "\n", + "model.eval()\n", + "eval_input = torch.randint(0, vocab_size, (256, seq_length))\n", + "with torch.no_grad():\n", + " fp32_probs = model(eval_input)\n", + " int8_probs = quantized_model(eval_input)\n", + "\n", + "agreement = (fp32_probs.argmax(dim=1) == int8_probs.argmax(dim=1)).float().mean().item()\n", + "print(f\"Weights: {state_dict_mb(model):.2f} MB (fp32) -> {state_dict_mb(quantized_model):.2f} MB (int8)\")\n", + "print(f\"Top-1 agreement with the float model: {agreement:.1%}\")\n", + "print(f\"Max probability difference: {(fp32_probs - int8_probs).abs().max().item():.2e}\")" + ] } ], "metadata": { diff --git a/v3/alignment-training/gradient-checkpointing/gradient-checkpointing.ipynb b/v3/alignment-training/gradient-checkpointing/gradient-checkpointing.ipynb index 26d28f3..bc10810 100644 --- a/v3/alignment-training/gradient-checkpointing/gradient-checkpointing.ipynb +++ b/v3/alignment-training/gradient-checkpointing/gradient-checkpointing.ipynb @@ -54,7 +54,8 @@ " 💡 Hint\n", "\n", " - In `CheckpointFunction.forward`, use `ctx.save_for_backward(*inputs)` to save only the inputs. Detach the inputs when running the function to prevent PyTorch from building a graph for intermediates.\n", - " - In `CheckpointFunction.backward`, re-enable gradients with `torch.enable_grad()`, rerun the function on the saved inputs, then call `torch.autograd.grad()` to compute gradients of the recomputed outputs w.r.t. the inputs.\n", + " - In `CheckpointFunction.backward`, re-enable gradients with `torch.enable_grad()`, rerun the function on detached copies of the saved inputs, then call `torch.autograd.backward(outputs, grad_outputs)`. Unlike `torch.autograd.grad()` on the inputs alone, this also fills in `.grad` for the parameters `fn` closes over; read the input gradients back off the detached copies.\n", + " - Autograd only calls `backward` if some tensor *argument* requires grad. The first block is usually fed plain data, so in `checkpoint` pass an extra dummy tensor with `requires_grad=True` when no argument requires grad, or that block's weights never train.\n", " - Save the function itself on the context: `ctx.fn = fn`.\n", " - Memory estimation: count the number of tensors stored in the autograd graph, or use `torch.cuda.memory_allocated()` on GPU.\n", "\n", @@ -165,7 +166,9 @@ " # TODO: Retrieve saved inputs and fn from ctx\n", " # TODO: Enable gradients and recompute forward pass on the saved inputs\n", " # (need to set requires_grad on inputs and use torch.enable_grad)\n", - " # TODO: Use torch.autograd.grad to compute gradients of outputs w.r.t. inputs\n", + " # TODO: Backpropagate grad_outputs through the recomputed outputs with\n", + " # torch.autograd.backward, so fn's parameters get gradients too,\n", + " # then read the input gradients off the detached copies\n", " # TODO: Return (None, *grads) — None for the fn argument\n", " ...\n", "\n", @@ -181,7 +184,10 @@ " Returns:\n", " Output of fn(*args) with gradient checkpointing applied.\n", " \"\"\"\n", - " # TODO: Call CheckpointFunction.apply with fn and args\n", + " # TODO: Call CheckpointFunction.apply with fn and args.\n", + " # If no argument requires grad (plain input data), autograd will never\n", + " # call backward, so fn's weights would never train: pass an extra dummy\n", + " # tensor with requires_grad=True in that case.\n", " ...\n", "\n", "\n", @@ -248,7 +254,12 @@ "\n", "grad_match = True\n", "for (n1, p1), (n2, p2) in zip(model_no_ckpt.named_parameters(), model_ckpt.named_parameters()):\n", - " if p1.grad is not None and p2.grad is not None:\n", + " # A missing gradient is a FAILURE, not something to skip: checkpointed\n", + " # parameters silently receiving no gradient is the classic bug here.\n", + " if p1.grad is None or p2.grad is None:\n", + " print(f\" MISSING gradient at {n1}: no_ckpt={p1.grad is not None}, ckpt={p2.grad is not None}\")\n", + " grad_match = False\n", + " else:\n", " if not torch.allclose(p1.grad, p2.grad, atol=1e-5):\n", " print(f\" Gradient mismatch at {n1}: max diff = {(p1.grad - p2.grad).abs().max():.2e}\")\n", " grad_match = False\n", @@ -275,6 +286,13 @@ "X = torch.randn(32, 64)\n", "y = torch.randn(32, 1)\n", "\n", + "# X is plain data (requires_grad=False), exactly like a batch from a DataLoader.\n", + "# Every checkpointed block must still receive gradients, not just the head.\n", + "F.mse_loss(model_train(X), y).backward()\n", + "no_grad = [n for n, p in model_train.named_parameters() if p.grad is None or not p.grad.abs().sum()]\n", + "assert not no_grad, f\"{len(no_grad)} parameters got no gradient, e.g. {no_grad[:3]}\"\n", + "print(\"Every parameter receives a gradient with plain (no-grad) input\")\n", + "\n", "losses = []\n", "for step in range(30):\n", " optimizer.zero_grad()\n", diff --git a/v3/alignment-training/gradient-checkpointing/gradient-checkpointing_SOLN.ipynb b/v3/alignment-training/gradient-checkpointing/gradient-checkpointing_SOLN.ipynb index 381d65e..fb9f85f 100644 --- a/v3/alignment-training/gradient-checkpointing/gradient-checkpointing_SOLN.ipynb +++ b/v3/alignment-training/gradient-checkpointing/gradient-checkpointing_SOLN.ipynb @@ -44,7 +44,8 @@ " 💡 Hint\n", "\n", " - In `CheckpointFunction.forward`, use `ctx.save_for_backward(*inputs)` to save only the inputs. Detach the inputs when running the function to prevent PyTorch from building a graph for intermediates.\n", - " - In `CheckpointFunction.backward`, re-enable gradients with `torch.enable_grad()`, rerun the function on the saved inputs, then call `torch.autograd.grad()` to compute gradients of the recomputed outputs w.r.t. the inputs.\n", + " - In `CheckpointFunction.backward`, re-enable gradients with `torch.enable_grad()`, rerun the function on detached copies of the saved inputs, then call `torch.autograd.backward(outputs, grad_outputs)`. Unlike `torch.autograd.grad()` on the inputs alone, this also fills in `.grad` for the parameters `fn` closes over; read the input gradients back off the detached copies.\n", + " - Autograd only calls `backward` if some tensor *argument* requires grad. The first block is usually fed plain data, so in `checkpoint` pass an extra dummy tensor with `requires_grad=True` when no argument requires grad, or that block's weights never train.\n", " - Save the function itself on the context: `ctx.fn = fn`.\n", " - Memory estimation: count the number of tensors stored in the autograd graph, or use `torch.cuda.memory_allocated()` on GPU.\n", "\n", @@ -163,7 +164,11 @@ " # leaf reached during recomputation - parameters included - and we read\n", " # the input gradients back off the detached copies.\n", " outputs = output if isinstance(output, tuple) else (output,)\n", - " torch.autograd.backward(outputs, grad_outputs)\n", + " # Skip outputs that do not depend on anything trainable (e.g. a frozen fn).\n", + " pairs = [(o, g) for o, g in zip(outputs, grad_outputs)\n", + " if isinstance(o, torch.Tensor) and o.requires_grad]\n", + " if pairs:\n", + " torch.autograd.backward([o for o, _ in pairs], [g for _, g in pairs])\n", "\n", " # Build full gradient tuple: None for fn, then grad or None for each input\n", " result = [None] # None for fn\n", @@ -177,6 +182,16 @@ " \"\"\"\n", " Apply gradient checkpointing to fn(*args).\n", " \"\"\"\n", + " # Autograd only records CheckpointFunction (and later calls its backward) if at\n", + " # least one tensor *argument* requires grad. fn's parameters are not arguments,\n", + " # so when the input is plain data (the first block, fed a batch straight from\n", + " # the dataset) backward would never run and fn's weights would never train.\n", + " # Passing a dummy tensor that requires grad keeps the checkpoint in the graph.\n", + " if torch.is_grad_enabled() and not any(\n", + " isinstance(a, torch.Tensor) and a.requires_grad for a in args\n", + " ):\n", + " dummy = torch.empty(0, requires_grad=True)\n", + " return CheckpointFunction.apply(lambda _dummy, *a: fn(*a), dummy, *args)\n", " return CheckpointFunction.apply(fn, *args)\n", "\n", "\n", @@ -275,6 +290,13 @@ "X = torch.randn(32, 64)\n", "y = torch.randn(32, 1)\n", "\n", + "# X is plain data (requires_grad=False), exactly like a batch from a DataLoader.\n", + "# Every checkpointed block must still receive gradients, not just the head.\n", + "F.mse_loss(model_train(X), y).backward()\n", + "no_grad = [n for n, p in model_train.named_parameters() if p.grad is None or not p.grad.abs().sum()]\n", + "assert not no_grad, f\"{len(no_grad)} parameters got no gradient, e.g. {no_grad[:3]}\"\n", + "print(\"Every parameter receives a gradient with plain (no-grad) input\")\n", + "\n", "losses = []\n", "for step in range(30):\n", " optimizer.zero_grad()\n", diff --git a/v3/alignment-training/gradient-checkpointing/problem.toml b/v3/alignment-training/gradient-checkpointing/problem.toml index a95209e..79f13e9 100644 --- a/v3/alignment-training/gradient-checkpointing/problem.toml +++ b/v3/alignment-training/gradient-checkpointing/problem.toml @@ -12,8 +12,9 @@ tracks = ["advanced"] # GENERATED from the grader spec — edit python/torchleet/problems/*.py, then run: generate.py hints hints = [ "In forward: stash fn on ctx, ctx.save_for_backward(*args), and run fn(*args) inside torch.no_grad() so no intermediate activation is kept.", - "In backward: detach the saved inputs, requires_grad_ them, re-run fn inside torch.enable_grad(), then torch.autograd.grad(output, inputs, grad_outputs).", - "backward must return one gradient per forward argument — None for fn itself, then a gradient (or None) for each tensor input.", + "In backward: detach the saved inputs, requires_grad_ them, re-run fn inside torch.enable_grad(), then torch.autograd.backward(outputs, grad_outputs) — unlike autograd.grad on the inputs, that also reaches fn's own parameters.", + "backward must return one gradient per forward argument — None for fn itself, then a gradient (or None) for each tensor input (read them off the detached copies).", + "Autograd only calls backward if some tensor argument requires grad. When none does (plain input data), have checkpoint() pass an extra dummy tensor with requires_grad=True, or the wrapped layer's weights never train.", ] [notebooks]