Codeseys/composer-replication-framework
0
1"""Gradient-flow tests for compose_loss channels (Wave 16b).2 3Wave 14-15 verified compose_loss returns correct numeric values and that4channel disables behave correctly. This file closes the gap by verifying5that gradients actually flow back through each enabled channel and reach6model parameters when the channel is on, AND that disabled channels7produce zero side-effects on the autograd graph.8 9Coverage:10 1. test_alpha_sdpo_routes_grad_to_params11 — alpha_sdpo=1.0 + SDPO inputs => non-zero finite grads on params12 2. test_beta_replay_routes_grad_to_params13 — beta_replay=1.0 + DPO inputs => non-zero finite grads on params14 3. test_alpha_zero_blocks_sdpo_grad15 — alpha_sdpo=0.0: SDPO inputs present vs absent yields BIT-IDENTICAL16 param.grad on every parameter (catches phantom-gradient leaks17 from disabled channels)18 4. test_taid_grad_flows_through_sdpo_path19 — sdpo_wrapper="taid", taid_t=0.5 still routes grads through20 the SDPO channel under autograd21 22Same TinyLM scaffold as test_compose_loss_integration.py — no HF / TRL,23all tests run in milliseconds.24"""25from __future__ import annotations26 27import math28 29import torch30import torch.nn as nn31 32from composer_replication import compose_loss33 34 35# ----------------------------------------------------------------------36# Tiny LM stand-in (mirrors test_compose_loss_integration.py)37# ----------------------------------------------------------------------38 39 40class TinyLM(nn.Module):41 """Minimal nn.Module with HF-style ``model(input_ids=...).logits`` API."""42 43 def __init__(self, vocab: int = 32, hidden: int = 16, seed: int = 0):44 super().__init__()45 torch.manual_seed(seed)46 self.embed = nn.Embedding(vocab, hidden)47 self.fc = nn.Linear(hidden, hidden)48 self.head = nn.Linear(hidden, vocab)49 50 def forward(self, input_ids: torch.Tensor):51 h = torch.tanh(self.fc(self.embed(input_ids)))52 logits = self.head(h)53 54 class _Out:55 pass56 out = _Out()57 out.logits = logits58 return out59 60 61# ----------------------------------------------------------------------62# Batch fixtures (mirror test_compose_loss_integration.py shape)63# ----------------------------------------------------------------------64 65VOCAB = 3266B = 267T = 868 69 70def _make_inputs(seed: int = 7, *, with_sdpo: bool, with_dpo: bool) -> dict:71 """Build a deterministic input batch with optional channel inputs.72 73 SDPO and DPO inputs can be independently included or excluded so we74 can exercise the channel-disable code paths cleanly.75 """76 g = torch.Generator().manual_seed(seed)77 inputs: dict[str, torch.Tensor] = {78 "input_ids": torch.randint(0, VOCAB, (B, T), generator=g),79 "response_mask": torch.zeros(B, T, dtype=torch.long),80 }81 inputs["response_mask"][:, T // 2:] = 182 83 if with_sdpo:84 inputs["ctx_teacher_input_ids"] = torch.randint(0, VOCAB, (B, T), generator=g)85 inputs["sdpo_loss_mask"] = torch.zeros(B, T, dtype=torch.long)86 inputs["sdpo_loss_mask"][:, T // 2:] = 187 88 if with_dpo:89 inputs["dpo_chosen_input_ids"] = torch.randint(0, VOCAB, (B, T), generator=g)90 inputs["dpo_chosen_response_mask"] = torch.ones(B, T, dtype=torch.long)91 inputs["dpo_rejected_input_ids"] = torch.randint(0, VOCAB, (B, T), generator=g)92 inputs["dpo_rejected_response_mask"] = torch.ones(B, T, dtype=torch.long)93 inputs["dpo_chosen_ref_logprobs"] = torch.randn(B, generator=g)94 inputs["dpo_rejected_ref_logprobs"] = torch.randn(B, generator=g)95 96 return inputs97 98 99def _grad_norm(model: nn.Module) -> float:100 """Sum of |grad| across all params with non-None grad."""101 return sum(102 p.grad.detach().abs().sum().item()103 for p in model.parameters()104 if p.grad is not None105 )106 107 108def _grad_is_finite(model: nn.Module) -> bool:109 """All param grads are finite (no inf, no nan)."""110 for p in model.parameters():111 if p.grad is None:112 continue113 if not torch.isfinite(p.grad).all():114 return False115 return True116 117 118def _model() -> TinyLM:119 """Fresh TinyLM with deterministic init."""120 return TinyLM(vocab=VOCAB, hidden=16, seed=0)121 122 123# ----------------------------------------------------------------------124# Test 1 — SDPO channel routes grads to params when alpha_sdpo > 0125# ----------------------------------------------------------------------126 127 128def test_alpha_sdpo_routes_grad_to_params():129 """When alpha_sdpo > 0 and SDPO inputs are present, calling130 out.total.backward() must produce non-zero finite gradients on131 model parameters.132 """133 model = _model()134 inputs = _make_inputs(with_sdpo=True, with_dpo=False)135 136 out = compose_loss(137 model,138 inputs,139 alpha_sdpo=1.0,140 beta_replay=0.0,141 )142 143 # Sanity: SDPO actually fired (channel is non-zero).144 assert float(out.sdpo_jsd) != 0.0, (145 "alpha_sdpo=1.0 with SDPO inputs should produce a non-zero sdpo_jsd; "146 f"got {float(out.sdpo_jsd)}"147 )148 149 out.total.backward()150 g = _grad_norm(model)151 assert g > 0.0, f"Expected non-zero grad sum from SDPO channel; got {g}"152 assert math.isfinite(g), f"Grad sum is not finite: {g}"153 assert _grad_is_finite(model), "Some grads are inf/nan"154 155 156# ----------------------------------------------------------------------157# Test 2 — Replay-DPO channel routes grads to params when beta_replay > 0158# ----------------------------------------------------------------------159 160 161def test_beta_replay_routes_grad_to_params():162 """When beta_replay > 0 and DPO inputs are present, backward must163 produce non-zero finite gradients on model parameters.164 165 Note: response_mask is set to all-zeros so the LM-CE channel is166 exactly zero — any non-zero grad must come from the DPO channel.167 """168 model = _model()169 inputs = _make_inputs(with_sdpo=False, with_dpo=True)170 # Zero out response_mask so LM-CE contributes nothing — isolates DPO.171 inputs["response_mask"] = torch.zeros(B, T, dtype=torch.long)172 173 out = compose_loss(174 model,175 inputs,176 alpha_sdpo=0.0,177 beta_replay=1.0,178 )179 180 assert float(out.lm_ce) == 0.0, "LM-CE should be zero with empty response_mask"181 assert float(out.trace_replay_dpo) != 0.0, (182 "beta_replay=1.0 with DPO inputs should produce a non-zero "183 f"trace_replay_dpo; got {float(out.trace_replay_dpo)}"184 )185 186 out.total.backward()187 g = _grad_norm(model)188 assert g > 0.0, f"Expected non-zero grad sum from DPO channel; got {g}"189 assert math.isfinite(g), f"Grad sum is not finite: {g}"190 assert _grad_is_finite(model), "Some grads are inf/nan"191 192 193# ----------------------------------------------------------------------194# Test 3 — Disabled SDPO channel produces ZERO side-effects on autograd195# ----------------------------------------------------------------------196 197 198def test_alpha_zero_blocks_sdpo_grad():199 """With alpha_sdpo=0.0, providing SDPO inputs vs omitting them must200 produce bit-identical parameter gradients.201 202 This catches a class of bug where a disabled channel leaks a phantom203 contribution into the autograd graph (e.g. if the SDPO branch ran a204 forward pass even when alpha=0 and somehow scaled the result by205 alpha=0 incorrectly).206 """207 inputs_with_sdpo = _make_inputs(with_sdpo=True, with_dpo=False)208 inputs_no_sdpo = _make_inputs(with_sdpo=False, with_dpo=False)209 210 # Trial A: SDPO inputs present, alpha=0 — channel should be silent.211 model_a = _model()212 out_a = compose_loss(model_a, inputs_with_sdpo, alpha_sdpo=0.0, beta_replay=0.0)213 out_a.total.backward()214 grads_a = {215 name: p.grad.detach().clone() if p.grad is not None else None216 for name, p in model_a.named_parameters()217 }218 219 # Trial B: SDPO inputs absent, alpha=0.220 model_b = _model() # Same seed -> bit-identical init.221 out_b = compose_loss(model_b, inputs_no_sdpo, alpha_sdpo=0.0, beta_replay=0.0)222 out_b.total.backward()223 grads_b = {224 name: p.grad.detach().clone() if p.grad is not None else None225 for name, p in model_b.named_parameters()226 }227 228 # Bit-identical grads on every parameter.229 assert set(grads_a.keys()) == set(grads_b.keys())230 for name in grads_a:231 ga, gb = grads_a[name], grads_b[name]232 if ga is None and gb is None:233 continue234 assert ga is not None and gb is not None, (235 f"Param {name}: grad_a={ga is not None}, grad_b={gb is not None}"236 )237 # atol=0, rtol=0 -> bit-exact equality. SDPO inputs with alpha=0238 # must not perturb the autograd graph by even one ULP.239 assert torch.equal(ga, gb), (240 f"Param {name}: disabled SDPO channel leaked phantom gradient. "241 f"|diff|.max()={float((ga - gb).abs().max())}"242 )243 244 245# ----------------------------------------------------------------------246# Test 4 — TAID-wrapped SDPO channel still routes grads under autograd247# ----------------------------------------------------------------------248 249 250def test_taid_grad_flows_through_sdpo_path():251 """The Wave 15 TAID rewrite (logit-space mix, current-student-detached252 anchor) must remain differentiable. With sdpo_wrapper='taid' and253 taid_t=0.5, backward must produce non-zero finite gradients on254 model parameters.255 """256 model = _model()257 inputs = _make_inputs(with_sdpo=True, with_dpo=False)258 259 out = compose_loss(260 model,261 inputs,262 alpha_sdpo=1.0,263 beta_replay=0.0,264 sdpo_wrapper="taid",265 taid_t=0.5,266 )267 268 assert float(out.sdpo_jsd) != 0.0, (269 f"taid_t=0.5 should still produce a non-zero sdpo_jsd; "270 f"got {float(out.sdpo_jsd)}"271 )272 273 out.total.backward()274 g = _grad_norm(model)275 assert g > 0.0, (276 f"Expected non-zero grad sum from TAID-wrapped SDPO channel; got {g}"277 )278 assert math.isfinite(g), f"Grad sum is not finite: {g}"279 assert _grad_is_finite(model), "Some grads are inf/nan"280 281 282# ----------------------------------------------------------------------283# Test 5 — Both channels enabled simultaneously route grads correctly284# (Wave 18 — closes the implicit-additivity gap from Wave 16's coverage)285# ----------------------------------------------------------------------286 287 288def test_both_channels_enabled_route_grad_to_params():289 """When alpha_sdpo > 0 AND beta_replay > 0 simultaneously, both channels290 must contribute to the gradient.291 292 Wave 16's tests covered each channel in isolation. This pins the293 additivity property at the gradient-norm level: with both channels294 enabled the gradient norm should be at least comparable to (and295 typically larger than) either channel alone.296 """297 inputs = _make_inputs(with_sdpo=True, with_dpo=True)298 299 def grads_and_norm(alpha, beta):300 m = _model() # seed=0 — same init every call301 out = compose_loss(m, inputs, alpha_sdpo=alpha, beta_replay=beta)302 out.total.backward()303 return _grad_norm(m)304 305 g_sdpo_only = grads_and_norm(alpha=1.0, beta=0.0)306 g_dpo_only = grads_and_norm(alpha=0.0, beta=1.0)307 g_both = grads_and_norm(alpha=1.0, beta=1.0)308 309 assert g_both > 0.0, f"Both-channels grad sum is zero: {g_both}"310 assert math.isfinite(g_both), f"Both-channels grad sum is not finite: {g_both}"311 # Smoke property: enabling both channels produces a finite, non-zero312 # gradient. We deliberately do NOT assert any lower bound relative to313 # individual-channel norms — there's no mathematical floor on the314 # composed gradient (the channels operate on different inputs and315 # their gradients can cancel arbitrarily on shared parameters). The316 # additivity property of autograd holds at the per-tensor level317 # (∂(αL1 + βL2)/∂θ = α∂L1/∂θ + β∂L2/∂θ exactly) but L1 norms of318 # vector sums need not be ≥ either summand's L1 norm.319 #320 # The companion test below verifies the per-channel grad-flow321 # property: alpha=1,beta=0 routes grad through SDPO; alpha=0,beta=1322 # routes grad through DPO. Both being non-zero in isolation + this323 # test's assertion that they jointly produce finite non-zero grads324 # is sufficient to pin "both channels contribute" without overclaiming.325 # Compute the single-channel norms purely as diagnostic context for326 # debugging when this test fails (no assertion uses them).327 _diagnostic = (g_sdpo_only, g_dpo_only) # noqa: F841 — kept for debug328 329 330# ----------------------------------------------------------------------331# Test 6 — entropy_opd wrapper routes grads through SDPO path332# (Wave 18 — Wave 15 added entropy_aware_opd_loss without an autograd test)333# ----------------------------------------------------------------------334 335 336def test_entropy_opd_grad_flows_through_sdpo_path():337 """sdpo_wrapper='entropy_opd' must remain differentiable.338 339 Wave 15 plumbed entropy_aware_opd_loss through compose_loss's340 sdpo_wrapper switch. Wave 16 tested the 'taid' wrapper under autograd341 but didn't exercise 'entropy_opd'. This test pins the entropy_opd342 path is differentiable end-to-end.343 """344 model = _model()345 inputs = _make_inputs(with_sdpo=True, with_dpo=False)346 347 out = compose_loss(348 model,349 inputs,350 alpha_sdpo=1.0,351 beta_replay=0.0,352 sdpo_wrapper="entropy_opd",353 )354 355 assert math.isfinite(float(out.sdpo_jsd)), (356 f"entropy_opd produced non-finite sdpo_jsd: {float(out.sdpo_jsd)}"357 )358 359 out.total.backward()360 g = _grad_norm(model)361 assert g > 0.0, (362 f"Expected non-zero grad sum from entropy_opd-wrapped SDPO; got {g}"363 )364 assert math.isfinite(g), f"Grad sum is not finite: {g}"365 assert _grad_is_finite(model), "Some grads are inf/nan"366 