diff --git a/training/evaluate.py b/training/evaluate.py index ec665e9f..786fdbe4 100644 --- a/training/evaluate.py +++ b/training/evaluate.py @@ -82,6 +82,8 @@ def evaluate_pair( """Replay one seeded state sequence with the models on opposite sides.""" if episodes < 2 or episodes % 2 != 0: raise ValueError("--episodes must be an even number of at least 2 for paired side swaps") + if team_size not in (1, 2): + raise ValueError("team_size must be 1 or 2") episodes_per_side = episodes // 2 first = run_half( diff --git a/training/test_evaluate.py b/training/test_evaluate.py index 55b5fe46..ae6736a0 100644 --- a/training/test_evaluate.py +++ b/training/test_evaluate.py @@ -45,6 +45,15 @@ class EvaluatePairTests(unittest.TestCase): def test_rejects_unsupported_team_size_before_launch(self) -> None: with self.assertRaisesRegex(ValueError, "team_size must be 1 or 2"): evaluate.run_half("godot", "a", "b", 2, 16, 1, team_size=3) + with self.assertRaisesRegex(ValueError, "team_size must be 1 or 2"): + evaluate.evaluate_pair("godot", "a", "b", 2, 16, 1, team_size=3) + + @patch("evaluate.subprocess.run") + def test_2v2_run_passes_team_size_to_godot(self, run_process) -> None: + run_process.return_value.stdout = 'EVAL_RESULT {"episodes": 2, "goals_a": 1, "goals_b": 0, "draws": 1}\n' + evaluate.run_half("godot", "a", "b", 2, 16, 9, team_size=2) + command = run_process.call_args.args[0] + self.assertIn("--eval_team_size=2", command) @patch("evaluate.run_half") def test_2v2_evaluation_preserves_side_swap_and_team_size(self, run_half) -> None: