diff --git a/deep_quoridor/src/train_v2.py b/deep_quoridor/src/train_v2.py index 57db5e62..bb4fcbef 100644 --- a/deep_quoridor/src/train_v2.py +++ b/deep_quoridor/src/train_v2.py @@ -1,5 +1,6 @@ import argparse import multiprocessing as mp +import os import subprocess import time from pathlib import Path @@ -7,6 +8,9 @@ from v2 import benchmarks, load_config_and_setup_run, self_play, train from v2.common import ShutdownSignal +# Prevents getting messages in the console every few lines telling you to install weave +os.environ["WANDB_DISABLE_WEAVE"] = "true" + if __name__ == "__main__": parser = argparse.ArgumentParser(description="Train Quoridor agent") parser.add_argument("config_file", type=str, help="Path to YAML configuration file") diff --git a/deep_quoridor/src/v2/common.py b/deep_quoridor/src/v2/common.py index 9637cd5d..bb93091b 100644 --- a/deep_quoridor/src/v2/common.py +++ b/deep_quoridor/src/v2/common.py @@ -129,6 +129,16 @@ def alphazero_params_dict_from_config( "mcts_ucb_c": config.alphazero.mcts_c_puct, } + if config.training.initial_model: + im = config.training.initial_model + if im.file: + params_dict["model_filename"] = im.file + if im.wandb_alias: + params_dict["wandb_alias"] = im.wandb_alias + params_dict["wandb_project"] = im.wandb_project or ( + config.wandb.project if config.wandb else "deep_quoridor" + ) + # Add network config if config.alphazero.network.type == "mlp": params_dict.update( diff --git a/deep_quoridor/src/v2/config.py b/deep_quoridor/src/v2/config.py index 46dcd362..7f0deeab 100644 --- a/deep_quoridor/src/v2/config.py +++ b/deep_quoridor/src/v2/config.py @@ -71,6 +71,19 @@ class SelfPlayConfig(StrictBaseModel): rust_selfplay_binary: Optional[str] = None +class InitialModel(StrictBaseModel): + file: Optional[str] = None + wandb_project: Optional[str] = None + wandb_alias: Optional[str] = None + + @field_validator("wandb_alias") + @classmethod + def file_and_wandb_mutually_exclusive(cls, v, info): + if v is not None and info.data.get("file") is not None: + raise ValueError("Cannot specify both 'file' and 'wandb_alias' in initial_model") + return v + + class TrainingConfig(StrictBaseModel): games_per_training_step: float learning_rate: float @@ -80,6 +93,7 @@ class TrainingConfig(StrictBaseModel): model_save_timing: bool = False save_onnx: bool = False finish_after: Optional[str] = None + initial_model: Optional[InitialModel] = None class TournamentBenchmarkConfig(StrictBaseModel): diff --git a/deep_quoridor/src/v2/trainer.py b/deep_quoridor/src/v2/trainer.py index 1f14f592..63f87f50 100644 --- a/deep_quoridor/src/v2/trainer.py +++ b/deep_quoridor/src/v2/trainer.py @@ -31,7 +31,6 @@ def model_uploader(config: Config, every: str, model_id: str, wandb_run, shutdow def train(config: Config): batch_size = config.training.batch_size - alphazero_agent = create_alphazero(config, config.self_play.alphazero, overrides={"training_mode": True}) upload_model_thread = None