Codeseys/composer-replication-framework
0
1"""Numerical parity test against the upstream OPSD reference.2 3Loads `OPSDTrainer.generalized_jsd_loss` from a clone of siyan-zhao/OPSD at4/tmp/opsd-clone (override with $OPSD_CLONE) and asserts our re-implementation5in `composer_replication.opsd` matches it byte-for-byte across a grid of6shapes and β values. Skips cleanly when the upstream clone is absent.7 8Why this lives in `tests/` rather than docs: numerical parity is the9contract for this lift. If a future refactor of `generalized_jsd_loss`10silently shifts gradients again, this test fails immediately.11"""12 13from __future__ import annotations14 15import importlib.util16import os17import sys18from pathlib import Path19 20import pytest21import torch22 23from composer_replication.opsd import generalized_jsd_loss24 25# ----------------------------------------------------------------------26# Locate upstream OPSDTrainer.generalized_jsd_loss27# ----------------------------------------------------------------------28 29_OPSD_CLONE = Path(os.environ.get("OPSD_CLONE", "/tmp/opsd-clone"))30_OPSD_TRAINER_PATH = _OPSD_CLONE / "opsd_trainer.py"31 32 33def _load_upstream():34 """Import OPSDTrainer.generalized_jsd_loss from a local clone, isolated.35 36 The upstream `opsd_trainer.py` imports heavyweight TRL / transformers37 machinery at module scope, which we do not want to drag into the test38 process. We instead extract the static method by parsing the source39 text and exec-ing only that function body — it depends only on40 `torch` and `torch.nn.functional`, which are already importable.41 """42 if not _OPSD_TRAINER_PATH.exists():43 return None44 45 text = _OPSD_TRAINER_PATH.read_text()46 # Pull out the function block. It starts with `def generalized_jsd_loss(`47 # under `class OPSDTrainer` and ends at the next top-of-class `def `.48 start = text.find("def generalized_jsd_loss(")49 if start < 0:50 return None51 # Walk forward to the start of the next sibling method (4-space indent52 # `def ` or class-end) — they all start with exactly 4 spaces of indent.53 rest = text[start:]54 # Skip past the function header and find the next `\n def ` or55 # `\n @staticmethod` boundary.56 end_marker_offsets = []57 for marker in ("\n @", "\n def ", "\nclass "):58 idx = rest.find(marker, len("def generalized_jsd_loss("))59 if idx > 0:60 end_marker_offsets.append(idx)61 if not end_marker_offsets:62 return None63 fn_text = rest[: min(end_marker_offsets)]64 65 # Dedent (the source lines are 4-space indented as a class method).66 fn_text = "\n".join(67 line[4:] if line.startswith(" ") else line for line in fn_text.splitlines()68 )69 70 # Exec into a fresh namespace with torch + F available.71 import torch.nn.functional as F # noqa: F401 (used by exec'd code)72 73 namespace: dict = {"torch": torch, "F": F}74 exec(compile(fn_text, str(_OPSD_TRAINER_PATH), "exec"), namespace)75 fn = namespace.get("generalized_jsd_loss")76 return fn77 78 79_UPSTREAM_FN = _load_upstream()80_SKIP_REASON = (81 f"upstream OPSD clone not found at {_OPSD_TRAINER_PATH} "82 f"(set $OPSD_CLONE or `git clone --depth 1 https://github.com/siyan-zhao/OPSD {_OPSD_CLONE}`)"83)84 85 86# ----------------------------------------------------------------------87# Parity grid88# ----------------------------------------------------------------------89 90_SHAPES = [91 (1, 4, 16),92 (2, 8, 32),93 (3, 5, 64),94 (1, 16, 8),95 (4, 3, 24),96]97_BETAS = [0.0, 0.5, 1.0]98 99 100@pytest.mark.skipif(_UPSTREAM_FN is None, reason=_SKIP_REASON)101@pytest.mark.parametrize("shape", _SHAPES)102@pytest.mark.parametrize("beta", _BETAS)103def test_parity_unmasked(shape, beta):104 """Our `generalized_jsd_loss` must match upstream within 1e-5 atol."""105 B, T, V = shape106 g = torch.Generator().manual_seed(13 + B * 31 + T * 17 + V)107 student = torch.randn(B, T, V, generator=g, dtype=torch.float64)108 teacher = torch.randn(B, T, V, generator=g, dtype=torch.float64)109 110 ours = generalized_jsd_loss(student, teacher, beta=beta)111 theirs = _UPSTREAM_FN(student, teacher, beta=beta) # type: ignore[misc]112 113 assert torch.allclose(ours, theirs, atol=1e-5, rtol=1e-5), (114 f"mismatch at shape={shape} beta={beta}: ours={ours.item()} theirs={theirs.item()}"115 )116 117 118@pytest.mark.skipif(_UPSTREAM_FN is None, reason=_SKIP_REASON)119@pytest.mark.parametrize("shape", _SHAPES)120@pytest.mark.parametrize("beta", _BETAS)121def test_parity_masked(shape, beta):122 """Same parity but with a labels mask that ignores ~half the tokens."""123 B, T, V = shape124 g = torch.Generator().manual_seed(101 + B * 7 + T * 11 + V)125 student = torch.randn(B, T, V, generator=g, dtype=torch.float64)126 teacher = torch.randn(B, T, V, generator=g, dtype=torch.float64)127 # Random valid/ignored mask: -100 for ignored, anything else for valid.128 labels = torch.randint(0, 2, (B, T), generator=g)129 labels = torch.where(labels == 0, torch.full_like(labels, -100), labels)130 131 ours = generalized_jsd_loss(student, teacher, labels=labels, beta=beta)132 theirs = _UPSTREAM_FN(student, teacher, labels=labels, beta=beta) # type: ignore[misc]133 134 assert torch.allclose(ours, theirs, atol=1e-5, rtol=1e-5), (135 f"mismatch at shape={shape} beta={beta}: ours={ours.item()} theirs={theirs.item()}"136 )137 138 139@pytest.mark.skipif(_UPSTREAM_FN is None, reason=_SKIP_REASON)140def test_parity_temperature_and_topk():141 """Spot-check the temperature + top_k branches against upstream."""142 g = torch.Generator().manual_seed(42)143 student = torch.randn(2, 6, 32, generator=g, dtype=torch.float64)144 teacher = torch.randn(2, 6, 32, generator=g, dtype=torch.float64)145 146 for beta in (0.0, 0.3, 0.5, 0.7, 1.0):147 ours = generalized_jsd_loss(student, teacher, beta=beta, temperature=2.0, top_k=8)148 theirs = _UPSTREAM_FN( # type: ignore[misc]149 student, teacher, beta=beta, temperature=2.0, top_k=8150 )151 assert torch.allclose(ours, theirs, atol=1e-5, rtol=1e-5), (152 f"temp+topk parity failed at beta={beta}: ours={ours.item()} theirs={theirs.item()}"153 )154 