"""test_loss_composition_smoke.py — end-to-end gradient step on a tiny model. Verifies the integration architecture's central claim — *all three channels can run simultaneously, ablate cleanly via α/β weights, and produce finite gradients on a real model* — without depending on TRL/VeRL being installed. We use a tiny custom nn.Module (a 2-layer MLP language head wrapper around an embedding) instead of `GRPOTrainer` because: 1. TRL's GRPOTrainer requires a full distributed setup (Accelerate, vLLM, real model) that's overkill for a wiring smoke test. 2. The integration claim is about LOSS COMPOSITION, not the GRPO inner loop. We can verify channel 2 (SDPO) and channel 3 (DPO) compose correctly with a stand-in channel 1 (a placeholder GRPO loss that's just `-log_prob.mean()`). What this test guarantees: - α=0, β=0 reduces to placeholder GRPO loss exactly - α=1, β=0 adds SDPO with correct gradient flow - α=0, β=1 adds DPO with correct gradient flow - α=1, β=1 sums all three; gradient is finite - No NaN/Inf in gradients across 5 sequential gradient steps - The optimizer can decrease the loss when α/β are set non-zero (i.e., the auxiliary terms aren't degenerate) Run: pytest spikes/005-integrated-trainer-skeleton/tests/test_loss_composition_smoke.py -v """ from __future__ import annotations import sys from pathlib import Path import pytest import torch import torch.nn as nn import torch.nn.functional as F sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from opsd_loss import generalized_jsd_loss # noqa: E402 # ---------------------------------------------------------------------------- # Tiny stand-in language model (~10K params) # ---------------------------------------------------------------------------- class TinyLM(nn.Module): """Two-layer MLP that takes input_ids -> logits over vocab. Vocab is intentionally tiny (V=64) so per-step compute is microseconds. """ def __init__(self, vocab_size: int = 64, hidden: int = 32) -> None: super().__init__() self.emb = nn.Embedding(vocab_size, hidden) self.fc1 = nn.Linear(hidden, hidden) self.fc2 = nn.Linear(hidden, vocab_size) def forward(self, input_ids: torch.Tensor) -> torch.Tensor: h = self.emb(input_ids) h = torch.relu(self.fc1(h)) return self.fc2(h) # ---------------------------------------------------------------------------- # Loss composition under test (mirror of ComposerReplicationTrainer logic) # ---------------------------------------------------------------------------- def placeholder_grpo_loss(logits: torch.Tensor, targets: torch.Tensor) -> torch.Tensor: """Stand-in for the parent GRPOTrainer's loss. Real GRPO depends on rollouts, group baselines, and reward shaping — none of which we have without TRL. As a stand-in we use a simple cross-entropy over a synthetic target sequence. The only property we need from this function is "differentiable scalar that reflects model quality" — that's enough to test loss composition. """ B, T, V = logits.shape return F.cross_entropy( logits.reshape(B * T, V), targets.reshape(B * T), ignore_index=-100, ) def composer_total_loss( model: nn.Module, inputs: dict[str, torch.Tensor], *, alpha_sdpo: float, beta_replay: float, ) -> dict[str, torch.Tensor]: """Mirror of ComposerReplicationTrainer._compute_loss for testing. Returns dict of (grpo, sdpo, dpo, total) so individual channels can be inspected. """ logits = model(inputs["input_ids"]) grpo_loss = placeholder_grpo_loss(logits, inputs["targets"]) # Channel 2: SDPO if alpha_sdpo > 0 and "ctx_teacher_input_ids" in inputs: student_logits = logits # student already computed above with torch.no_grad(): teacher_logits = model(inputs["ctx_teacher_input_ids"]) # Pad/truncate to align if shapes differ — should match in real use T = min(student_logits.shape[1], teacher_logits.shape[1]) sdpo_loss = generalized_jsd_loss( student_logits=student_logits[:, :T, :], teacher_logits=teacher_logits[:, :T, :], labels=inputs["sdpo_loss_mask"][:, :T] if "sdpo_loss_mask" in inputs else None, beta=0.5, ) else: sdpo_loss = torch.tensor(0.0, device=logits.device) # Channel 3: trace-replay DPO if beta_replay > 0 and "dpo_chosen_input_ids" in inputs: chosen_lp = _seq_logprob(model, inputs["dpo_chosen_input_ids"], inputs["dpo_chosen_response_mask"]) rejected_lp = _seq_logprob(model, inputs["dpo_rejected_input_ids"], inputs["dpo_rejected_response_mask"]) ref_chosen_lp = inputs["dpo_chosen_ref_logprobs"] ref_rejected_lp = inputs["dpo_rejected_ref_logprobs"] beta_dpo = 0.1 dpo_logits = beta_dpo * ( (chosen_lp - ref_chosen_lp) - (rejected_lp - ref_rejected_lp) ) dpo_loss = -F.logsigmoid(dpo_logits).mean() else: dpo_loss = torch.tensor(0.0, device=logits.device) total = grpo_loss + alpha_sdpo * sdpo_loss + beta_replay * dpo_loss return {"grpo": grpo_loss, "sdpo": sdpo_loss, "dpo": dpo_loss, "total": total} def _seq_logprob(model: nn.Module, input_ids: torch.Tensor, response_mask: torch.Tensor) -> torch.Tensor: logits = model(input_ids) log_probs = F.log_softmax(logits[:, :-1, :], dim=-1) targets = input_ids[:, 1:] token_lp = log_probs.gather(-1, targets.unsqueeze(-1)).squeeze(-1) masked = token_lp * response_mask[:, 1:].float() return masked.sum(dim=-1) # ---------------------------------------------------------------------------- # Fixtures: synthetic batch with all three channels populated # ---------------------------------------------------------------------------- @pytest.fixture def model(): torch.manual_seed(42) return TinyLM(vocab_size=64, hidden=32) @pytest.fixture def batch(): """Synthetic batch with all three channels: input_ids, ctx_teacher_input_ids, dpo pairs.""" torch.manual_seed(0) B, T = 2, 8 return { "input_ids": torch.randint(1, 64, (B, T)), "targets": torch.randint(0, 64, (B, T)), "ctx_teacher_input_ids": torch.randint(1, 64, (B, T)), "sdpo_loss_mask": torch.tensor([[1, 1, -100, -100, -100, -100, -100, -100], [-100, 1, 1, -100, -100, -100, -100, -100]]), "dpo_chosen_input_ids": torch.randint(1, 64, (B, T)), "dpo_chosen_response_mask": torch.tensor([[0, 0, 0, 1, 1, 1, 1, 1]] * B), "dpo_rejected_input_ids": torch.randint(1, 64, (B, T)), "dpo_rejected_response_mask": torch.tensor([[0, 0, 0, 1, 1, 1, 1, 1]] * B), "dpo_chosen_ref_logprobs": torch.randn(B), "dpo_rejected_ref_logprobs": torch.randn(B), } # ---------------------------------------------------------------------------- # Tests # ---------------------------------------------------------------------------- def test_alpha0_beta0_equals_grpo_only(model, batch): """With α=0, β=0, total_loss must equal grpo_loss exactly.""" out = composer_total_loss(model, batch, alpha_sdpo=0.0, beta_replay=0.0) assert torch.isclose(out["total"], out["grpo"]), \ f"Expected total == grpo with α=β=0, got total={out['total']}, grpo={out['grpo']}" def test_alpha_only_adds_sdpo(model, batch): """With α=1, β=0, total_loss = grpo + sdpo (and sdpo > 0).""" out = composer_total_loss(model, batch, alpha_sdpo=1.0, beta_replay=0.0) assert out["sdpo"].item() > 0, "SDPO loss should be positive on random init" expected = out["grpo"] + out["sdpo"] assert torch.isclose(out["total"], expected, atol=1e-5) def test_beta_only_adds_dpo(model, batch): """With α=0, β=1, total_loss = grpo + dpo.""" out = composer_total_loss(model, batch, alpha_sdpo=0.0, beta_replay=1.0) assert torch.isfinite(out["dpo"]), "DPO loss must be finite" expected = out["grpo"] + out["dpo"] assert torch.isclose(out["total"], expected, atol=1e-5) def test_full_composition_is_sum(model, batch): """All three channels active: total = grpo + α·sdpo + β·dpo.""" out = composer_total_loss(model, batch, alpha_sdpo=0.5, beta_replay=0.3) expected = out["grpo"] + 0.5 * out["sdpo"] + 0.3 * out["dpo"] assert torch.isclose(out["total"], expected, atol=1e-5) def test_all_channels_produce_finite_gradients(model, batch): """Backprop succeeds, no NaN/Inf in any model parameter's gradient.""" out = composer_total_loss(model, batch, alpha_sdpo=0.5, beta_replay=0.3) out["total"].backward() for name, param in model.named_parameters(): assert param.grad is not None, f"{name} got no gradient" assert torch.isfinite(param.grad).all(), \ f"{name} has NaN/Inf in grad: max={param.grad.abs().max()}" def test_5_step_train_decreases_loss(): """Run 5 gradient steps with all 3 channels; total loss should monotonically or near-monotonically decrease — channels are not actively fighting each other.""" torch.manual_seed(7) model = TinyLM(vocab_size=64, hidden=32) optimizer = torch.optim.Adam(model.parameters(), lr=1e-2) # Build a fixed batch we'll re-use across steps (overfitting check) B, T = 2, 8 fixed_batch = { "input_ids": torch.randint(1, 64, (B, T)), "targets": torch.randint(0, 64, (B, T)), "ctx_teacher_input_ids": torch.randint(1, 64, (B, T)), "sdpo_loss_mask": torch.tensor([[1, 1, -100, -100, -100, -100, -100, -100]] * B), "dpo_chosen_input_ids": torch.randint(1, 64, (B, T)), "dpo_chosen_response_mask": torch.tensor([[0, 0, 0, 1, 1, 1, 1, 1]] * B), "dpo_rejected_input_ids": torch.randint(1, 64, (B, T)), "dpo_rejected_response_mask": torch.tensor([[0, 0, 0, 1, 1, 1, 1, 1]] * B), "dpo_chosen_ref_logprobs": torch.randn(B), "dpo_rejected_ref_logprobs": torch.randn(B), } losses: list[float] = [] for _step in range(5): optimizer.zero_grad() out = composer_total_loss(model, fixed_batch, alpha_sdpo=0.1, beta_replay=0.05) out["total"].backward() optimizer.step() losses.append(out["total"].item()) # No NaN at any step assert torch.isfinite(out["total"]), f"Loss is NaN/Inf at step {_step}" # Loss at step 4 should be lower than at step 0 (overfitting check) assert losses[-1] < losses[0], \ f"Loss did not decrease over 5 steps: {[round(l, 4) for l in losses]}" def test_sdpo_only_run_reduces_to_grpo_when_no_error_sites(): """Sanity check: even with α=1, if the data collator emits no SDPO fields (no error sites), the loss still reduces to GRPO-only.""" torch.manual_seed(1) model = TinyLM(vocab_size=64, hidden=32) B, T = 2, 4 batch = { "input_ids": torch.randint(1, 64, (B, T)), "targets": torch.randint(0, 64, (B, T)), # Note: NO ctx_teacher_input_ids — this is what the collator does # when there are no error turns in the batch. } out = composer_total_loss(model, batch, alpha_sdpo=1.0, beta_replay=0.0) assert out["sdpo"].item() == 0.0, "SDPO must be 0 when no SDPO inputs in batch" assert torch.isclose(out["total"], out["grpo"])