Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
test_compose_loss_integration.py398 linesDownload Raw Back to tests
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