From a1215959a70ea6590d1d91bf02ae10bf44795a17 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=EC=A0=95=EC=8B=9C=EC=9B=90?= Date: Thu, 7 May 2026 17:03:28 +0900 Subject: [PATCH] Warn when max_iterations is too small for any eval to run If eval_every is positive but max_iterations falls before the next scheduled eval iteration (including resume cases where current_iteration is already past the last eval boundary), log a one-time warning at run start so the user notices the misconfiguration. We deliberately do not force an end-of-run eval, which would distort time budgets and reproducibility. Co-Authored-By: Claude Opus 4.7 (1M context) --- .../games/classic/deep_cfr/trainer.py | 30 +++++ tests/games/classic/test_deep_cfr_trainer.py | 103 +++++++++++++++++- 2 files changed, 132 insertions(+), 1 deletion(-) diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py b/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py index 7465c89..a127947 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py @@ -156,6 +156,29 @@ def _format_summary_value(value: float | int) -> str: return str(value) +def _next_eval_iteration(current_iteration: int, eval_every: int) -> int | None: + if eval_every <= 0: + return None + return ((current_iteration // eval_every) + 1) * eval_every + + +def eval_skipped_warning( + current_iteration: int, + max_iterations: int | None, + eval_every: int, +) -> str | None: + if eval_every <= 0 or max_iterations is None: + return None + next_eval = _next_eval_iteration(current_iteration, eval_every) + if next_eval is None or next_eval <= max_iterations: + return None + return ( + f"WARNING evaluation will not run: eval_every={eval_every} " + f"but max_iterations={max_iterations} (current iteration={current_iteration}). " + f"Next scheduled eval at iteration {next_eval}." + ) + + class DeepCFRTrainer: def __init__( self, @@ -612,6 +635,13 @@ class DeepCFRTrainer: self.tracker.log_event( f"Deep CFR run start iteration={self.iteration} seed={self.config.run.seed}" ) + warning = eval_skipped_warning( + self.iteration, + self.config.run.max_iterations, + self.config.evaluation.eval_every, + ) + if warning is not None: + self.tracker.log_event(warning) def _append_metrics(self, metrics: IterationMetrics, iteration_seconds: float) -> None: data = metrics.to_dict() diff --git a/tests/games/classic/test_deep_cfr_trainer.py b/tests/games/classic/test_deep_cfr_trainer.py index a9770b1..e92b4c9 100644 --- a/tests/games/classic/test_deep_cfr_trainer.py +++ b/tests/games/classic/test_deep_cfr_trainer.py @@ -24,7 +24,10 @@ from coolrl_lost_cities.games.classic.deep_cfr.config import DeepCFRConfig, load from coolrl_lost_cities.games.classic.deep_cfr.evaluate import evaluate_strategy_network from coolrl_lost_cities.games.classic.deep_cfr.memory import ReservoirMemory, TrainingSample from coolrl_lost_cities.games.classic.deep_cfr.networks import DeepCFRMLP -from coolrl_lost_cities.games.classic.deep_cfr.trainer import DeepCFRTrainer +from coolrl_lost_cities.games.classic.deep_cfr.trainer import ( + DeepCFRTrainer, + eval_skipped_warning, +) def _deep_cfr_config(data: dict) -> DeepCFRConfig: @@ -857,3 +860,101 @@ def test_deep_cfr_weighted_self_play_league_uses_snapshot_bucket(tmp_path) -> No assert len(metrics) == 2 assert len(trainer.self_play_league_snapshots) == 2 assert metrics[1].traversal_nodes > 0 + + +def test_eval_skipped_warning_returns_none_when_eval_disabled() -> None: + assert eval_skipped_warning(0, max_iterations=10, eval_every=0) is None + assert eval_skipped_warning(0, max_iterations=10, eval_every=-5) is None + + +def test_eval_skipped_warning_returns_none_when_max_iterations_unbounded() -> None: + assert eval_skipped_warning(0, max_iterations=None, eval_every=50) is None + + +def test_eval_skipped_warning_returns_none_when_eval_will_run_within_budget() -> None: + assert eval_skipped_warning(0, max_iterations=50, eval_every=50) is None + assert eval_skipped_warning(0, max_iterations=51, eval_every=50) is None + assert eval_skipped_warning(0, max_iterations=200, eval_every=50) is None + + +def test_eval_skipped_warning_returns_none_when_starting_just_before_an_eval() -> None: + # Resuming at iteration=49: next iter (50) is an eval. + assert eval_skipped_warning(49, max_iterations=50, eval_every=50) is None + + +def test_eval_skipped_warning_warns_when_max_below_first_eval() -> None: + warning = eval_skipped_warning(0, max_iterations=10, eval_every=50) + assert warning is not None + assert "max_iterations=10" in warning + assert "iteration 50" in warning + + +def test_eval_skipped_warning_warns_on_resume_when_no_more_evals_fit() -> None: + # Resumed at iter 100 (just past the last scheduled eval). + # Next eval is 150, but we only have budget up to 110. + warning = eval_skipped_warning(100, max_iterations=110, eval_every=50) + assert warning is not None + assert "iteration 150" in warning + + +def test_eval_skipped_warning_no_warn_when_resume_lands_exactly_on_next_eval() -> None: + # iter=50 means iteration 50 already evaluated; next is 100. Budget 100 fits. + assert eval_skipped_warning(50, max_iterations=100, eval_every=50) is None + + +def test_eval_skipped_warning_warns_when_eval_every_one_but_max_is_zero() -> None: + # Pathological: max_iterations=0 means no iterations will run. + warning = eval_skipped_warning(0, max_iterations=0, eval_every=1) + assert warning is not None + + +def test_deep_cfr_trainer_logs_eval_skipped_warning_on_start(tmp_path) -> None: + trainer = DeepCFRTrainer( + _deep_cfr_config( + { + "run": {"max_iterations": 1, "seed": 90}, + "network": {"hidden_size": 16}, + "traversal": {"traversals_per_player": 1, "max_depth": 1}, + "optimization": { + "advantage_batch_size": 2, + "strategy_batch_size": 2, + "advantage_updates_per_iteration": 1, + "strategy_updates_per_iteration": 1, + }, + "checkpoint": {"save_every": 0, "save_latest": False}, + "evaluation": {"eval_every": 50, "games": 2, "opponents": ["random"]}, + } + ), + LostCitiesConfig(seed=90), + run_dir=tmp_path, + ) + + trainer.train() + log_text = (tmp_path / "train.log").read_text() + assert "WARNING evaluation will not run" in log_text + + +def test_deep_cfr_trainer_does_not_log_eval_warning_when_eval_disabled(tmp_path) -> None: + trainer = DeepCFRTrainer( + _deep_cfr_config( + { + "run": {"max_iterations": 1, "seed": 91}, + "network": {"hidden_size": 16}, + "traversal": {"traversals_per_player": 1, "max_depth": 1}, + "optimization": { + "advantage_batch_size": 2, + "strategy_batch_size": 2, + "advantage_updates_per_iteration": 1, + "strategy_updates_per_iteration": 1, + }, + "checkpoint": {"save_every": 0, "save_latest": False}, + "evaluation": {"eval_every": 0}, + } + ), + LostCitiesConfig(seed=91), + run_dir=tmp_path, + ) + + trainer.train() + log_text = (tmp_path / "train.log").read_text() + assert "WARNING evaluation will not run" not in log_text