Use generic train config overrides

This commit is contained in:
2026-05-07 15:53:43 +09:00
parent 44dd97876e
commit 0fea13786f
3 changed files with 47 additions and 129 deletions
+23 -13
View File
@@ -66,8 +66,11 @@ Short fixed-iteration run:
```bash ```bash
uv run lost-cities-deep-cfr train \ uv run lost-cities-deep-cfr train \
--config configs/deep_cfr/deep_cfr_selfplay_full_depth_slot_playability.yaml \ --config configs/deep_cfr/deep_cfr_selfplay_full_depth_slot_playability.yaml \
--iterations 100 \ --set run.iterations=100 \
--save-latest-only --set run.max_hours=null \
--set run.max_iterations=null \
--set checkpoint.save_latest_only=true \
--set checkpoint.save_every_iteration=false
``` ```
Use explicit run directories for experiments. Put Deep CFR runs under Use explicit run directories for experiments. Put Deep CFR runs under
@@ -77,8 +80,8 @@ Use explicit run directories for experiments. Put Deep CFR runs under
RUN_DIR="runs/deep_cfr/$(date +%Y-%m-%d_%H%M%S)_deep_cfr_experiment_name" RUN_DIR="runs/deep_cfr/$(date +%Y-%m-%d_%H%M%S)_deep_cfr_experiment_name"
uv run lost-cities-deep-cfr train \ uv run lost-cities-deep-cfr train \
--config configs/deep_cfr/deep_cfr_selfplay_full_depth_slot_playability.yaml \ --config configs/deep_cfr/deep_cfr_selfplay_full_depth_slot_playability.yaml \
--checkpoint-dir "$RUN_DIR" \ --set checkpoint.directory="$RUN_DIR" \
--max-iterations 100 --set run.max_iterations=100
``` ```
Date-prefixed examples: Date-prefixed examples:
@@ -86,16 +89,23 @@ Date-prefixed examples:
- `runs/deep_cfr/YYYY-MM-DD_HHMMSS_deep_cfr_100iter` - `runs/deep_cfr/YYYY-MM-DD_HHMMSS_deep_cfr_100iter`
- `runs/deep_cfr/YYYY-MM-DD_HHMMSS_deep_cfr_unbounded` - `runs/deep_cfr/YYYY-MM-DD_HHMMSS_deep_cfr_unbounded`
Useful train overrides: Useful train controls:
- `--resume`: resume from `<checkpoint-dir>/latest.pt`. - `--resume`: resume from `<checkpoint-dir>/latest.pt`.
- `--resume PATH`: resume from a specific checkpoint. - `--resume PATH`: resume from a specific checkpoint.
- `--exact-resume`: require checkpoint config compatibility for exact resume. - `--device DEVICE`: override the trainer device for this invocation.
- `--no-save`: disable checkpoint writes. - `--set PATH=VALUE`: override config fields. It is repeatable and parses
- `--save-latest-only`: keep only `latest.pt`. values as YAML, e.g. `--set traversal.num_workers=4` or
- `--save-iteration-interval N`: archive every N iterations. `--set run.max_hours=null`.
- `--set PATH=VALUE`: override arbitrary config fields, e.g.
`--set traversal.num_workers=4`. Common `--set` overrides:
- `--set checkpoint.exact_resume=true`: require checkpoint config compatibility.
- `--set checkpoint.save_latest=false --set checkpoint.save_every_iteration=false
--set checkpoint.save_iteration_interval=0`: disable checkpoint writes.
- `--set checkpoint.save_latest_only=true --set checkpoint.save_every_iteration=false`:
keep only `latest.pt`.
- `--set checkpoint.save_iteration_interval=N`: archive every N iterations.
## Long Runs ## Long Runs
@@ -147,8 +157,8 @@ checkpoint:
``` ```
`latest.pt` is updated continuously; archive checkpoints are written every 100 `latest.pt` is updated continuously; archive checkpoints are written every 100
iterations. If disk is tight, prefer `--save-latest-only` or increase iterations. If disk is tight, prefer `--set checkpoint.save_latest_only=true` or
`save_iteration_interval`. increase `save_iteration_interval`.
## Evaluation And Analysis ## Evaluation And Analysis
@@ -80,52 +80,6 @@ def _resolve_resume_path(config: DeepCFRConfig, resume: str | None) -> str | Non
def _train_overrides_from_args(args: argparse.Namespace) -> dict[str, Any]: def _train_overrides_from_args(args: argparse.Namespace) -> dict[str, Any]:
overrides: dict[str, Any] = {} overrides: dict[str, Any] = {}
if args.no_save and (args.save_latest_only or args.save_iteration_interval is not None):
raise ValueError("--no-save cannot be combined with checkpoint save overrides")
run_overrides = overrides.setdefault("run", {})
if args.iterations is not None:
run_overrides["iterations"] = args.iterations
if args.max_hours is None and args.max_iterations is None:
run_overrides["max_hours"] = None
run_overrides["max_iterations"] = None
if args.max_hours is not None:
run_overrides["max_hours"] = args.max_hours
if args.max_iterations is not None:
run_overrides["max_iterations"] = args.max_iterations
if args.seed is not None:
run_overrides["seed"] = args.seed
if args.traversals_per_iteration is not None:
traversal_overrides = overrides.setdefault("traversal", {})
traversal_overrides["traversals_per_iteration"] = args.traversals_per_iteration
traversal_overrides["traversals_per_player"] = None
if args.num_workers is not None:
overrides.setdefault("traversal", {})["num_workers"] = args.num_workers
if args.checkpoint_dir is not None:
overrides.setdefault("checkpoint", {})["directory"] = args.checkpoint_dir
if args.eval_every is not None:
overrides.setdefault("evaluation", {})["eval_every"] = args.eval_every
if args.eval_games is not None:
overrides.setdefault("evaluation", {})["games"] = args.eval_games
if args.regret_fallback is not None:
overrides.setdefault("regret_matching", {})["all_negative_fallback"] = args.regret_fallback
if args.training_weighting is not None:
overrides.setdefault("training_weighting", {})["mode"] = args.training_weighting
if args.no_save:
checkpoint_overrides = overrides.setdefault("checkpoint", {})
checkpoint_overrides["save_latest"] = False
checkpoint_overrides["save_every_iteration"] = False
checkpoint_overrides["save_iteration_interval"] = 0
if args.save_latest_only:
checkpoint_overrides = overrides.setdefault("checkpoint", {})
checkpoint_overrides["save_latest"] = True
checkpoint_overrides["save_latest_only"] = True
checkpoint_overrides["save_every_iteration"] = False
if args.save_iteration_interval is not None:
overrides.setdefault("checkpoint", {})["save_iteration_interval"] = (
args.save_iteration_interval
)
if args.exact_resume:
overrides.setdefault("checkpoint", {})["exact_resume"] = True
for assignment in getattr(args, "config_overrides", None) or (): for assignment in getattr(args, "config_overrides", None) or ():
_set_path_override(overrides, assignment) _set_path_override(overrides, assignment)
return overrides return overrides
@@ -249,39 +203,19 @@ def main(argv: list[str] | None = None) -> None:
train = subparsers.add_parser("train") train = subparsers.add_parser("train")
train.add_argument("--config") train.add_argument("--config")
train.add_argument("--iterations", type=int)
train.add_argument("--max-hours", type=float)
train.add_argument("--max-iterations", type=int)
train.add_argument("--traversals-per-iteration", type=int)
train.add_argument("--num-workers")
train.add_argument("--checkpoint-dir")
train.add_argument("--resume", nargs="?", const=_RESUME_LATEST, default=None) train.add_argument("--resume", nargs="?", const=_RESUME_LATEST, default=None)
train.add_argument("--exact-resume", action="store_true")
train.add_argument("--device") train.add_argument("--device")
train.add_argument("--eval-every", type=int)
train.add_argument("--eval-games", type=int)
train.add_argument(
"--regret-fallback",
choices=("uniform", "argmax_tiebreak"),
help="Override regret_matching.all_negative_fallback.",
)
train.add_argument(
"--training-weighting",
choices=("none", "lcfr", "dcfr"),
help="Override training_weighting.mode.",
)
train.add_argument("--seed", type=int)
train.add_argument( train.add_argument(
"--set", "--set",
action="append", action="append",
default=[], default=[],
dest="config_overrides", dest="config_overrides",
metavar="PATH=VALUE", metavar="PATH=VALUE",
help="Override any config field using dotted paths, e.g. traversal.sampling_mode=external.", help=(
"Override config fields using dotted paths. Repeatable. VALUE is parsed as "
"YAML, e.g. --set traversal.num_workers=4 --set run.max_hours=null."
),
) )
train.add_argument("--no-save", action="store_true")
train.add_argument("--save-latest-only", action="store_true")
train.add_argument("--save-iteration-interval", type=int)
train.add_argument( train.add_argument(
"--wandb", "--wandb",
action="store_true", action="store_true",
+20 -46
View File
@@ -74,26 +74,24 @@ def test_deep_cfr_loads_mapped_legacy_reproduction_config() -> None:
) )
def test_deep_cfr_train_cli_count_overrides_disable_duration_limits() -> None: def test_deep_cfr_train_cli_accepts_run_and_traversal_config_overrides() -> None:
args = type( args = type(
"Args", "Args",
(), (),
{ {
"iterations": 1, "config_overrides": [
"max_hours": None, "run.iterations=1",
"max_iterations": None, "run.max_hours=null",
"seed": None, "run.max_iterations=null",
"traversals_per_iteration": 1, "traversal.traversals_per_iteration=1",
"num_workers": "0", "traversal.traversals_per_player=null",
"checkpoint_dir": None, "traversal.num_workers=0",
"eval_every": None, "regret_matching.all_negative_fallback=argmax_tiebreak",
"eval_games": None, "training_weighting.mode=lcfr",
"regret_fallback": "argmax_tiebreak", "checkpoint.save_latest=false",
"training_weighting": "lcfr", "checkpoint.save_every_iteration=false",
"no_save": True, "checkpoint.save_iteration_interval=0",
"save_latest_only": False, ],
"save_iteration_interval": None,
"exact_resume": False,
}, },
)() )()
config = load_config("configs/deep_cfr/deep_cfr_selfplay_full_depth_slot_playability.yaml") config = load_config("configs/deep_cfr/deep_cfr_selfplay_full_depth_slot_playability.yaml")
@@ -123,21 +121,12 @@ def test_deep_cfr_train_cli_checkpoint_save_overrides() -> None:
"Args", "Args",
(), (),
{ {
"iterations": None, "config_overrides": [
"max_hours": None, "checkpoint.save_latest=true",
"max_iterations": None, "checkpoint.save_latest_only=true",
"seed": None, "checkpoint.save_every_iteration=false",
"traversals_per_iteration": None, "checkpoint.save_iteration_interval=1",
"num_workers": None, ],
"checkpoint_dir": None,
"eval_every": None,
"eval_games": None,
"regret_fallback": None,
"training_weighting": None,
"no_save": False,
"save_latest_only": True,
"save_iteration_interval": 1,
"exact_resume": False,
}, },
)() )()
@@ -154,21 +143,6 @@ def test_deep_cfr_train_cli_accepts_generic_config_overrides() -> None:
"Args", "Args",
(), (),
{ {
"iterations": None,
"max_hours": None,
"max_iterations": None,
"seed": None,
"traversals_per_iteration": None,
"num_workers": None,
"checkpoint_dir": None,
"eval_every": None,
"eval_games": None,
"regret_fallback": None,
"training_weighting": None,
"no_save": False,
"save_latest_only": False,
"save_iteration_interval": None,
"exact_resume": False,
"config_overrides": [ "config_overrides": [
"traversal.sampling_mode=external", "traversal.sampling_mode=external",
"traversal.max_depth=null", "traversal.max_depth=null",