feat(training): PPO training script with vectorized LanderCliEnv
This commit is contained in:
96
Training/train.py
Normal file
96
Training/train.py
Normal file
@@ -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()
|
||||
Reference in New Issue
Block a user