Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 6 additions & 4 deletions mcp-server/src/data/hints.generated.ts
Original file line number Diff line number Diff line change
Expand Up @@ -455,13 +455,15 @@ export const authoredHints: Record<string, string[]> = {
],
"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.",
Expand Down
5 changes: 3 additions & 2 deletions problems.json
Original file line number Diff line number Diff line change
Expand Up @@ -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."
]
},
{
Expand Down
76 changes: 76 additions & 0 deletions python/test_grader.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)),
Expand All @@ -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), {}),
]


Expand Down
63 changes: 50 additions & 13 deletions python/torchleet/problems/gradient_checkpointing.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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 = []
Expand All @@ -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
Expand Down Expand Up @@ -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)
Expand All @@ -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, \
Expand All @@ -182,11 +188,42 @@ 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,
check_gradients_match_uncheckpointed,
check_chained_checkpoints_match,
check_intermediates_are_not_kept,
check_module_parameter_gradients,
check_plain_input_still_trains_parameters,
]
27 changes: 25 additions & 2 deletions python/torchleet/problems/rotary_positional_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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)
Expand All @@ -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]
5 changes: 3 additions & 2 deletions space/problems.json
Original file line number Diff line number Diff line change
Expand Up @@ -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."
]
},
{
Expand Down
Loading
Loading