"""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()