Team Ai
Modelpublic

Codeseys/composer-replication-framework

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