Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
test_po_objective_menu.py88 linesDownload Raw Back to tests
1"""Tests for the policy-optimization objective menu (make_po_config, ADR-014).2 3These build real trl GRPOConfigs, so they require trl installed (the framework's4.venv has trl==1.5.0). Skips cleanly if trl is absent.5"""6from __future__ import annotations7 8import pytest9 10trl = pytest.importorskip("trl")11 12from composer_replication.trainer.composer_trainer import (  # noqa: E40213    PO_OBJECTIVES,14    make_po_config,15)16 17 18def test_menu_lists_expected_objectives():19    assert set(PO_OBJECTIVES) == {"grpo", "dr_grpo", "bnpo", "dapo", "gspo", "cispo"}20 21 22def test_unknown_objective_raises_with_menu(tmp_path):23    with pytest.raises(ValueError) as ei:24        make_po_config("nope", output_dir=str(tmp_path))25    msg = str(ei.value)26    assert "Unknown PO objective" in msg and "dapo" in msg and "gspo" in msg27 28 29def test_grpo_preset(tmp_path):30    cfg = make_po_config("grpo", output_dir=str(tmp_path))31    assert str(cfg.loss_type) == "grpo"32    assert str(cfg.importance_sampling_level) == "token"33    # group scaling = std-normalized advantage (vanilla GRPO)34    assert str(cfg.scale_rewards).lower() in ("group", "true")35 36 37def test_dr_grpo_preset_matches_legacy(tmp_path):38    cfg = make_po_config("dr_grpo", output_dir=str(tmp_path))39    assert str(cfg.loss_type) == "dr_grpo"40    # no std-normalization (the Dr.GRPO fix)41    assert str(cfg.scale_rewards).lower() in ("none", "false")42 43 44def test_dapo_preset_sets_decoupled_clip(tmp_path):45    cfg = make_po_config("dapo", output_dir=str(tmp_path))46    assert str(cfg.loss_type) == "dapo"47    # clip-higher: epsilon_high strictly above epsilon48    assert cfg.epsilon_high is not None49    assert float(cfg.epsilon_high) > float(cfg.epsilon)50    assert bool(cfg.mask_truncated_completions) is True51    assert float(cfg.beta) == 0.0  # DAPO removes KL52 53 54def test_gspo_is_sequence_level(tmp_path):55    cfg = make_po_config("gspo", output_dir=str(tmp_path))56    # GSPO = grpo loss + SEQUENCE-level importance ratio57    assert str(cfg.loss_type) == "grpo"58    assert str(cfg.importance_sampling_level) == "sequence"59 60 61def test_gspo_guard_rejects_token_override(tmp_path):62    # Overriding back to token-level would silently degrade GSPO to GRPO -> guard.63    with pytest.raises(AssertionError):64        make_po_config(65            "gspo", output_dir=str(tmp_path), importance_sampling_level="token"66        )67 68 69def test_cispo_preset(tmp_path):70    cfg = make_po_config("cispo", output_dir=str(tmp_path))71    assert str(cfg.loss_type) == "cispo"72    # eps_max (ScaleRL recommended 5.0) carried via epsilon_high73    assert cfg.epsilon_high is not None and float(cfg.epsilon_high) >= 5.074 75 76def test_overrides_apply_on_top(tmp_path):77    cfg = make_po_config(78        "dr_grpo", output_dir=str(tmp_path), beta=0.05, num_generations=479    )80    assert float(cfg.beta) == 0.0581    assert int(cfg.num_generations) == 482    assert str(cfg.loss_type) == "dr_grpo"  # preset preserved under overrides83 84 85def test_default_objective_is_dr_grpo(tmp_path):86    cfg = make_po_config(output_dir=str(tmp_path))87    assert str(cfg.loss_type) == "dr_grpo"88