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) <noreply@anthropic.com>
This commit is contained in:
2026-05-07 17:03:28 +09:00
co-authored by Claude Opus 4.7
parent b561fa8457
commit a1215959a7
2 changed files with 132 additions and 1 deletions
@@ -156,6 +156,29 @@ def _format_summary_value(value: float | int) -> str:
return str(value) 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: class DeepCFRTrainer:
def __init__( def __init__(
self, self,
@@ -612,6 +635,13 @@ class DeepCFRTrainer:
self.tracker.log_event( self.tracker.log_event(
f"Deep CFR run start iteration={self.iteration} seed={self.config.run.seed}" 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: def _append_metrics(self, metrics: IterationMetrics, iteration_seconds: float) -> None:
data = metrics.to_dict() data = metrics.to_dict()
+102 -1
View File
@@ -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.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.memory import ReservoirMemory, TrainingSample
from coolrl_lost_cities.games.classic.deep_cfr.networks import DeepCFRMLP 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: 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(metrics) == 2
assert len(trainer.self_play_league_snapshots) == 2 assert len(trainer.self_play_league_snapshots) == 2
assert metrics[1].traversal_nodes > 0 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