Wire W&B tracking into ISMCTS trainer
This commit is contained in:
@@ -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,
|
||||||
)
|
)
|
||||||
trainer.train()
|
try:
|
||||||
|
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
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user