From 5dc8ad74a48685046b18f89d6feae956a228ea25 Mon Sep 17 00:00:00 2001 From: Saksham Adhikari Date: Thu, 24 Sep 2026 13:44:23 -0500 Subject: [PATCH 1/3] Fix quantize-lm on macOS/torch>=2.6 and the double softmax in its training loop (#21) - Pick a quantized engine when the default is 'none' (Apple Silicon), the same guard the grader already uses. Previously quantize_dynamic raised "Didn't find engine for operation quantized::linear_prepack NoQEngine". - Save the float state_dict and re-quantize after loading. A dynamically quantized LSTM state_dict holds torch.ScriptObject packed params, which torch.load(weights_only=True) - the default since 2.6 - refuses to unpickle. Dynamic quantization is deterministic, so the reloaded model is identical. - The model ends in a softmax, so train with NLLLoss on log-probabilities instead of CrossEntropyLoss (a second softmax), which pinned the loss at ~ln(50): 3.9118 -> 3.9097 over 5 epochs, now 3.9032 -> 3.8048. - Import quantize_dynamic from torch.ao.quantization (torch.quantization is the deprecated alias). - Solution: add the float-vs-int8 comparison the problem statement asks for (weight size and top-1 agreement). Clear outputs produced by the old code. --- .../quantize-lm/quantize-language-model.ipynb | 81 +++++-------- .../quantize-language-model_SOLN.ipynb | 108 ++++++++++-------- 2 files changed, 89 insertions(+), 100 deletions(-) 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": { From d78661d658a06e1effe9b13bb70d6a2bd0f44df3 Mon Sep 17 00:00:00 2001 From: Saksham Adhikari Date: Thu, 24 Sep 2026 13:44:23 -0500 Subject: [PATCH 2/3] Fix gradient checkpointing leaving checkpointed blocks untrained autograd only records a custom Function, and later calls its backward, when at least one tensor argument requires grad. The block parameters fn closes over are not arguments, so with plain input data (the notebook's own training test, and the first layer of any real network) backward never ran: only the head trained and all 96 block parameters had grad None, yet the test passed because the head alone brought the loss down. - Solution: checkpoint() passes a dummy requires_grad tensor when no argument requires grad; backward skips outputs that need no grad. - Training test now asserts every parameter receives a gradient with plain input (the previous solution fails it). - Question: hint and TODOs no longer steer to autograd.grad w.r.t. the inputs (which never reaches fn's parameters), and the validation treats a missing gradient as a failure instead of skipping it, matching the solution. - Grader: a checkpoint whose wrapped module gets no gradient now fails instead of being reported as "not verified" (a pass); new check for plain input. Hints updated and regenerated into problem.toml, problems.json and the MCP server's hint table. --- mcp-server/src/data/hints.generated.ts | 10 +-- problems.json | 5 +- .../problems/gradient_checkpointing.py | 63 +++++++++++++++---- space/problems.json | 5 +- .../gradient-checkpointing.ipynb | 26 ++++++-- .../gradient-checkpointing_SOLN.ipynb | 26 +++++++- .../gradient-checkpointing/problem.toml | 5 +- 7 files changed, 111 insertions(+), 29 deletions(-) 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/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/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/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] From 7b12efddd70115f4e0c6d19f1b87b093f42d5b43 Mon Sep 17 00:00:00 2001 From: Saksham Adhikari Date: Thu, 24 Sep 2026 13:44:23 -0500 Subject: [PATCH 3/3] RoPE grader: reject implementations that do not rotate Every existing value check was satisfied by returning (q, k) unchanged: the zero-angle check expects exactly that, and a no-op trivially preserves norms. Add check_quarter_turn: at cos=0, sin=1 the result must equal the solver's own rotate_half(q) and rotate_half(k), which also catches forgetting to rotate k. test_grader.py gains wrong-answer cases for RoPE (no-op, query-only) and gradient checkpointing (no dummy input, autograd.grad on inputs only) so these checks cannot regress to accepting broken code. --- python/test_grader.py | 76 +++++++++++++++++++ .../problems/rotary_positional_embedding.py | 27 ++++++- 2 files changed, 101 insertions(+), 2 deletions(-) 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/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]