diff --git a/src/coolrl_lost_cities/games/classic/ismcts/cli.py b/src/coolrl_lost_cities/games/classic/ismcts/cli.py index cbd3639..3ed96fd 100644 --- a/src/coolrl_lost_cities/games/classic/ismcts/cli.py +++ b/src/coolrl_lost_cities/games/classic/ismcts/cli.py @@ -8,6 +8,8 @@ from typing import Any import yaml +from coolrl_lost_cities.games.classic.deep_cfr.tracking import WandbRunTracker + from .config import IsMctsConfig, load_config from .trainer import IsMctsTrainer @@ -57,13 +59,31 @@ def train_command(args: argparse.Namespace) -> None: config = load_config(args.config) if args.config else IsMctsConfig() config = _with_overrides(config, args.config_overrides) run_dir = _resolve_run_dir(config, keep=args.keep) + tracker = None + if args.wandb: + tracker = WandbRunTracker( + project=args.wandb_project, + name=args.wandb_name or config.run.experiment_name, + mode=args.wandb_mode, + run_dir=run_dir, + config=config.to_dict() if hasattr(config, "to_dict") else config.model_dump(), + group=args.wandb_group, + job_type=args.wandb_job_type, + tags=list(args.wandb_tag) if args.wandb_tag else None, + notes=args.wandb_notes, + ) trainer = IsMctsTrainer( config, config.rules.to_lost_cities_config(seed=config.run.seed), run_dir=run_dir, device=config.run.device, + tracker=tracker, ) - trainer.train() + try: + trainer.train() + finally: + if tracker is not None: + tracker.close() def main(argv: list[str] | None = None) -> None: @@ -79,7 +99,11 @@ def main(argv: list[str] | None = None) -> None: dest="config_overrides", metavar="PATH=VALUE", ) - train.add_argument("--wandb", action="store_true", help="Accepted for CLI parity; ignored.") + train.add_argument( + "--wandb", + action="store_true", + help="Mirror metrics to Weights & Biases (requires wandb extra).", + ) train.add_argument("--wandb-project", default="coolrl-lost-cities") train.add_argument("--wandb-name") train.add_argument("--wandb-mode", choices=("online", "offline", "disabled"), default="online") diff --git a/src/coolrl_lost_cities/games/classic/ismcts/trainer.py b/src/coolrl_lost_cities/games/classic/ismcts/trainer.py index bbad61d..ed52112 100644 --- a/src/coolrl_lost_cities/games/classic/ismcts/trainer.py +++ b/src/coolrl_lost_cities/games/classic/ismcts/trainer.py @@ -57,11 +57,13 @@ class IsMctsTrainer: *, run_dir: str | Path, device: torch.device | str = "cpu", + tracker: object | None = None, ) -> None: self.config = config self.game_config = game_config self.run_dir = Path(run_dir) self.device = self._resolve_device(device) + self.tracker = tracker probe = GameState.new_game(game_config, seed=config.run.seed) self.input_dim = input_dim(probe, config.encoding) self.action_size = probe.action_size @@ -101,6 +103,11 @@ class IsMctsTrainer: metrics.append(item) self._append_metrics(item) self._save_checkpoints(iteration, item) + if self.tracker is not None: + try: + self.tracker.log_metrics(item.to_dict(), step=iteration) + except Exception as exc: # pragma: no cover + print(f"tracker.log_metrics failed: {exc}", flush=True) print(json.dumps(item.to_dict(), sort_keys=True)) return metrics