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). Saves inside the matching PPO_X subfolder as created by TensorBoard. """ def __init__(self, save_path: str, name_prefix: str = "ppo_jackbot", step_interval: int = 100_000, verbose: int = 1): super().__init__(verbose) self.base_save_path = save_path self.run_save_path = save_path self.name_prefix = name_prefix self.step_interval = step_interval self.last_milestone = 0 def _on_training_start(self) -> None: """Executed right before training loop starts. Resolves TensorBoard's run folder name (e.g. PPO_1).""" if self.logger and self.logger.dir: run_folder_name = Path(self.logger.dir).name # Extracts "PPO_1", "PPO_2", etc. self.run_save_path = os.path.join(self.base_save_path, run_folder_name) else: self.run_save_path = self.base_save_path os.makedirs(self.run_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.run_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: return True def _on_rollout_end(self) -> bool: 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) 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, 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 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) policy_kwargs = dict( log_std_init=-1.5, net_arch=dict(pi=[256, 256], vf=[256, 256]) ) model = PPO( "MlpPolicy", env, verbose=1, seed=seed, learning_rate=1.5e-4, n_steps=256, batch_size=256, n_epochs=10, gamma=0.99, gae_lambda=0.95, clip_range=0.2, target_kl=0.03, ent_coef=0.03, vf_coef=0.5, max_grad_norm=0.5, device=device, policy_kwargs=policy_kwargs, tensorboard_log=str(Path(__file__).resolve().parent / "tensorboard"), ) # Base folder where model runs will be stored 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]) # Save the final model inside the matching PPO_X directory as well if model.logger and model.logger.dir: run_folder_name = Path(model.logger.dir).name final_dir = Path(model_path).parent / run_folder_name else: final_dir = Path(model_path).parent final_dir.mkdir(parents=True, exist_ok=True) final_save_path = final_dir / f"{model_prefix}_final.zip" model.save(str(final_save_path)) if model.verbose > 0: print(f"[Training Complete] Saved final model to -> {final_save_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, )