acb3d671be
new reward/penalty system learning phases with curriculum learning new training parameters cleanup of old code better logging while training multiple environments instead of robots (they could bumb into each other)
230 lines
9.4 KiB
Python
230 lines
9.4 KiB
Python
import os
|
|
import argparse
|
|
from pathlib import Path
|
|
|
|
import numpy as np
|
|
from stable_baselines3.common.callbacks import BaseCallback
|
|
from stable_baselines3.common.vec_env import SubprocVecEnv, DummyVecEnv
|
|
from .env import JackBotEnv
|
|
|
|
|
|
class MilestoneCheckpointCallback(BaseCallback):
|
|
"""
|
|
Saves a model checkpoint the FIRST time total_timesteps
|
|
crosses every multiple of step_interval (e.g., 100,000).
|
|
"""
|
|
def __init__(self, save_path: str, name_prefix: str = "ppo_jackbot", step_interval: int = 100_000, verbose: int = 1):
|
|
super().__init__(verbose)
|
|
self.save_path = save_path
|
|
self.name_prefix = name_prefix
|
|
self.step_interval = step_interval
|
|
self.last_milestone = 0
|
|
os.makedirs(self.save_path, exist_ok=True)
|
|
|
|
def _on_step(self) -> bool:
|
|
current_milestone = self.num_timesteps // self.step_interval
|
|
|
|
if current_milestone > self.last_milestone:
|
|
self.last_milestone = current_milestone
|
|
milestone_step = current_milestone * self.step_interval
|
|
|
|
save_file = os.path.join(
|
|
self.save_path,
|
|
f"{self.name_prefix}_{milestone_step}_steps.zip"
|
|
)
|
|
self.model.save(save_file)
|
|
|
|
if self.verbose > 0:
|
|
print(f"\n[Checkpoint] Saved milestone model at {self.num_timesteps} steps -> {save_file}\n")
|
|
|
|
return True
|
|
|
|
class JackBotMetricsCallback(BaseCallback):
|
|
"""
|
|
Tracks the best current alive-robot performance for the most recent rollout,
|
|
instead of logging lifetime maxima from the entire training run.
|
|
"""
|
|
def __init__(self, verbose=0):
|
|
super().__init__(verbose)
|
|
self.best_alive_speed = 0.0
|
|
self.best_alive_yaw_rate = 0.0
|
|
self.best_alive_distance = 0.0
|
|
self.best_alive_survival_steps = 0.0
|
|
self.best_alive_reward = -float('inf')
|
|
|
|
def _on_step(self) -> bool:
|
|
"""Required by SB3 BaseCallback; no-op here because the rollout summary is emitted at rollout end."""
|
|
return True
|
|
|
|
def _on_rollout_end(self) -> bool:
|
|
"""Executed right before PPO outputs the log table to console."""
|
|
try:
|
|
vec_env = self.training_env
|
|
alive_metrics = vec_env.env_method("get_current_robot_metrics")
|
|
|
|
self.best_alive_speed = 0.0
|
|
self.best_alive_yaw_rate = 0.0
|
|
self.best_alive_distance = 0.0
|
|
self.best_alive_survival_steps = 0.0
|
|
self.best_alive_reward = -float('inf')
|
|
|
|
for worker_res in alive_metrics:
|
|
for metrics in worker_res:
|
|
if not metrics.get("alive", False):
|
|
continue
|
|
|
|
if metrics["reward"] > self.best_alive_reward:
|
|
self.best_alive_reward = float(metrics["reward"])
|
|
if metrics["speed"] > self.best_alive_speed:
|
|
self.best_alive_speed = float(metrics["speed"])
|
|
if metrics["yaw_rate"] > self.best_alive_yaw_rate:
|
|
self.best_alive_yaw_rate = float(metrics["yaw_rate"])
|
|
if metrics["distance_from_start"] > self.best_alive_distance:
|
|
self.best_alive_distance = float(metrics["distance_from_start"])
|
|
if metrics["survival_steps"] > self.best_alive_survival_steps:
|
|
self.best_alive_survival_steps = float(metrics["survival_steps"])
|
|
|
|
self.logger.record("custom/best_alive_speed_mps", float(self.best_alive_speed))
|
|
self.logger.record("custom/best_alive_yaw_rate_rads", float(self.best_alive_yaw_rate))
|
|
self.logger.record("custom/best_alive_distance_from_start_m", float(self.best_alive_distance))
|
|
self.logger.record("custom/best_alive_survival_steps", float(self.best_alive_survival_steps))
|
|
self.logger.record("custom/best_alive_reward", float(self.best_alive_reward) if np.isfinite(self.best_alive_reward) else 0.0)
|
|
|
|
# Backward-compatible aliases so old dashboards keep a stable field name.
|
|
self.logger.record("custom/max_speed_mps", float(self.best_alive_speed))
|
|
self.logger.record("custom/max_yaw_rate_rads", float(self.best_alive_yaw_rate))
|
|
self.logger.record("custom/max_distance_from_start_m", float(self.best_alive_distance))
|
|
self.logger.record("custom/max_survival_steps", float(self.best_alive_survival_steps))
|
|
self.logger.record("custom/best_episode_reward", float(self.best_alive_reward) if np.isfinite(self.best_alive_reward) else 0.0)
|
|
except Exception:
|
|
pass
|
|
|
|
return True
|
|
|
|
def parse_args():
|
|
parser = argparse.ArgumentParser(description="Train a joint-command policy for JackBot.")
|
|
parser.add_argument("--timesteps", type=int, default=500_000, help="Total training timesteps")
|
|
parser.add_argument("--model-path", type=str, default="ml/checkpoints/ppo_joint_command", help="Where to save the trained model")
|
|
parser.add_argument("--seed", type=int, default=0, help="Random seed")
|
|
parser.add_argument("--device", type=str, default="auto", help="Device to use: 'cpu', 'cuda', or 'auto' to autodetect")
|
|
parser.add_argument("--gui", "--use-gui", dest="use_gui", action="store_true", help="Enable PyBullet GUI during training")
|
|
parser.add_argument("--num-workers", type=int, default=8, help="Number of parallel CPU worker processes")
|
|
parser.add_argument("--robot-spacing", type=float, default=3.0, help="Spacing between robots in meters")
|
|
parser.add_argument("--start-pose", type=str, choices=["init_deg", "init90_deg"], default="init_deg", help="Initial robot pose at reset")
|
|
return parser.parse_args()
|
|
|
|
|
|
def make_env(robot_spacing, start_pose, use_gui, rank, seed=0):
|
|
def _init():
|
|
env = JackBotEnv(
|
|
use_gui=use_gui if rank == 0 else False, # Only rank 0 gets GUI if requested
|
|
random_command=True,
|
|
robot_spacing=robot_spacing,
|
|
start_pose=start_pose,
|
|
)
|
|
env.reset(seed=seed + rank)
|
|
return env
|
|
return _init
|
|
|
|
|
|
def train(
|
|
total_timesteps: int,
|
|
model_path: str,
|
|
seed: int = 0,
|
|
device: str = "auto",
|
|
use_gui: bool = False,
|
|
num_workers: int = 8,
|
|
robot_spacing: float = 0.5,
|
|
start_pose: str = "init_deg",
|
|
):
|
|
try:
|
|
from stable_baselines3 import PPO
|
|
except ImportError as exc:
|
|
raise ImportError("stable-baselines3 is required. Install with: pip install stable-baselines3") from exc
|
|
|
|
# Create multi-process vector environment
|
|
if num_workers > 1:
|
|
env_fns = [
|
|
make_env(robot_spacing, start_pose, use_gui, rank=i, seed=seed)
|
|
for i in range(num_workers)
|
|
]
|
|
env = SubprocVecEnv(env_fns)
|
|
else:
|
|
env = DummyVecEnv([
|
|
make_env(robot_spacing, start_pose, use_gui, rank=0, seed=seed)
|
|
])
|
|
|
|
def resolve_device(requested_device: str) -> str:
|
|
try:
|
|
import torch
|
|
except ImportError:
|
|
if requested_device != "cpu":
|
|
raise RuntimeError("PyTorch is not installed.")
|
|
return "cpu"
|
|
|
|
hip_supported = getattr(torch.version, "hip", None) is not None
|
|
cuda_available = torch.cuda.is_available()
|
|
hip_available = hip_supported and getattr(torch.backends, "hip", None) is not None and torch.backends.hip.is_available()
|
|
|
|
if requested_device == "auto":
|
|
return "cuda" if (hip_available or cuda_available) else "cpu"
|
|
|
|
if requested_device in {"cuda", "gpu", "hip"}:
|
|
if hip_available or cuda_available:
|
|
return "cuda"
|
|
raise RuntimeError(f"GPU requested ({requested_device}) but not available.")
|
|
|
|
return "cpu"
|
|
|
|
device = resolve_device(device)
|
|
|
|
model = PPO(
|
|
"MlpPolicy",
|
|
env,
|
|
verbose=1,
|
|
seed=seed,
|
|
learning_rate=3.5e-4, # Cut LR in half (from 3e-4) to smooth out updates
|
|
n_steps=2048, # Larger rollout buffer per env for stable gradients
|
|
batch_size=128, # Larger minibatches reduce noise
|
|
n_epochs=10, # Number of epoch updates per rollout
|
|
gamma=0.99, # Discount factor
|
|
gae_lambda=0.95, # GAE smoothing
|
|
clip_range=0.2, # Standard PPO clipping
|
|
target_kl=0.03, # EARLY STOPPING: Halts policy update if KL > 0.015!
|
|
ent_coef=0.03, # Entropy coefficient to encourage exploration
|
|
vf_coef=0.5,
|
|
max_grad_norm=0.5,
|
|
device=device,
|
|
tensorboard_log=str(Path(__file__).resolve().parent / "tensorboard"),
|
|
)
|
|
|
|
save_dir = str(Path(model_path).parent)
|
|
model_prefix = Path(model_path).stem
|
|
|
|
milestone_cb = MilestoneCheckpointCallback(
|
|
save_path=save_dir,
|
|
name_prefix=model_prefix,
|
|
step_interval=100_000
|
|
)
|
|
|
|
metrics_callback = JackBotMetricsCallback()
|
|
|
|
model.learn(total_timesteps=total_timesteps, callback=[milestone_cb, metrics_callback])
|
|
|
|
Path(model_path).parent.mkdir(parents=True, exist_ok=True)
|
|
model.save(model_path)
|
|
env.close()
|
|
|
|
|
|
if __name__ == "__main__":
|
|
args = parse_args()
|
|
train(
|
|
args.timesteps,
|
|
args.model_path,
|
|
seed=args.seed,
|
|
device=args.device,
|
|
use_gui=args.use_gui,
|
|
num_workers=args.num_workers,
|
|
robot_spacing=args.robot_spacing,
|
|
start_pose=args.start_pose,
|
|
) |