diff --git a/Training/train.py b/Training/train.py new file mode 100644 index 0000000..14821e2 --- /dev/null +++ b/Training/train.py @@ -0,0 +1,96 @@ +"""Train a PPO policy on LanderCliEnv. + +Usage: + python train.py --steps 100000 # short smoke run + python train.py --steps 2000000 --n-envs 8 # full training +""" +from __future__ import annotations + +import argparse +from pathlib import Path + +from stable_baselines3 import PPO +from stable_baselines3.common.callbacks import CheckpointCallback +from stable_baselines3.common.vec_env import SubprocVecEnv, DummyVecEnv + +# sys.path shim so this script can be run from any CWD. +import sys +sys.path.insert(0, str(Path(__file__).resolve().parent)) + +from lander_cli_env import LanderCliEnv + + +REPO_ROOT = Path(__file__).resolve().parents[1] +CLI_BINARY = REPO_ROOT / "publish" / "GameCli" / "GameCli" + + +def make_env(seed: int, max_episode_steps: int): + """Return a thunk that SubprocVecEnv can call to construct one env.""" + def _fn(): + env = LanderCliEnv( + cli_binary=str(CLI_BINARY), + max_episode_steps=max_episode_steps, + ) + env.reset(seed=seed) + return env + return _fn + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--steps", type=int, default=2_000_000, + help="total env steps to train for") + parser.add_argument("--n-envs", type=int, default=8, + help="parallel envs; each spawns one GameCli subprocess") + parser.add_argument("--max-episode-steps", type=int, default=1000) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--checkpoint-dir", type=Path, + default=REPO_ROOT / "checkpoints") + parser.add_argument("--tb-dir", type=Path, + default=REPO_ROOT / "tensorboard") + parser.add_argument("--use-dummy", action="store_true", + help="run envs in-process (slower, easier to debug)") + args = parser.parse_args() + + if not CLI_BINARY.exists(): + raise SystemExit( + f"CLI binary not found at {CLI_BINARY}. Run: " + f"dotnet publish GameCli/GameCli.csproj -c Release -o publish/GameCli" + ) + + args.checkpoint_dir.mkdir(parents=True, exist_ok=True) + args.tb_dir.mkdir(parents=True, exist_ok=True) + + env_fns = [ + make_env(seed=args.seed + i, max_episode_steps=args.max_episode_steps) + for i in range(args.n_envs) + ] + vec_env_cls = DummyVecEnv if args.use_dummy else SubprocVecEnv + vec_env = vec_env_cls(env_fns) + + model = PPO( + "MlpPolicy", + vec_env, + verbose=1, + seed=args.seed, + tensorboard_log=str(args.tb_dir), + policy_kwargs=dict(net_arch=[64, 64]), + ) + + checkpoint_cb = CheckpointCallback( + save_freq=max(args.steps // 10, 1) // args.n_envs, + save_path=str(args.checkpoint_dir), + name_prefix="ppo_lander", + ) + + try: + model.learn(total_timesteps=args.steps, callback=checkpoint_cb) + final_path = args.checkpoint_dir / "ppo_lander_final.zip" + model.save(str(final_path)) + print(f"saved final model to {final_path}") + finally: + vec_env.close() + + +if __name__ == "__main__": + main()