402c20dfb5
rewards adjustment pybullet logic contained in SimManager
236 lines
9.5 KiB
Python
236 lines
9.5 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)
|
|
|
|
policy_kwargs = dict(
|
|
log_std_init=-1.5, # Sets initial std ~ 0.22 instead of 1.0
|
|
net_arch=dict(pi=[256, 256], vf=[256, 256])
|
|
)
|
|
|
|
model = PPO(
|
|
"MlpPolicy",
|
|
env,
|
|
verbose=1,
|
|
seed=seed,
|
|
learning_rate=1.5e-4, # Cut LR in half (from 3e-4) to smooth out updates
|
|
n_steps=1024, # 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,
|
|
policy_kwargs=policy_kwargs,
|
|
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,
|
|
) |