Wire W&B tracking into ISMCTS trainer

This commit is contained in:
2026-05-10 23:13:03 +09:00
parent 812dace1e3
commit 0999d34277
2 changed files with 33 additions and 2 deletions
@@ -8,6 +8,8 @@ from typing import Any
import yaml import yaml
from coolrl_lost_cities.games.classic.deep_cfr.tracking import WandbRunTracker
from .config import IsMctsConfig, load_config from .config import IsMctsConfig, load_config
from .trainer import IsMctsTrainer 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 = load_config(args.config) if args.config else IsMctsConfig()
config = _with_overrides(config, args.config_overrides) config = _with_overrides(config, args.config_overrides)
run_dir = _resolve_run_dir(config, keep=args.keep) 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( trainer = IsMctsTrainer(
config, config,
config.rules.to_lost_cities_config(seed=config.run.seed), config.rules.to_lost_cities_config(seed=config.run.seed),
run_dir=run_dir, run_dir=run_dir,
device=config.run.device, device=config.run.device,
tracker=tracker,
) )
try:
trainer.train() trainer.train()
finally:
if tracker is not None:
tracker.close()
def main(argv: list[str] | None = None) -> None: def main(argv: list[str] | None = None) -> None:
@@ -79,7 +99,11 @@ def main(argv: list[str] | None = None) -> None:
dest="config_overrides", dest="config_overrides",
metavar="PATH=VALUE", 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-project", default="coolrl-lost-cities")
train.add_argument("--wandb-name") train.add_argument("--wandb-name")
train.add_argument("--wandb-mode", choices=("online", "offline", "disabled"), default="online") train.add_argument("--wandb-mode", choices=("online", "offline", "disabled"), default="online")
@@ -57,11 +57,13 @@ class IsMctsTrainer:
*, *,
run_dir: str | Path, run_dir: str | Path,
device: torch.device | str = "cpu", device: torch.device | str = "cpu",
tracker: object | None = None,
) -> None: ) -> None:
self.config = config self.config = config
self.game_config = game_config self.game_config = game_config
self.run_dir = Path(run_dir) self.run_dir = Path(run_dir)
self.device = self._resolve_device(device) self.device = self._resolve_device(device)
self.tracker = tracker
probe = GameState.new_game(game_config, seed=config.run.seed) probe = GameState.new_game(game_config, seed=config.run.seed)
self.input_dim = input_dim(probe, config.encoding) self.input_dim = input_dim(probe, config.encoding)
self.action_size = probe.action_size self.action_size = probe.action_size
@@ -101,6 +103,11 @@ class IsMctsTrainer:
metrics.append(item) metrics.append(item)
self._append_metrics(item) self._append_metrics(item)
self._save_checkpoints(iteration, 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)) print(json.dumps(item.to_dict(), sort_keys=True))
return metrics return metrics