Reworked training

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)
This commit is contained in:
2026-08-03 22:27:26 +02:00
parent 5317ef1299
commit acb3d671be
8 changed files with 585 additions and 163 deletions
+81 -8
View File
@@ -1,6 +1,8 @@
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
@@ -37,6 +39,67 @@ class MilestoneCheckpointCallback(BaseCallback):
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.")
@@ -44,19 +107,18 @@ def parse_args():
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("--use-gui", action="store_true", help="Enable PyBullet GUI during training")
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(num_robots, robot_spacing, start_pose, use_gui, rank, seed=0):
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,
num_robots=num_robots,
robot_spacing=robot_spacing,
start_pose=start_pose,
)
@@ -71,7 +133,6 @@ def train(
seed: int = 0,
device: str = "auto",
use_gui: bool = False,
num_robots: int = 1,
num_workers: int = 8,
robot_spacing: float = 0.5,
start_pose: str = "init_deg",
@@ -84,13 +145,13 @@ def train(
# Create multi-process vector environment
if num_workers > 1:
env_fns = [
make_env(num_robots, robot_spacing, start_pose, use_gui, rank=i, seed=seed)
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(num_robots, robot_spacing, start_pose, use_gui, rank=0, seed=seed)
make_env(robot_spacing, start_pose, use_gui, rank=0, seed=seed)
])
def resolve_device(requested_device: str) -> str:
@@ -122,6 +183,17 @@ def train(
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"),
)
@@ -135,7 +207,9 @@ def train(
step_interval=100_000
)
model.learn(total_timesteps=total_timesteps, callback=milestone_cb)
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)
@@ -150,7 +224,6 @@ if __name__ == "__main__":
seed=args.seed,
device=args.device,
use_gui=args.use_gui,
num_robots=args.num_robots,
num_workers=args.num_workers,
robot_spacing=args.robot_spacing,
start_pose=args.start_pose,