Team Ai
Modelpublic

Codeseys/composer-replication-framework

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