Codeseys/composer-replication-framework
0
1"""Integration tests for ADR-007 distillation kwargs in compose_loss.2 3These tests exercise the wiring between `compose_loss` and the three4pluggable losses (SimPO, TAID, Entropy-Aware OPD). They use a tiny5hand-rolled language model wrapper (no HF, no TRL) so the tests run6in <1s on CPU and are isolated from external library churn.7 8Coverage requirements:9 (a) defaults reproduce existing compose_loss output bit-exact10 (b) dpo_variant='simpo' produces a different total than dpo11 (c) sdpo_wrapper='taid' with t=0 differs from t=1 (interpolation works)12 (d) sdpo_wrapper='taid' with t=1 reproduces upstream forward-KL13 (e) sdpo_wrapper='entropy_opd' returns a finite differentiable scalar14 (f) error case: sdpo_wrapper='taid' without taid_t raises ValueError15"""16from __future__ import annotations17 18import pytest19import torch20import torch.nn as nn21 22from composer_replication import LossComponents, compose_loss23 24 25# ----------------------------------------------------------------------26# Tiny LM stand-in27# ----------------------------------------------------------------------28 29class TinyLM(nn.Module):30 """Minimal `nn.Module` with the HF-style `model(input_ids=...).logits` API.31 32 Vocab=32, hidden=16, two-layer MLP head. Tiny enough that all tests33 run in milliseconds on CPU.34 """35 36 def __init__(self, vocab: int = 32, hidden: int = 16, seed: int = 0):37 super().__init__()38 torch.manual_seed(seed)39 self.embed = nn.Embedding(vocab, hidden)40 self.fc = nn.Linear(hidden, hidden)41 self.head = nn.Linear(hidden, vocab)42 43 def forward(self, input_ids: torch.Tensor):44 h = torch.tanh(self.fc(self.embed(input_ids)))45 logits = self.head(h)46 47 class _Out:48 pass49 out = _Out()50 out.logits = logits51 return out52 53 54# ----------------------------------------------------------------------55# Batch fixtures56# ----------------------------------------------------------------------57 58VOCAB = 3259B = 260T = 861 62 63def _base_batch(seed: int = 7, *, with_dpo: bool = True) -> dict[str, torch.Tensor]:64 """Build a deterministic input batch with all 3 channels populated."""65 g = torch.Generator().manual_seed(seed)66 inputs: dict[str, torch.Tensor] = {67 "input_ids": torch.randint(0, VOCAB, (B, T), generator=g),68 "response_mask": torch.zeros(B, T, dtype=torch.long),69 "ctx_teacher_input_ids": torch.randint(0, VOCAB, (B, T), generator=g),70 "sdpo_loss_mask": torch.zeros(B, T, dtype=torch.long),71 }72 # Mark the second half as response tokens so the LM-CE channel is non-trivial.73 inputs["response_mask"][:, T // 2:] = 174 inputs["sdpo_loss_mask"][:, T // 2:] = 175 76 if with_dpo:77 inputs["dpo_chosen_input_ids"] = torch.randint(0, VOCAB, (B, T), generator=g)78 inputs["dpo_chosen_response_mask"] = torch.ones(B, T, dtype=torch.long)79 inputs["dpo_rejected_input_ids"] = torch.randint(0, VOCAB, (B, T), generator=g)80 inputs["dpo_rejected_response_mask"] = torch.ones(B, T, dtype=torch.long)81 # Standard DPO needs ref logprobs; SimPO ignores them.82 inputs["dpo_chosen_ref_logprobs"] = torch.randn(B, generator=g)83 inputs["dpo_rejected_ref_logprobs"] = torch.randn(B, generator=g)84 return inputs85 86 87def _model_seeded(seed: int = 0) -> TinyLM:88 m = TinyLM(vocab=VOCAB, hidden=16, seed=seed)89 m.eval() # Deterministic forward — no dropout.90 return m91 92 93# ----------------------------------------------------------------------94# (a) Defaults reproduce existing output bit-exact95# ----------------------------------------------------------------------96 97def test_defaults_bit_exact_with_legacy_kwargs():98 """Calling compose_loss with new kwargs at their defaults must equal99 calling it with only the legacy kwargs. Bit-exact: every channel +100 total agree to 0 ULPs because the code path is identical.101 """102 inputs = _base_batch()103 104 model_a = _model_seeded(seed=0)105 out_legacy = compose_loss(106 model_a,107 inputs,108 alpha_sdpo=0.1,109 beta_replay=0.05,110 sdpo_jsd_beta=0.5,111 sdpo_temperature=1.0,112 replay_dpo_beta=0.1,113 )114 115 model_b = _model_seeded(seed=0)116 out_new = compose_loss(117 model_b,118 inputs,119 alpha_sdpo=0.1,120 beta_replay=0.05,121 sdpo_jsd_beta=0.5,122 sdpo_temperature=1.0,123 replay_dpo_beta=0.1,124 dpo_variant="dpo",125 sdpo_wrapper="none",126 )127 128 assert isinstance(out_new, LossComponents)129 assert torch.equal(out_legacy.lm_ce, out_new.lm_ce)130 assert torch.equal(out_legacy.sdpo_jsd, out_new.sdpo_jsd)131 assert torch.equal(out_legacy.trace_replay_dpo, out_new.trace_replay_dpo)132 assert torch.equal(out_legacy.total, out_new.total)133 134 135# ----------------------------------------------------------------------136# (b) dpo_variant='simpo' produces a different total than dpo137# ----------------------------------------------------------------------138 139def test_simpo_variant_changes_total():140 """SimPO uses average-logprob and drops the reference subtraction, so141 it must produce a different (and finite) trace_replay_dpo + total."""142 inputs = _base_batch()143 144 model_a = _model_seeded(seed=0)145 out_dpo = compose_loss(146 model_a, inputs,147 alpha_sdpo=0.0, # isolate channel 3148 beta_replay=0.05,149 dpo_variant="dpo",150 )151 152 model_b = _model_seeded(seed=0)153 out_simpo = compose_loss(154 model_b, inputs,155 alpha_sdpo=0.0,156 beta_replay=0.05,157 dpo_variant="simpo",158 )159 160 assert torch.isfinite(out_simpo.total)161 assert torch.isfinite(out_simpo.trace_replay_dpo)162 # Different formulae => different values.163 assert not torch.allclose(164 out_dpo.trace_replay_dpo, out_simpo.trace_replay_dpo165 )166 assert not torch.allclose(out_dpo.total, out_simpo.total)167 # Gradient flow check.168 out_simpo.total.backward()169 assert any(170 p.grad is not None and torch.isfinite(p.grad).all()171 for p in model_b.parameters()172 )173 174 175def test_simpo_does_not_require_ref_logprobs():176 """SimPO is reference-free; compose_loss should run when those keys are177 absent from `inputs` (only when dpo_variant='simpo')."""178 inputs = _base_batch()179 inputs.pop("dpo_chosen_ref_logprobs")180 inputs.pop("dpo_rejected_ref_logprobs")181 182 model = _model_seeded(seed=0)183 out = compose_loss(184 model, inputs,185 alpha_sdpo=0.0,186 beta_replay=0.05,187 dpo_variant="simpo",188 )189 assert torch.isfinite(out.total)190 assert torch.isfinite(out.trace_replay_dpo)191 192 193# ----------------------------------------------------------------------194# (c) TAID with t=1 reproduces upstream forward-KL on the masked tokens195# ----------------------------------------------------------------------196 197def test_taid_t_one_matches_upstream_forward_kl():198 """At t=1, taid_loss reduces to forward-KL with target = softmax(teacher).199 compose_loss should plumb through to that exact value (modulo the200 sdpo_loss_mask token-mean denominator).201 """202 import torch.nn.functional as F203 204 inputs = _base_batch(with_dpo=False)205 206 model = _model_seeded(seed=1)207 208 # Run compose_loss with TAID at t=1.209 out_taid = compose_loss(210 model, inputs,211 alpha_sdpo=1.0, # so out.sdpo_jsd is added straight to total212 beta_replay=0.0,213 sdpo_wrapper="taid",214 taid_t=1.0,215 )216 217 # Manually compute the same forward-KL on the masked tokens.218 student_logits = model(input_ids=inputs["input_ids"]).logits219 with torch.no_grad():220 teacher_logits = model(input_ids=inputs["ctx_teacher_input_ids"]).logits221 mask = inputs["sdpo_loss_mask"].float()222 p_teacher = F.softmax(teacher_logits, dim=-1, dtype=torch.float32)223 log_q = F.log_softmax(student_logits, dim=-1, dtype=torch.float32)224 per_token = -(p_teacher * log_q).sum(dim=-1)225 flat = per_token.reshape(-1)226 fmask = mask.reshape(-1).to(flat.dtype)227 expected = (flat * fmask).sum() / fmask.sum().clamp_min(1.0)228 229 # Bit-exact assertion. The TAID-loss path at t=1 is mathematically230 # identical to the manual `-(p_teacher * log_q).sum(...)` cross-entropy231 # below: at t=1, TAID's logit-space mix collapses to `teacher_logits`,232 # `softmax(teacher_logits)` is computed bit-identically inside233 # `taid_loss`, and the masked-mean reduction matches. So `torch.equal`234 # succeeds — and asserting `equal` rather than `allclose` catches any235 # future refactor that re-introduces a softmax→log roundtrip with236 # ULP drift.237 #238 # If a future change forces a roundtrip we cannot eliminate, drop to239 # `torch.testing.assert_close(out_taid.sdpo_jsd, expected,240 # atol=1e-7, rtol=0)` — that is the strict-but-feasible bound for241 # softmax→log→softmax in float32 (one ULP at the scale of the loss,242 # ~3.5e-7 here, dominated by the log_softmax LSE accumulation).243 assert torch.equal(out_taid.sdpo_jsd, expected), (244 f"TAID t=1 must equal upstream forward-KL bit-exact; "245 f"got out={out_taid.sdpo_jsd.item()!r}, "246 f"expected={expected.item()!r}, "247 f"diff={(out_taid.sdpo_jsd - expected).abs().item():.3e}"248 )249 250 251# ----------------------------------------------------------------------252# (d) TAID interpolates: t=0 differs from t=1253# ----------------------------------------------------------------------254 255def test_taid_interpolates_with_t():256 """Different t values give different sdpo_jsd. Differentiable end-to-end."""257 inputs = _base_batch(with_dpo=False)258 259 model_zero = _model_seeded(seed=2)260 out_zero = compose_loss(261 model_zero, inputs,262 alpha_sdpo=0.1, beta_replay=0.0,263 sdpo_wrapper="taid",264 taid_t=0.0,265 )266 267 model_mid = _model_seeded(seed=2)268 out_mid = compose_loss(269 model_mid, inputs,270 alpha_sdpo=0.1, beta_replay=0.0,271 sdpo_wrapper="taid",272 taid_t=0.5,273 )274 275 model_one = _model_seeded(seed=2)276 out_one = compose_loss(277 model_one, inputs,278 alpha_sdpo=0.1, beta_replay=0.0,279 sdpo_wrapper="taid",280 taid_t=1.0,281 )282 283 for out in (out_zero, out_mid, out_one):284 assert torch.isfinite(out.total)285 assert torch.isfinite(out.sdpo_jsd)286 287 assert not torch.allclose(out_zero.sdpo_jsd, out_one.sdpo_jsd, atol=1e-5)288 assert not torch.allclose(out_mid.sdpo_jsd, out_one.sdpo_jsd, atol=1e-5)289 290 out_mid.total.backward()291 assert any(292 p.grad is not None and torch.isfinite(p.grad).all()293 for p in model_mid.parameters()294 )295 296 297# ----------------------------------------------------------------------298# (e) Entropy-Aware OPD returns a finite differentiable scalar299# ----------------------------------------------------------------------300 301def test_entropy_opd_returns_finite_differentiable_scalar():302 inputs = _base_batch(with_dpo=False)303 304 model = _model_seeded(seed=3)305 out = compose_loss(306 model, inputs,307 alpha_sdpo=0.1,308 beta_replay=0.0,309 sdpo_wrapper="entropy_opd",310 )311 312 assert isinstance(out, LossComponents)313 assert out.total.shape == ()314 assert torch.isfinite(out.total)315 assert torch.isfinite(out.sdpo_jsd)316 assert out.total.requires_grad317 318 out.total.backward()319 grads = [p.grad for p in model.parameters() if p.grad is not None]320 assert len(grads) > 0321 assert all(torch.isfinite(g).all() for g in grads)322 323 324# ----------------------------------------------------------------------325# (f) Error: sdpo_wrapper='taid' without taid_t326# ----------------------------------------------------------------------327 328def test_taid_requires_t():329 inputs = _base_batch(with_dpo=False)330 model = _model_seeded(seed=4)331 with pytest.raises(ValueError, match="taid_t"):332 compose_loss(333 model, inputs,334 alpha_sdpo=0.1, beta_replay=0.0,335 sdpo_wrapper="taid",336 # taid_t omitted on purpose337 )338 339 340def test_taid_t_out_of_range_raises():341 inputs = _base_batch(with_dpo=False)342 model = _model_seeded(seed=4)343 with pytest.raises(ValueError, match=r"taid_t must be in \[0, 1\]"):344 compose_loss(345 model, inputs,346 alpha_sdpo=0.1, beta_replay=0.0,347 sdpo_wrapper="taid",348 taid_t=1.5,349 )350 351 352def test_invalid_dpo_variant_raises():353 inputs = _base_batch()354 model = _model_seeded(seed=5)355 with pytest.raises(ValueError, match="dpo_variant"):356 compose_loss(357 model, inputs,358 dpo_variant="bogus", # type: ignore[arg-type]359 )360 361 362def test_invalid_sdpo_wrapper_raises():363 inputs = _base_batch()364 model = _model_seeded(seed=5)365 with pytest.raises(ValueError, match="sdpo_wrapper"):366 compose_loss(367 model, inputs,368 sdpo_wrapper="bogus", # type: ignore[arg-type]369 )370 371 372# ----------------------------------------------------------------------373# Bonus: TAIDScheduler integration374# ----------------------------------------------------------------------375 376def test_taid_compose_with_scheduler():377 """End-to-end: TAIDScheduler drives taid_t into compose_loss."""378 from composer_replication.distillation import TAIDScheduler379 380 inputs = _base_batch(with_dpo=False)381 model = _model_seeded(seed=6)382 sched = TAIDScheduler(num_train_steps=100, t_start=0.4)383 384 for step in range(3):385 out = compose_loss(386 model, inputs,387 alpha_sdpo=0.1, beta_replay=0.0,388 sdpo_wrapper="taid",389 taid_t=sched.t,390 )391 assert torch.isfinite(out.total)392 sched.update_t(out.sdpo_jsd.detach(), global_step=step)393 394 # t may have advanced past t_start after some steps (or stayed the same395 # given small num_train_steps and only 3 iters; just check it's still396 # in-range).397 assert 0.4 <= sched.t <= 1.0398 