mirror of
https://github.com/jcreek/CosmicClash.git
synced 2026-09-11 08:23:45 +00:00
feat(*): Add self-play RL training pipeline with PPO trainer, in-game GDScript policy inference, and bot opponent support in Match mode
This commit is contained in:
@@ -0,0 +1,105 @@
|
||||
"""Train the Cosmic Clash self-play PPO policy.
|
||||
|
||||
Example (smoke run):
|
||||
.venv/bin/python train.py --experiment smoke --timesteps 100000
|
||||
|
||||
Long run on the Linux/CUDA box:
|
||||
GODOT_BIN=~/godot/Godot_v4.7.1-stable_linux.x86_64 \
|
||||
.venv/bin/python train.py --experiment run01 --timesteps 20000000 \
|
||||
--n-parallel 6 --speedup 16
|
||||
|
||||
See TRAINING.md at the repo root for the full workflow.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import pathlib
|
||||
|
||||
from stable_baselines3 import PPO
|
||||
from stable_baselines3.common.callbacks import CheckpointCallback
|
||||
from stable_baselines3.common.vec_env.vec_monitor import VecMonitor
|
||||
|
||||
from cosmic_env import CosmicClashVecEnv
|
||||
|
||||
TRAINING_DIR = pathlib.Path(__file__).resolve().parent
|
||||
DEFAULT_GODOT_MACOS = "/Applications/Godot.app/Contents/MacOS/Godot"
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument(
|
||||
"--godot_bin",
|
||||
default=os.environ.get("GODOT_BIN", DEFAULT_GODOT_MACOS),
|
||||
help="Path to the Godot binary (or set GODOT_BIN)",
|
||||
)
|
||||
parser.add_argument("--experiment", default="default", help="Run name for logs/checkpoints")
|
||||
parser.add_argument("--timesteps", type=int, default=200_000)
|
||||
parser.add_argument("--n-parallel", type=int, default=2, help="Parallel Godot instances (2 agents each)")
|
||||
parser.add_argument("--speedup", type=int, default=8, help="Physics speedup factor inside Godot")
|
||||
parser.add_argument("--port", type=int, default=11008, help="Base TCP port (one per instance)")
|
||||
parser.add_argument("--seed", type=int, default=0)
|
||||
parser.add_argument("--resume", default=None, help="Checkpoint .zip to resume from")
|
||||
parser.add_argument("--checkpoint-every", type=int, default=100_000, help="Timesteps between checkpoints")
|
||||
parser.add_argument("--viz", action="store_true", help="Show game windows (debugging; slow)")
|
||||
parser.add_argument("--wandb", action="store_true", help="Also log to Weights & Biases")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
log_dir = TRAINING_DIR / "logs"
|
||||
checkpoint_dir = TRAINING_DIR / "checkpoints" / args.experiment
|
||||
checkpoint_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
if args.wandb:
|
||||
import wandb
|
||||
|
||||
wandb.init(project="cosmic-clash-rl", name=args.experiment, sync_tensorboard=True)
|
||||
|
||||
env = CosmicClashVecEnv(
|
||||
godot_bin=args.godot_bin,
|
||||
n_parallel=args.n_parallel,
|
||||
seed=args.seed,
|
||||
port=args.port,
|
||||
show_window=args.viz,
|
||||
speedup=args.speedup,
|
||||
)
|
||||
env = VecMonitor(env)
|
||||
|
||||
if args.resume:
|
||||
model = PPO.load(args.resume, env=env, tensorboard_log=str(log_dir))
|
||||
print(f"Resumed from {args.resume} at {model.num_timesteps} timesteps")
|
||||
else:
|
||||
model = PPO(
|
||||
"MultiInputPolicy",
|
||||
env,
|
||||
verbose=1,
|
||||
ent_coef=0.0001,
|
||||
n_steps=256,
|
||||
batch_size=256,
|
||||
learning_rate=3e-4,
|
||||
tensorboard_log=str(log_dir),
|
||||
)
|
||||
|
||||
checkpoint_callback = CheckpointCallback(
|
||||
save_freq=max(args.checkpoint_every // env.num_envs, 1),
|
||||
save_path=str(checkpoint_dir),
|
||||
name_prefix="ppo",
|
||||
)
|
||||
|
||||
try:
|
||||
model.learn(
|
||||
args.timesteps,
|
||||
callback=checkpoint_callback,
|
||||
tb_log_name=args.experiment,
|
||||
reset_num_timesteps=not args.resume,
|
||||
)
|
||||
finally:
|
||||
final_path = checkpoint_dir / "final.zip"
|
||||
model.save(str(final_path))
|
||||
print(f"Saved {final_path}")
|
||||
env.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user