Codeseys/composer-replication-framework
0
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 