#!/usr/bin/env python3
"""
PlateBalance RL: Advanced MuJoCo + PyTorch PPO Object Balancing System
A high-performance MuJoCo + PyTorch PPO reinforcement learning system where
a 2-axis tilting plate learns to balance 21 diverse rigid and multi-body objects:
- Standard Shapes: sphere, disk, egg, hollow cup, coin, stick, tall block,
triangular prism, cube block, puck.
- Harder Challenge Shapes: cone, horizontal capsule, ramp wedge, tetrahedron,
long flat bar, cross/plus, asymmetric L-shape, wide tile block, heavy bowling ball,
and off-center mass block.
- Multi-Body Crumbling Cookie: A breakable cookie with 5 independent physical crumb
fragments that scatter and slide independently, challenging the agent to keep
EVERY individual crumb on the plate!
Installation:
pip install -U torch numpy mujoco imageio imageio-ffmpeg
Usage:
# Train policy on all objects (default)
python balance_plate_rl.py
# Train or evaluate on a specific object (e.g., crumbling cookie or cone)
python balance_plate_rl.py --object cookie
python balance_plate_rl.py --object cone
# Resume training from latest checkpoint
python balance_plate_rl.py --resume
# Run full benchmark evaluation across all 21 objects
python balance_plate_rl.py --eval --checkpoint runs/plate_balance_v1/checkpoints/latest
# Interactive real-time 3D viewer (watch the cookie crumble and balance in real-time)
python balance_plate_rl.py --human-view --object cookie --checkpoint runs/plate_balance_v1/checkpoints/latest
# Record evaluation video
python balance_plate_rl.py --record-video --checkpoint runs/plate_balance_v1/checkpoints/latest
"""
from __future__ import annotations
import argparse
import json
import math
import os
import random
import re
import sys
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
import numpy as np
import torch
import torch.nn as nn
from torch.distributions import Normal
try:
import mujoco
except ImportError as exc:
raise SystemExit("MuJoCo is required. Install with: pip install -U mujoco") from exc
try:
import imageio.v2 as imageio
except ImportError:
try:
import imageio
except ImportError:
imageio = None # Optional for headless training/eval
# =============================================================================
# CONFIGURATION — DEFAULT VALUES
# =============================================================================
# ----------------------------- Experiment -----------------------------------
SEED = 42
EXPERIMENT_NAME = "plate_balance_v1"
OUTPUT_DIR = "runs"
RESUME = True
RESUME_PATH = r"runs\plate_balance_v1\checkpoints\step_03825064_final" # Empty = automatically use latest checkpoint.
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
# ----------------------------- Architecture ---------------------------------
OBSERVATION_SIZE = 64
HIDDEN_SIZE = 128 # First/second hidden layer width.
INTERMEDIATE_SIZE = 128 # Third hidden layer width.
BOTTLENECK_SIZE = 64 # Fourth hidden layer width.
ACTION_SIZE = 2 # Roll torque, Pitch torque
# Architecture: 64 -> 192 -> 192 -> 64 -> {actor: 2, critic: 1}
ACTOR_LOG_STD_INIT = -0.75
ORTHOGONAL_INIT = True
# ----------------------------- PPO ------------------------------------------
NUM_ENVS = 8
ROLLOUT_STEPS = 512
PPO_EPOCHS = 6
MINIBATCH_SIZE = 2048
GAMMA = 0.99
GAE_LAMBDA = 0.95
CLIP_COEF = 0.20
VALUE_COEF = 0.50
ENTROPY_COEF = 0.005
MAX_GRAD_NORM = 0.50
LEARNING_RATE = 3e-4
ADAM_EPS = 1e-5
ANNEAL_LR = True # Linearly anneal learning rate to 0.
# ----------------------------- Training -------------------------------------
NUM_EPISODES = 30_000
MAX_EPISODE_STEPS = 10_000
LOGGING_STEPS = 32240
SAVE_STEPS = 1_000_000
# ----------------------------- Video / Visualization ------------------------
VISUALIZE = True # True = save rollout videos periodically during training.
VIDEO_EVERY_STEPS = 1_000_000
VIDEO_LENGTH_STEPS = 3000
VIDEO_FPS = 60
RENDER_WIDTH = 640
RENDER_HEIGHT = 480
HUMAN_VIEW = False # Optional interactive MuJoCo viewer.
# ----------------------------- Physics & Environment ------------------------
PHYSICS_TIMESTEP = 0.0025 # 400 Hz physics simulation.
CONTROL_DECIMATION = 4 # 4 physics sub-steps per control step => 100 Hz RL control.
GRAVITY = 9.81
PLATE_HALF_SIZE = 0.90 # Half-width and half-length of square plate (meters).
PLATE_THICKNESS = 0.08 # Half-thickness in z (box geom size z = 0.08).
PLATE_FRICTION = (1.0, 0.02, 0.002) # Sliding, torsional, and rolling friction.
PLATE_DAMPING = 0.08
MAX_PLATE_ANGLE = math.radians(15.0) # Maximum tilt angle in radians (~0.2618 rad).
MAX_PLATE_TORQUE = 15.0 # Maximum plate torque in N*m.
OBJECT_START_HEIGHT = 0.02 # Gentle initial clearance above plate when spawning (2cm).
OBJECT_START_POS_RANGE = 0.35 # Spawn xy radius from plate center.
OBJECT_START_VELOCITY = 0.15 # Initial linear velocity standard deviation.
OBJECT_START_ANGULAR_VELOCITY = 0.20 # Initial angular velocity standard deviation.
# Curriculum & Domain Randomization
RANDOMIZE_MASS = True
RANDOMIZE_FRICTION = True
RANDOMIZE_OBJECT_SIZE = True
MASS_LOG_RANGE = (0.50, 2.00) # Multiplicative log-uniform mass range.
FRICTION_RANGE = (0.45, 1.25) # Multiplicative friction coefficient range.
SIZE_RANGE = (0.85, 1.15) # Multiplicative linear geometry scale range.
RANDOMIZE_PLATE_ANGLE = True
INITIAL_PLATE_ANGLE_RANGE = math.radians(3.0)
# Reward Shaping
CENTER_REWARD_SCALE = 2.0 # Reward for centering object on plate.
VELOCITY_REWARD_SCALE = 0.40 # Reward for damping object linear velocity.
ANGLE_PENALTY_SCALE = 0.20 # Penalty for excessive plate tilt.
ACTION_PENALTY_SCALE = 0.005 # Penalty for excessive torque commands.
ACTION_RATE_PENALTY_SCALE = 0.01 # Penalty for rapid actuator jerk (smooth control).
EDGE_PENALTY_SCALE = 0.50 # Progressive penalty near plate perimeter.
FALL_PENALTY = -10.0 # Terminal penalty when object falls off.
SURVIVAL_REWARD = 0.05 # Step reward for keeping object on plate.
# Fall detection margins
FALL_MARGIN = 0.10 # Beyond plate half-size + margin triggers fall.
FALL_Z_THRESHOLD = 0.90 # Height below which object is deemed fallen.
# ----------------------------- Object Classification -----------------------
STANDARD_OBJECT_TYPES = [
"sphere",
"disk",
"egg",
"cup",
"coin",
"stick",
"tall",
"triangle",
"block",
"puck",
]
HARDER_OBJECT_TYPES = [
"cone",
"capsule",
"wedge",
"tetra",
"flat_bar",
"cross",
"L_shape",
"wide_block",
"heavy_ball",
"offcenter_block",
]
MULTI_BODY_OBJECT_TYPES = [
"cookie", # Multi-body crumbling cookie with 5 independent crumb fragments
]
OBJECT_TYPES = STANDARD_OBJECT_TYPES + HARDER_OBJECT_TYPES + MULTI_BODY_OBJECT_TYPES
EVAL_OBJECT_TYPES = OBJECT_TYPES.copy()
PRINT_OBJECT_COUNTS = True
PRINT_HYPERPARAMS = True
# =============================================================================
# UTILITIES & REPRODUCIBILITY
# =============================================================================
def set_seed(seed: int) -> None:
"""Set random seeds across Python, NumPy, and PyTorch for reproducibility."""
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
def extract_step_number(path: Path) -> int:
"""Safely extract the integer step number from checkpoint folder names."""
matches = re.findall(r"\d+", path.name)
if matches:
return int(matches[0])
return -1
def latest_checkpoint(root: Path) -> Optional[Path]:
"""Find the latest checkpoint directory in a robust, crash-free manner."""
if not root.exists():
return None
latest_alias = root / "latest"
if latest_alias.exists() and (latest_alias / "model.pt").exists():
return latest_alias
candidates = [p for p in root.glob("step_*") if p.is_dir() and (p / "model.pt").exists()]
if not candidates:
return None
candidates.sort(key=extract_step_number)
return candidates[-1]
def safe_float(x: Any) -> float:
return float(np.asarray(x).item())
# =============================================================================
# INERTIA CALCULATIONS & PHYSICAL DEFINITIONS
# =============================================================================
OBJECT_INFO: Dict[str, Dict[str, float]] = {
# Baseline 10 Shapes
"sphere": {"radius": 0.16, "mass": 0.55},
"disk": {"radius": 0.22, "height": 0.055, "mass": 0.60},
"egg": {"radius": 0.17, "height": 0.34, "mass": 0.58},
"cup": {"radius": 0.18, "height": 0.25, "mass": 0.48},
"coin": {"radius": 0.18, "height": 0.025, "mass": 0.25},
"stick": {"radius": 0.035, "length": 0.48, "mass": 0.22},
"tall": {"radius": 0.045, "height": 0.52, "mass": 0.35},
"triangle": {"size": 0.34, "height": 0.16, "mass": 0.50},
"block": {"size": 0.28, "mass": 0.70},
"puck": {"radius": 0.25, "height": 0.08, "mass": 0.70},
# Harder 10 Shapes
"cone": {"radius": 0.16, "height": 0.32, "mass": 0.50},
"capsule": {"radius": 0.065, "length": 0.36, "mass": 0.45},
"wedge": {"size": 0.36, "height": 0.20, "mass": 0.55},
"tetra": {"size": 0.36, "height": 0.30, "mass": 0.40},
"flat_bar": {"length": 0.76, "height": 0.04, "mass": 0.60},
"cross": {"size": 0.52, "height": 0.08, "mass": 0.65},
"L_shape": {"size": 0.32, "height": 0.08, "mass": 0.60},
"wide_block": {"size": 0.64, "height": 0.04, "mass": 0.75},
"heavy_ball": {"radius": 0.16, "mass": 3.50},
"offcenter_block": {"size": 0.36, "height": 0.20, "mass": 0.75},
# Multi-Body Crumbling Cookie Challenge (5 Crumb Fragments)
"cookie": {"radius": 0.22, "height": 0.04, "mass": 0.45},
}
def cube_inertia(m: float, sx: float, sy: float, sz: float) -> np.ndarray:
return np.array([
m * (sy * sy + sz * sz) / 12.0,
m * (sx * sx + sz * sz) / 12.0,
m * (sx * sx + sy * sy) / 12.0,
], dtype=np.float64)
def sphere_inertia(m: float, r: float) -> np.ndarray:
i = 0.4 * m * r * r
return np.array([i, i, i], dtype=np.float64)
def cylinder_inertia(m: float, r: float, h: float) -> np.ndarray:
axial = 0.5 * m * r * r
transverse = m * (3.0 * r * r + h * h) / 12.0
return np.array([transverse, transverse, axial], dtype=np.float64)
def capsule_inertia(m: float, r: float, length: float) -> np.ndarray:
rod = max(0.55 * m, 1e-6)
cyl = max(m - rod, 1e-6)
i_trans = rod * (length * length) / 12.0 + cyl * (3.0 * r * r + length * length) / 12.0
i_axial = 0.5 * cyl * r * r + rod * r * r / 2.0
return np.array([i_trans, i_trans, i_axial], dtype=np.float64)
# Custom 3D Surface Meshes
TRIANGLE_MESH = """
"""
CONE_MESH = """
"""
WEDGE_MESH = """
"""
TETRA_MESH = """
"""
def build_model_xml() -> str:
"""Build and return the comprehensive MuJoCo XML containing all 21 objects and cookie crumbs."""
spawn_z = 1.05 + PLATE_THICKNESS + OBJECT_START_HEIGHT
return f"""
{TRIANGLE_MESH}
{CONE_MESH}
{WEDGE_MESH}
{TETRA_MESH}
"""
# =============================================================================
# MUJOCO ENVIRONMENT (SINGLE & MULTI-BODY CRUMBLING COOKIE)
# =============================================================================
class PlateBalanceEnv:
"""
Robust single-instance MuJoCo environment supporting 21 single & multi-body
objects with accurate physical scaling and multi-crumb tracking.
"""
OBJECT_GEOM_NAMES = {
# Baseline 10
"sphere": ["g_sphere"],
"disk": ["g_disk"],
"egg": ["g_egg"],
"cup": ["g_cup_bottom", "g_cup_w1", "g_cup_w2", "g_cup_w3", "g_cup_w4"],
"coin": ["g_coin"],
"stick": ["g_stick"],
"tall": ["g_tall"],
"triangle": ["g_triangle"],
"block": ["g_block"],
"puck": ["g_puck"],
# Harder 10
"cone": ["g_cone"],
"capsule": ["g_capsule"],
"wedge": ["g_wedge"],
"tetra": ["g_tetra"],
"flat_bar": ["g_flat_bar"],
"cross": ["g_cross_1", "g_cross_2"],
"L_shape": ["g_lshape_1", "g_lshape_2"],
"wide_block": ["g_wide_block"],
"heavy_ball": ["g_heavy_ball"],
"offcenter_block": ["g_offcenter_block"],
}
def __init__(self, seed: int, render: bool = False, active_objects: Optional[List[str]] = None):
self.rng = np.random.default_rng(seed)
self.model = mujoco.MjModel.from_xml_string(build_model_xml())
self.data = mujoco.MjData(self.model)
self.model.opt.timestep = PHYSICS_TIMESTEP
self.render_enabled = render
self.renderer: Optional[mujoco.Renderer] = None
self.active_objects = list(active_objects) if active_objects else list(OBJECT_TYPES)
if render:
self.renderer = mujoco.Renderer(self.model, height=RENDER_HEIGHT, width=RENDER_WIDTH)
# Primary Single Object
self.qpos_object = mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_JOINT, "object_free")
self.object_qpos_addr = self.model.jnt_qposadr[self.qpos_object]
self.object_dof_addr = self.model.jnt_dofadr[self.qpos_object]
self.object_body_id = mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_BODY, "object")
# Plate Joints and Sites
self.plate_roll_joint = mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_JOINT, "plate_roll")
self.plate_pitch_joint = mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_JOINT, "plate_pitch")
self.plate_roll_qpos = self.model.jnt_qposadr[self.plate_roll_joint]
self.plate_pitch_qpos = self.model.jnt_qposadr[self.plate_pitch_joint]
self.plate_roll_dof = self.model.jnt_dofadr[self.plate_roll_joint]
self.plate_pitch_dof = self.model.jnt_dofadr[self.plate_pitch_joint]
self.plate_body_id = mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_BODY, "plate_y")
self.plate_center_site = mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_SITE, "plate_center")
self.plate_geom_id = mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_GEOM, "plate")
# Cookie Crumb Bodies (5 fragments)
self.crumb_body_ids = [mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_BODY, f"crumb_{i}") for i in range(5)]
self.crumb_qpos_addrs = [self.model.jnt_qposadr[mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_JOINT, f"crumb_free_{i}")] for i in range(5)]
self.crumb_dof_addrs = [self.model.jnt_dofadr[mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_JOINT, f"crumb_free_{i}")] for i in range(5)]
self.crumb_geom_ids = [mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_GEOM, f"g_crumb_{i}") for i in range(5)]
# Map candidate geoms
self.object_geom_ids = {
name: [mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_GEOM, g) for g in geoms]
for name, geoms in self.OBJECT_GEOM_NAMES.items()
}
# Cache baseline sizes and positions for physical scaling
self.base_geom_sizes: Dict[int, np.ndarray] = {}
self.base_geom_positions: Dict[int, np.ndarray] = {}
for ids in self.object_geom_ids.values():
for gid in ids:
self.base_geom_sizes[gid] = self.model.geom_size[gid].copy()
self.base_geom_positions[gid] = self.model.geom_pos[gid].copy()
for gid in self.crumb_geom_ids:
self.base_geom_sizes[gid] = self.model.geom_size[gid].copy()
self.base_geom_positions[gid] = self.model.geom_pos[gid].copy()
self.steps = 0
self.episode_reward = 0.0
self.current_object = "sphere"
self.object_mass = 0.55
self.object_scale = 1.0
self.base_friction = 0.90
self.last_action = np.zeros(ACTION_SIZE, dtype=np.float32)
def _set_object_geometry(self, object_name: str) -> None:
"""Activate the selected single or multi-body object and scale its physics."""
# Deactivate all single object geoms
for ids in self.object_geom_ids.values():
for gid in ids:
self.model.geom_size[gid, :] = 1e-6
self.model.geom_pos[gid, :] = np.array([0.0, 0.0, -999.0])
self.model.geom_rgba[gid, 3] = 0.0
# Deactivate all crumb geoms by default
for gid in self.crumb_geom_ids:
self.model.geom_size[gid, :] = 1e-6
self.model.geom_pos[gid, :] = np.array([0.0, 0.0, -999.0])
self.model.geom_rgba[gid, 3] = 0.0
s = self.object_scale
info = OBJECT_INFO[object_name]
m = self.object_mass
if object_name == "cookie":
# Activate 5 cookie crumb fragments
for i, gid in enumerate(self.crumb_geom_ids):
self.model.geom_friction[gid, 0] = self.base_friction
self.model.geom_rgba[gid, :] = np.array([0.82, 0.58, 0.32, 1.0])
self.model.geom_size[gid, :] = self.base_geom_sizes[gid] * s
self.model.geom_pos[gid, :] = self.base_geom_positions[gid] * s
self.model.body_mass[self.crumb_body_ids[i]] = (m / 5.0)
return
# Activate single candidate object
active = self.object_geom_ids[object_name]
rgba = np.array([0.35, 0.38, 0.45, 1.0]) if object_name == "heavy_ball" else np.array([0.95, 0.45, 0.16, 1.0])
for gid in active:
self.model.geom_friction[gid, 0] = self.base_friction
self.model.geom_rgba[gid, :] = rgba
self.model.geom_size[gid, :] = self.base_geom_sizes[gid] * s
self.model.geom_pos[gid, :] = self.base_geom_positions[gid] * s
# Reset body center of mass position (offset for offcenter_block)
if object_name == "offcenter_block":
self.model.body_ipos[self.object_body_id] = np.array([0.08 * s, 0.06 * s, -0.02 * s])
else:
self.model.body_ipos[self.object_body_id] = np.array([0.0, 0.0, 0.0])
# Compute accurate 3D moment of inertia tensor
if object_name in ("sphere", "heavy_ball"):
inertia = sphere_inertia(m, info["radius"] * s)
elif object_name in ("disk", "coin", "puck"):
inertia = cylinder_inertia(m, info["radius"] * s, info["height"] * s)
elif object_name == "egg":
rx = info["radius"] * 0.85 * s
rz = info["height"] * 0.55 * s
inertia = np.array([
m * (rx * rx + rz * rz) / 5.0,
m * (rx * rx + rz * rz) / 5.0,
m * (2.0 * rx * rx) / 5.0,
])
elif object_name == "cup":
inertia = cylinder_inertia(m, info["radius"] * s, info["height"] * s)
elif object_name in ("stick", "capsule"):
inertia = capsule_inertia(m, info["radius"] * s, info["length"] * s)
elif object_name == "tall":
side = 0.09 * s
height = info["height"] * s
inertia = cube_inertia(m, side, side, height)
elif object_name == "cone":
r = info["radius"] * s
h = info["height"] * s
i_trans = 0.6 * m * (0.25 * r * r + h * h)
i_axial = 0.3 * m * r * r
inertia = np.array([i_trans, i_trans, i_axial])
elif object_name in ("triangle", "wedge"):
side = info["size"] * s
height = info["height"] * s
inertia = cube_inertia(m, side, side, height) * 1.10
elif object_name == "tetra":
side = info["size"] * s
i_val = (m * side * side) / 20.0
inertia = np.array([i_val, i_val, i_val])
elif object_name == "flat_bar":
inertia = cube_inertia(m, info["length"] * s, 0.10 * s, info["height"] * s)
elif object_name == "cross":
side = info["size"] * s
h = info["height"] * s
inertia = cube_inertia(m, side, side, h) * 0.70
elif object_name == "L_shape":
side = info["size"] * s
h = info["height"] * s
inertia = cube_inertia(m, side, side, h) * 0.85
elif object_name == "wide_block":
side = info["size"] * s
h = info["height"] * s
inertia = cube_inertia(m, side, side, h)
elif object_name == "offcenter_block":
side = info["size"] * s
h = info["height"] * s
inertia = cube_inertia(m, side, side, h) * 1.20
else: # block
side = info["size"] * s
inertia = cube_inertia(m, side, side, side)
self.model.body_mass[self.object_body_id] = m
self.model.body_inertia[self.object_body_id, :] = np.maximum(inertia, 1e-6)
def reset(self, specific_object: Optional[str] = None) -> np.ndarray:
"""Reset the environment state for a new episode."""
if specific_object and specific_object in OBJECT_INFO:
self.current_object = specific_object
else:
self.current_object = str(self.rng.choice(self.active_objects))
info = OBJECT_INFO[self.current_object]
self.object_mass = info["mass"]
if RANDOMIZE_MASS:
mult = float(np.exp(self.rng.uniform(math.log(MASS_LOG_RANGE[0]), math.log(MASS_LOG_RANGE[1]))))
self.object_mass *= mult
self.object_scale = float(self.rng.uniform(*SIZE_RANGE)) if RANDOMIZE_OBJECT_SIZE else 1.0
self.base_friction = float(self.rng.uniform(*FRICTION_RANGE)) if RANDOMIZE_FRICTION else 0.90
self._set_object_geometry(self.current_object)
mujoco.mj_resetData(self.model, self.data)
# Initial plate tilt randomization
if RANDOMIZE_PLATE_ANGLE:
self.data.qpos[self.plate_roll_qpos] = self.rng.uniform(-INITIAL_PLATE_ANGLE_RANGE, INITIAL_PLATE_ANGLE_RANGE)
self.data.qpos[self.plate_pitch_qpos] = self.rng.uniform(-INITIAL_PLATE_ANGLE_RANGE, INITIAL_PLATE_ANGLE_RANGE)
else:
self.data.qpos[self.plate_roll_qpos] = 0.0
self.data.qpos[self.plate_pitch_qpos] = 0.0
x = float(self.rng.uniform(-OBJECT_START_POS_RANGE, OBJECT_START_POS_RANGE))
y = float(self.rng.uniform(-OBJECT_START_POS_RANGE, OBJECT_START_POS_RANGE))
if self.current_object == "cookie":
# Park single object far away
base_s = self.object_qpos_addr
self.data.qpos[base_s:base_s + 7] = np.array([0.0, 0.0, -999.0, 1.0, 0.0, 0.0, 0.0])
self.data.qvel[self.object_dof_addr:self.object_dof_addr + 6] = 0.0
# Spawn 5 cookie crumb fragments on the plate with scatter offsets & velocities
crumb_offsets = [
(0.0, 0.0),
(0.08 * self.object_scale, 0.0),
(-0.08 * self.object_scale, 0.0),
(0.0, 0.08 * self.object_scale),
(0.0, -0.08 * self.object_scale),
]
for i in range(5):
dx, dy = crumb_offsets[i]
dx += float(self.rng.uniform(-0.02, 0.02))
dy += float(self.rng.uniform(-0.02, 0.02))
c_addr = self.crumb_qpos_addrs[i]
c_dof = self.crumb_dof_addrs[i]
z_c = 1.05 + PLATE_THICKNESS + 0.02 + OBJECT_START_HEIGHT
self.data.qpos[c_addr:c_addr + 7] = np.array([x + dx, y + dy, z_c, 1.0, 0.0, 0.0, 0.0])
self.data.qvel[c_dof:c_dof + 6] = np.concatenate([
self.rng.normal(0.0, OBJECT_START_VELOCITY, 3),
self.rng.normal(0.0, OBJECT_START_ANGULAR_VELOCITY, 3),
])
else:
# Park all crumb fragments far away
for i in range(5):
c_addr = self.crumb_qpos_addrs[i]
c_dof = self.crumb_dof_addrs[i]
self.data.qpos[c_addr:c_addr + 7] = np.array([0.0, 0.0, -999.0, 1.0, 0.0, 0.0, 0.0])
self.data.qvel[c_dof:c_dof + 6] = 0.0
# Spawn single object with safe clearance
base = self.object_qpos_addr
if self.current_object in ("sphere", "heavy_ball"):
half_height = info["radius"] * self.object_scale
elif self.current_object in ("stick", "capsule", "flat_bar"):
half_height = (info["length"] / 2.0) * self.object_scale if "length" in info else (info["height"] / 2.0) * self.object_scale
elif self.current_object in ("block", "cross", "L_shape", "wide_block", "wedge", "tetra", "offcenter_block"):
half_height = (info.get("height", info.get("size", 0.20)) / 2.0) * self.object_scale
else:
half_height = (info["height"] / 2.0) * self.object_scale
z = 1.05 + PLATE_THICKNESS + half_height + OBJECT_START_HEIGHT
self.data.qpos[base:base + 7] = np.array([x, y, z, 1.0, 0.0, 0.0, 0.0])
self.data.qvel[self.object_dof_addr:self.object_dof_addr + 6] = np.concatenate([
self.rng.normal(0.0, OBJECT_START_VELOCITY, 3),
self.rng.normal(0.0, OBJECT_START_ANGULAR_VELOCITY, 3),
])
self.data.ctrl[:] = 0.0
mujoco.mj_forward(self.model, self.data)
self.steps = 0
self.episode_reward = 0.0
self.last_action = np.zeros(ACTION_SIZE, dtype=np.float32)
return self._observation()
def _observation(self) -> np.ndarray:
"""Construct the normalized 64-dimensional observation vector."""
plate_angles = np.array([self.data.qpos[self.plate_roll_qpos], self.data.qpos[self.plate_pitch_qpos]], dtype=np.float32)
plate_vel = np.array([self.data.qvel[self.plate_roll_dof], self.data.qvel[self.plate_pitch_dof]], dtype=np.float32)
plate_surface_center = self.data.site_xpos[self.plate_center_site]
info = OBJECT_INFO[self.current_object]
object_descriptor = np.array([
self.object_mass,
self.object_scale,
self.base_friction,
info.get("radius", info.get("size", 0.10)),
info.get("height", info.get("length", 0.10)),
], dtype=np.float32)
if self.current_object == "cookie":
# Multi-body crumb tracking
crumb_xys = [self.data.xpos[bid, :2] for bid in self.crumb_body_ids]
crumb_zs = [self.data.qpos[addr + 2] for addr in self.crumb_qpos_addrs]
on_plate_crumbs = [
p for p, z in zip(crumb_xys, crumb_zs)
if max(abs(p[0]), abs(p[1])) <= (PLATE_HALF_SIZE + FALL_MARGIN) and z >= FALL_Z_THRESHOLD
]
if on_plate_crumbs:
centroid_xy = np.mean(on_plate_crumbs, axis=0)
mean_z = float(np.mean([z for z in crumb_zs if z >= FALL_Z_THRESHOLD]))
spread = float(np.max([np.linalg.norm(p - centroid_xy) for p in on_plate_crumbs]))
else:
centroid_xy = np.zeros(2, dtype=np.float32)
mean_z = 1.15
spread = 0.0
object_pos = np.array([centroid_xy[0], centroid_xy[1], mean_z], dtype=np.float32)
rel_xy = (centroid_xy - plate_surface_center[:2]).astype(np.float32)
object_linvel = np.mean([self.data.qvel[dof:dof+3] for dof in self.crumb_dof_addrs], axis=0).astype(np.float32)
object_angvel = np.zeros(3, dtype=np.float32)
object_quat = np.array([1.0, 0.0, 0.0, 0.0], dtype=np.float32)
# Crumb specific relative features
crumb_ratio = float(len(on_plate_crumbs) / 5.0)
crumb_rel_1 = (crumb_xys[1] - centroid_xy).astype(np.float32)
crumb_rel_2 = (crumb_xys[2] - centroid_xy).astype(np.float32)
multi_body_feats = np.concatenate([
np.array([1.0, spread, crumb_ratio], dtype=np.float32),
crumb_rel_1,
crumb_rel_2,
])
else:
base_q = self.object_qpos_addr
base_v = self.object_dof_addr
object_pos = self.data.qpos[base_q:base_q + 3].astype(np.float32)
object_quat = self.data.qpos[base_q + 3:base_q + 7].astype(np.float32)
object_linvel = self.data.qvel[base_v:base_v + 3].astype(np.float32)
object_angvel = self.data.qvel[base_v + 3:base_v + 6].astype(np.float32)
rel_pos = object_pos - plate_surface_center
rel_xy = rel_pos[:2].astype(np.float32)
multi_body_feats = np.zeros(7, dtype=np.float32)
state = np.concatenate([
object_pos, # 3
rel_xy, # 2
object_linvel, # 3
object_angvel, # 3
object_quat, # 4
plate_angles, # 2
plate_vel, # 2
self.data.qfrc_actuator[:2], # 2
self.last_action, # 2
object_descriptor, # 5
multi_body_feats, # 7
]).astype(np.float32)
out = np.zeros(OBSERVATION_SIZE, dtype=np.float32)
n = min(len(state), OBSERVATION_SIZE)
out[:n] = state[:n]
out = np.clip(out, -10.0, 10.0)
return out
def step(self, action: np.ndarray) -> Tuple[np.ndarray, float, bool, bool, Dict[str, Any]]:
"""Advance simulation and compute reward for single or multi-crumb cookie objects."""
action = np.asarray(action, dtype=np.float64)
action = np.clip(action, -1.0, 1.0)
self.data.ctrl[0] = action[0]
self.data.ctrl[1] = action[1]
for _ in range(CONTROL_DECIMATION):
mujoco.mj_step(self.model, self.data)
self.steps += 1
plate_ang = np.array([self.data.qpos[self.plate_roll_qpos], self.data.qpos[self.plate_pitch_qpos]])
plate_ang_mag = float(np.linalg.norm(plate_ang))
action_rate_penalty = float(np.mean((action - self.last_action) ** 2))
self.last_action = action.astype(np.float32).copy()
if self.current_object == "cookie":
# Multi-body Crumbling Cookie Evaluation
crumb_xys = [self.data.xpos[bid, :2] for bid in self.crumb_body_ids]
crumb_zs = [self.data.qpos[addr + 2] for addr in self.crumb_qpos_addrs]
on_plate_flags = [
bool(max(abs(p[0]), abs(p[1])) <= (PLATE_HALF_SIZE + FALL_MARGIN) and z >= FALL_Z_THRESHOLD)
for p, z in zip(crumb_xys, crumb_zs)
]
crumbs_on = sum(on_plate_flags)
crumb_ratio = crumbs_on / 5.0
dists = [float(np.linalg.norm(p)) for p in crumb_xys]
max_dist = max(dists)
mean_dist = float(np.mean(dists))
box_dist = float(max(max(abs(p[0]), abs(p[1])) for p in crumb_xys))
vels = [float(np.linalg.norm(self.data.qvel[dof:dof+2])) for dof in self.crumb_dof_addrs]
mean_vel = float(np.mean(vels))
center_score = 0.5 * math.exp(-4.0 * mean_dist * mean_dist) + 0.5 * math.exp(-4.0 * max_dist * max_dist)
velocity_score = math.exp(-1.8 * mean_vel * mean_vel)
edge_fraction = np.clip(box_dist / PLATE_HALF_SIZE, 0.0, 1.5)
edge_penalty = max(0.0, float(edge_fraction) - 0.65) ** 2
# Reward scaled by the fraction of cookie crumbs kept on the plate
reward = (
SURVIVAL_REWARD * crumb_ratio
+ CENTER_REWARD_SCALE * center_score
+ VELOCITY_REWARD_SCALE * velocity_score
- ANGLE_PENALTY_SCALE * (plate_ang_mag / MAX_PLATE_ANGLE) ** 2
- ACTION_PENALTY_SCALE * float(np.mean(action ** 2))
- ACTION_RATE_PENALTY_SCALE * action_rate_penalty
- EDGE_PENALTY_SCALE * edge_penalty
)
# Failure occurs when any crumb falls off (or partial penalty for lost crumbs)
fallen = bool(crumbs_on < 5)
dist_xy = mean_dist
vel_xy = mean_vel
else:
# Single Rigid Object Evaluation
pos = self.data.xpos[self.object_body_id, :2]
dist_xy = float(np.linalg.norm(pos))
box_dist = float(max(abs(pos[0]), abs(pos[1])))
obj_vel = self.data.qvel[self.object_dof_addr:self.object_dof_addr + 3]
vel_xy = float(np.linalg.norm(obj_vel[:2]))
center_score = math.exp(-4.0 * dist_xy * dist_xy)
velocity_score = math.exp(-1.8 * vel_xy * vel_xy)
edge_fraction = np.clip(box_dist / PLATE_HALF_SIZE, 0.0, 1.5)
edge_penalty = max(0.0, float(edge_fraction) - 0.65) ** 2
reward = (
SURVIVAL_REWARD
+ CENTER_REWARD_SCALE * center_score
+ VELOCITY_REWARD_SCALE * velocity_score
- ANGLE_PENALTY_SCALE * (plate_ang_mag / MAX_PLATE_ANGLE) ** 2
- ACTION_PENALTY_SCALE * float(np.mean(action ** 2))
- ACTION_RATE_PENALTY_SCALE * action_rate_penalty
- EDGE_PENALTY_SCALE * edge_penalty
)
object_z = self.data.qpos[self.object_qpos_addr + 2]
fallen = bool(box_dist > (PLATE_HALF_SIZE + FALL_MARGIN) or object_z < FALL_Z_THRESHOLD)
crumbs_on = 1 if not fallen else 0
timeout = bool(self.steps >= MAX_EPISODE_STEPS)
terminated = fallen
truncated = timeout and not fallen
if fallen:
reward += FALL_PENALTY
self.episode_reward += reward
info = {
"object": self.current_object,
"distance": dist_xy,
"box_distance": box_dist,
"object_velocity": vel_xy,
"fallen": fallen,
"timeout": timeout,
"crumbs_on_plate": crumbs_on,
"episode_reward": self.episode_reward,
"episode_length": self.steps,
}
return self._observation(), float(reward), terminated, truncated, info
def render(self, camera: str = "track") -> np.ndarray:
"""Render RGB visual frame from the specified camera."""
if not self.render_enabled:
raise RuntimeError("Environment initialized with render=False")
assert self.renderer is not None
self.renderer.update_scene(self.data, camera=camera)
return self.renderer.render().copy()
def close(self) -> None:
"""Cleanly release rendering and physics resources."""
if self.renderer is not None:
self.renderer.close()
self.renderer = None
# =============================================================================
# ACTOR / CRITIC NEURAL NETWORK
# =============================================================================
class ActorCritic(nn.Module):
"""
Continuous-action Actor-Critic MLP with Tanh-squashed Gaussian policy
and exact, numerically stable log-probability computation.
"""
def __init__(self):
super().__init__()
self.backbone = nn.Sequential(
nn.Linear(OBSERVATION_SIZE, HIDDEN_SIZE),
nn.SiLU(),
nn.Linear(HIDDEN_SIZE, INTERMEDIATE_SIZE),
nn.SiLU(),
nn.Linear(INTERMEDIATE_SIZE, BOTTLENECK_SIZE),
nn.SiLU(),
)
self.actor = nn.Linear(BOTTLENECK_SIZE, ACTION_SIZE)
self.critic = nn.Linear(BOTTLENECK_SIZE, 1)
self.log_std = nn.Parameter(torch.full((ACTION_SIZE,), ACTOR_LOG_STD_INIT))
self._init_weights()
def _init_weights(self) -> None:
if not ORTHOGONAL_INIT:
return
for layer in self.backbone:
if isinstance(layer, nn.Linear):
nn.init.orthogonal_(layer.weight, gain=math.sqrt(2.0))
nn.init.zeros_(layer.bias)
nn.init.orthogonal_(self.actor.weight, gain=0.01)
nn.init.zeros_(self.actor.bias)
nn.init.orthogonal_(self.critic.weight, gain=1.0)
nn.init.zeros_(self.critic.bias)
def forward(self, obs: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
h = self.backbone(obs)
mean = self.actor(h)
value = self.critic(h).squeeze(-1)
std = self.log_std.exp().expand_as(mean)
return mean, std, value
def get_value(self, obs: torch.Tensor) -> torch.Tensor:
return self.critic(self.backbone(obs)).squeeze(-1)
def get_action_and_value(
self, obs: torch.Tensor, raw_action: Optional[torch.Tensor] = None, deterministic: bool = False
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
mean, std, value = self(obs)
dist = Normal(mean, std)
if raw_action is None:
raw_action = mean if deterministic else dist.sample()
squashed = torch.tanh(raw_action)
# Numerically stable exact log-prob correction for Tanh transformation
log_prob = dist.log_prob(raw_action).sum(-1)
log_prob -= torch.log(torch.clamp(1.0 - squashed.pow(2), min=1e-6)).sum(-1)
entropy = dist.entropy().sum(-1)
return squashed, log_prob, entropy, value, raw_action
# =============================================================================
# PPO TRAINER
# =============================================================================
@dataclass
class EpisodeStats:
reward: float = 0.0
length: int = 0
object_name: str = ""
success: bool = False
crumbs_on: int = 1
class PPOTrainer:
"""High-throughput PPO Trainer with robust GAE, metric logging, and checkpointing."""
def __init__(self, device: Optional[str] = None):
self.device = torch.device(device or DEVICE)
self.policy = ActorCritic().to(self.device)
self.optimizer = torch.optim.Adam(self.policy.parameters(), lr=LEARNING_RATE, eps=ADAM_EPS)
self.global_steps = 0
self.episodes = 0
self.updates = 0
self.last_log_step = 0
self.last_save_step = 0
self.last_video_step = 0
self.object_counts: Dict[str, int] = {name: 0 for name in OBJECT_TYPES}
def save(self, directory: Path, latest: bool = True, extra: Optional[Dict[str, Any]] = None) -> Path:
"""Save a complete resumable checkpoint."""
directory.mkdir(parents=True, exist_ok=True)
model_path = directory / "model.pt"
state_path = directory / "trainer_state.pt"
torch.save(self.policy.state_dict(), model_path)
trainer_state = {
"optimizer": self.optimizer.state_dict(),
"global_steps": self.global_steps,
"episodes": self.episodes,
"updates": self.updates,
"last_log_step": self.last_log_step,
"last_save_step": self.last_save_step,
"last_video_step": self.last_video_step,
"object_counts": self.object_counts,
"torch_rng_state": torch.get_rng_state(),
"numpy_rng_state": np.random.get_state(),
"python_rng_state": random.getstate(),
"config": {
k: v for k, v in globals().items()
if k.isupper() and isinstance(v, (int, float, str, bool, tuple, list))
},
"extra": extra or {},
}
if torch.cuda.is_available():
trainer_state["cuda_rng_state"] = torch.cuda.get_rng_state_all()
torch.save(trainer_state, state_path)
if latest:
latest_dir = directory.parent / "latest"
latest_dir.mkdir(parents=True, exist_ok=True)
torch.save(self.policy.state_dict(), latest_dir / "model.pt")
torch.save(trainer_state, latest_dir / "trainer_state.pt")
with open(latest_dir / "config.json", "w", encoding="utf-8") as f:
json.dump(trainer_state["config"], f, indent=2, default=str)
return directory
def load(self, directory: Path) -> None:
"""Restore policy weights, optimizer state, and RNGs from checkpoint."""
model_path = directory / "model.pt"
state_path = directory / "trainer_state.pt"
if not model_path.exists() or not state_path.exists():
raise FileNotFoundError(f"Checkpoint directory missing model.pt or trainer_state.pt: {directory}")
self.policy.load_state_dict(torch.load(model_path, map_location=self.device, weights_only=True))
state = torch.load(state_path, map_location="cpu", weights_only=False)
self.optimizer.load_state_dict(state["optimizer"])
self.global_steps = int(state.get("global_steps", 0))
self.episodes = int(state.get("episodes", 0))
self.updates = int(state.get("updates", 0))
self.last_log_step = int(state.get("last_log_step", 0))
self.last_save_step = int(state.get("last_save_step", 0))
self.last_video_step = int(state.get("last_video_step", 0))
self.object_counts.update(state.get("object_counts", {}))
try:
torch.set_rng_state(state["torch_rng_state"])
np.random.set_state(state["numpy_rng_state"])
random.setstate(state["python_rng_state"])
if torch.cuda.is_available() and "cuda_rng_state" in state:
torch.cuda.set_rng_state_all(state["cuda_rng_state"])
except Exception as exc:
print(f"[resume] Notice: Could not restore full RNG state: {exc}")
print(f"[resume] Loaded checkpoint {directory} | global_steps={self.global_steps:,} episodes={self.episodes:,}")
def update(self, batch: Dict[str, np.ndarray], total_steps_target: int = NUM_EPISODES * 100) -> Dict[str, float]:
"""Perform PPO mini-batch updates over the collected rollout."""
obs = torch.as_tensor(batch["obs"], dtype=torch.float32, device=self.device)
actions = torch.as_tensor(batch["raw_actions"], dtype=torch.float32, device=self.device)
old_logprobs = torch.as_tensor(batch["logprobs"], dtype=torch.float32, device=self.device)
advantages = torch.as_tensor(batch["advantages"], dtype=torch.float32, device=self.device)
returns = torch.as_tensor(batch["returns"], dtype=torch.float32, device=self.device)
old_values = torch.as_tensor(batch["values"], dtype=torch.float32, device=self.device)
# Standard advantage normalization
adv_std = advantages.std()
if adv_std > 1e-8:
advantages = (advantages - advantages.mean()) / (adv_std + 1e-8)
# Optional learning rate schedule annealing
if ANNEAL_LR:
frac = 1.0 - (self.global_steps / max(1, total_steps_target))
lr_now = max(1e-6, frac * LEARNING_RATE)
for param_group in self.optimizer.param_groups:
param_group["lr"] = lr_now
n = obs.shape[0]
minibatch = min(MINIBATCH_SIZE, n)
indices = np.arange(n)
metrics: Dict[str, List[float]] = {
"policy_loss": [],
"value_loss": [],
"entropy": [],
"approx_kl": [],
"clipfrac": [],
"explained_var": [],
}
for _ in range(PPO_EPOCHS):
np.random.shuffle(indices)
for start in range(0, n, minibatch):
mb = indices[start:start + minibatch]
_, new_logprob, entropy, new_value, _ = self.policy.get_action_and_value(obs[mb], actions[mb])
logratio = new_logprob - old_logprobs[mb]
ratio = logratio.exp()
mb_adv = advantages[mb]
# Policy loss (clipped surrogate)
pg_loss1 = -mb_adv * ratio
pg_loss2 = -mb_adv * torch.clamp(ratio, 1.0 - CLIP_COEF, 1.0 + CLIP_COEF)
policy_loss = torch.max(pg_loss1, pg_loss2).mean()
# Value loss (smooth MSE)
value_loss = 0.5 * ((new_value - returns[mb]) ** 2).mean()
entropy_loss = entropy.mean()
loss = policy_loss + VALUE_COEF * value_loss - ENTROPY_COEF * entropy_loss
self.optimizer.zero_grad(set_to_none=True)
loss.backward()
nn.utils.clip_grad_norm_(self.policy.parameters(), MAX_GRAD_NORM)
self.optimizer.step()
approx_kl = ((ratio - 1.0) - logratio).mean().detach().cpu().item()
clipfrac = ((ratio - 1.0).abs() > CLIP_COEF).float().mean().detach().cpu().item()
metrics["policy_loss"].append(policy_loss.detach().cpu().item())
metrics["value_loss"].append(value_loss.detach().cpu().item())
metrics["entropy"].append(entropy_loss.detach().cpu().item())
metrics["approx_kl"].append(approx_kl)
metrics["clipfrac"].append(clipfrac)
self.updates += 1
y_true = returns.cpu().numpy()
y_pred = old_values.cpu().numpy()
var_y = np.var(y_true)
explained_var = float(1.0 - np.var(y_true - y_pred) / (var_y + 1e-8)) if var_y > 1e-8 else 0.0
out_metrics = {k: float(np.mean(v)) for k, v in metrics.items() if len(v) > 0}
out_metrics["explained_var"] = explained_var
return out_metrics
# =============================================================================
# ROLLOUT COLLECTION & GAE
# =============================================================================
def collect_rollout(
policy: ActorCritic,
envs: List[PlateBalanceEnv],
current_obs: np.ndarray,
trainer: PPOTrainer,
) -> Tuple[Dict[str, np.ndarray], np.ndarray, List[EpisodeStats], int]:
"""Collect vector rollouts and compute generalized advantage estimations."""
n_envs = len(envs)
T = ROLLOUT_STEPS
obs_buf = np.zeros((T, n_envs, OBSERVATION_SIZE), dtype=np.float32)
actions_buf = np.zeros((T, n_envs, ACTION_SIZE), dtype=np.float32)
raw_actions_buf = np.zeros((T, n_envs, ACTION_SIZE), dtype=np.float32)
logprob_buf = np.zeros((T, n_envs), dtype=np.float32)
rewards_buf = np.zeros((T, n_envs), dtype=np.float32)
terminated_buf = np.zeros((T, n_envs), dtype=np.float32)
truncated_buf = np.zeros((T, n_envs), dtype=np.float32)
values_buf = np.zeros((T, n_envs), dtype=np.float32)
completed_episodes: List[EpisodeStats] = []
start_steps = trainer.global_steps
policy.eval()
for t in range(T):
obs_buf[t] = current_obs
obs_t = torch.as_tensor(current_obs, dtype=torch.float32, device=trainer.device)
with torch.no_grad():
action_t, logprob_t, _, value_t, raw_t = policy.get_action_and_value(obs_t)
actions = action_t.cpu().numpy()
raw_actions = raw_t.cpu().numpy()
actions_buf[t] = actions
raw_actions_buf[t] = raw_actions
logprob_buf[t] = logprob_t.cpu().numpy()
values_buf[t] = value_t.cpu().numpy()
next_obs = np.empty_like(current_obs)
for e, env in enumerate(envs):
next_obs[e], reward, terminated, truncated, info = env.step(actions[e])
rewards_buf[t, e] = reward
terminated_buf[t, e] = float(terminated)
truncated_buf[t, e] = float(truncated)
trainer.global_steps += 1
if terminated or truncated:
trainer.episodes += 1
trainer.object_counts[info["object"]] = trainer.object_counts.get(info["object"], 0) + 1
completed_episodes.append(EpisodeStats(
reward=float(info["episode_reward"]),
length=int(info["episode_length"]),
object_name=str(info["object"]),
success=bool(not info["fallen"] and info["episode_length"] >= MAX_EPISODE_STEPS),
crumbs_on=int(info.get("crumbs_on_plate", 1)),
))
# Reset environment immediately upon termination or truncation
next_obs[e] = env.reset()
current_obs = next_obs
if trainer.episodes >= NUM_EPISODES:
break
actual_T = t + 1
obs_buf = obs_buf[:actual_T]
actions_buf = actions_buf[:actual_T]
raw_actions_buf = raw_actions_buf[:actual_T]
logprob_buf = logprob_buf[:actual_T]
rewards_buf = rewards_buf[:actual_T]
terminated_buf = terminated_buf[:actual_T]
truncated_buf = truncated_buf[:actual_T]
values_buf = values_buf[:actual_T]
with torch.no_grad():
next_obs_t = torch.as_tensor(current_obs, dtype=torch.float32, device=trainer.device)
next_value = policy.get_value(next_obs_t).cpu().numpy()
# GAE Computation with proper termination vs. truncation bootstrapping
advantages = np.zeros_like(rewards_buf)
lastgaelam = np.zeros(n_envs, dtype=np.float32)
for t2 in reversed(range(actual_T)):
if t2 == actual_T - 1:
next_vals = next_value
else:
next_vals = values_buf[t2 + 1]
# Only true terminations (falling) zero out future value; timeouts bootstrap value
nonterminal = 1.0 - terminated_buf[t2]
delta = rewards_buf[t2] + GAMMA * next_vals * nonterminal - values_buf[t2]
lastgaelam = delta + GAMMA * GAE_LAMBDA * nonterminal * lastgaelam
advantages[t2] = lastgaelam
returns = advantages + values_buf
batch = {
"obs": obs_buf.reshape(-1, OBSERVATION_SIZE),
"actions": actions_buf.reshape(-1, ACTION_SIZE),
"raw_actions": raw_actions_buf.reshape(-1, ACTION_SIZE),
"logprobs": logprob_buf.reshape(-1),
"rewards": rewards_buf.reshape(-1),
"terminated": terminated_buf.reshape(-1),
"values": values_buf.reshape(-1),
"advantages": advantages.reshape(-1),
"returns": returns.reshape(-1),
}
policy.train()
return batch, current_obs, completed_episodes, start_steps
# =============================================================================
# EVALUATION & VIDEO VISUALIZATION
# =============================================================================
def record_video(
policy: ActorCritic,
out_path: Path,
seed: int = 1234,
max_steps: int = VIDEO_LENGTH_STEPS,
camera: str = "track",
specific_object: Optional[str] = None,
) -> None:
"""Record and save an MP4 demonstration video of the policy."""
if imageio is None:
print("[video] imageio / imageio-ffmpeg not installed. Skipping video recording.")
return
env = PlateBalanceEnv(seed=seed, render=True)
obs = env.reset(specific_object=specific_object)
frames: List[np.ndarray] = []
policy.eval()
try:
with torch.no_grad():
for _ in range(max_steps):
frames.append(env.render(camera=camera))
obs_t = torch.as_tensor(obs, dtype=torch.float32, device=DEVICE).unsqueeze(0)
action, _, _, _, _ = policy.get_action_and_value(obs_t, deterministic=True)
obs, _, terminated, truncated, _ = env.step(action[0].cpu().numpy())
if terminated or truncated:
obs = env.reset()
out_path.parent.mkdir(parents=True, exist_ok=True)
imageio.mimsave(out_path, frames, fps=VIDEO_FPS, codec="libx264", quality=7)
print(f"[video] Saved {len(frames)} frames to: {out_path}")
finally:
env.close()
policy.train()
def evaluate_policy(
policy: ActorCritic,
episodes_per_object: int = 5,
seed: int = 777,
deterministic: bool = True,
eval_objects: Optional[List[str]] = None,
) -> Dict[str, Any]:
"""Benchmark the policy across all 21 supported object geometries."""
results: Dict[str, Dict[str, float]] = {}
policy.eval()
active_eval_list = eval_objects or OBJECT_TYPES
print("\n" + "=" * 88)
print("EVALUATION BENCHMARK ACROSS ALL 21 OBJECT GEOMETRIES (INCL. CRUMBLING COOKIE)")
print("=" * 88)
print(f"{'Category':<14} | {'Object':<16} | {'Reward Mean':<12} | {'Len Mean':<10} | {'Survival %':<12} | {'Tracking Err':<12}")
print("-" * 88)
for obj_name in active_eval_list:
category = "Multi-Body" if obj_name == "cookie" else ("Harder" if obj_name in HARDER_OBJECT_TYPES else "Standard")
env = PlateBalanceEnv(seed=seed, render=False, active_objects=[obj_name])
rewards: List[float] = []
lengths: List[int] = []
distances: List[float] = []
survived = 0
for ep in range(episodes_per_object):
obs = env.reset(specific_object=obj_name)
done = False
ep_dist = []
while not done:
with torch.no_grad():
obs_t = torch.as_tensor(obs, dtype=torch.float32, device=DEVICE).unsqueeze(0)
action, _, _, _, _ = policy.get_action_and_value(obs_t, deterministic=deterministic)
obs, reward, terminated, truncated, info = env.step(action[0].cpu().numpy())
ep_dist.append(info["distance"])
if terminated or truncated:
rewards.append(info["episode_reward"])
lengths.append(info["episode_length"])
if not info["fallen"]:
survived += 1
done = True
distances.append(float(np.mean(ep_dist)))
env.close()
mean_r = float(np.mean(rewards)) if rewards else 0.0
mean_l = float(np.mean(lengths)) if lengths else 0.0
surv_rate = (survived / episodes_per_object) * 100.0
mean_d = float(np.mean(distances)) if distances else 0.0
results[obj_name] = {
"mean_reward": mean_r,
"mean_length": mean_l,
"survival_rate": surv_rate,
"tracking_error": mean_d,
}
print(f"{category:<14} | {obj_name:<16} | {mean_r:>12.2f} | {mean_l:>10.1f} | {surv_rate:>11.1f}% | {mean_d:>12.4f}m")
print("=" * 88 + "\n")
policy.train()
return results
def human_view(policy: ActorCritic, specific_object: Optional[str] = None) -> None:
"""Interactive real-time 3D MuJoCo viewer."""
try:
import mujoco.viewer
except ImportError:
print("[viewer] mujoco.viewer is unavailable on this system.")
return
env = PlateBalanceEnv(seed=SEED + 999, render=False)
obs = env.reset(specific_object=specific_object)
control_dt = PHYSICS_TIMESTEP * CONTROL_DECIMATION
print(f"\n[viewer] Launching interactive 3D viewer for: {specific_object or 'Random Objects'}")
print("[viewer] Controls: Space to pause/resume, Esc/close window to exit.")
with mujoco.viewer.launch_passive(env.model, env.data) as viewer:
policy.eval()
while viewer.is_running():
step_start = time.time()
with torch.no_grad():
obs_t = torch.as_tensor(obs, dtype=torch.float32, device=DEVICE).unsqueeze(0)
action, _, _, _, _ = policy.get_action_and_value(obs_t, deterministic=True)
obs, _, terminated, truncated, info = env.step(action[0].cpu().numpy())
viewer.sync()
if terminated or truncated:
status = "TIMEOUT" if truncated else "FALL"
crumbs_info = f" | Crumbs on plate: {info.get('crumbs_on_plate', 1)}/5" if info["object"] == "cookie" else ""
print(f"[viewer] Episode end: {status} | Object: {info['object']} | Length: {info['episode_length']}{crumbs_info}")
obs = env.reset(specific_object=specific_object)
# Precise real-time rate pacing
elapsed = time.time() - step_start
if elapsed < control_dt:
time.sleep(control_dt - elapsed)
policy.train()
env.close()
def export_onnx(policy: ActorCritic, out_path: Path) -> None:
"""Export the trained policy backbone and actor head to standard ONNX format."""
policy.eval()
dummy_input = torch.zeros(1, OBSERVATION_SIZE, dtype=torch.float32, device=DEVICE)
class ExportWrapper(nn.Module):
def __init__(self, p: ActorCritic):
super().__init__()
self.policy = p
def forward(self, x: torch.Tensor) -> torch.Tensor:
mean, _, _ = self.policy(x)
return torch.tanh(mean)
wrapper = ExportWrapper(policy)
out_path.parent.mkdir(parents=True, exist_ok=True)
torch.onnx.export(
wrapper,
dummy_input,
str(out_path),
input_names=["observation"],
output_names=["action"],
dynamic_axes={"observation": {0: "batch_size"}, "action": {0: "batch_size"}},
opset_version=14,
)
print(f"[export] Successfully exported ONNX model to: {out_path}")
policy.train()
# =============================================================================
# MAIN ENTRY POINT & TRAINING LOOP
# =============================================================================
def print_config(active_objects: List[str]) -> None:
if not PRINT_HYPERPARAMS:
return
print("=" * 88)
print("PlateBalance RL - Advanced MuJoCo + PyTorch PPO System (21 Objects & Cookie Crumble)")
print("=" * 88)
print(f"Device: {DEVICE} | Master Seed: {SEED}")
print(f"Policy: Obs({OBSERVATION_SIZE}) -> {HIDDEN_SIZE} -> {INTERMEDIATE_SIZE} -> {BOTTLENECK_SIZE} -> Act({ACTION_SIZE})")
print(f"PPO: Envs={NUM_ENVS} | Rollout={ROLLOUT_STEPS} | Epochs={PPO_EPOCHS} | Batch={MINIBATCH_SIZE} | LR={LEARNING_RATE}")
print(f"Target: Episodes={NUM_EPISODES:,} | MaxSteps/Ep={MAX_EPISODE_STEPS:,}")
print(f"Active Objects ({len(active_objects)}): {', '.join(active_objects)}")
print(f"Domain Randomization: Mass={RANDOMIZE_MASS}, Friction={RANDOMIZE_FRICTION}, Geometry={RANDOMIZE_OBJECT_SIZE}")
print(f"Actuator Torque: {MAX_PLATE_TORQUE} N*m | Control Rate: {1.0/(PHYSICS_TIMESTEP*CONTROL_DECIMATION):.0f} Hz")
print("=" * 88)
def main() -> None:
global SEED, DEVICE, NUM_EPISODES
parser = argparse.ArgumentParser(description="PlateBalance RL: MuJoCo PPO Object Balancing")
parser.add_argument("--eval", action="store_true", help="Run benchmark evaluation across all 21 object shapes")
parser.add_argument("--human-view", action="store_true", help="Launch interactive 3D viewer")
parser.add_argument("--record-video", action="store_true", help="Record evaluation video")
parser.add_argument("--export-onnx", type=str, default="", help="Export model to ONNX file path")
parser.add_argument("--checkpoint", type=str, default="", help="Path to checkpoint folder to load")
parser.add_argument("--resume", action="store_true", help="Auto-resume training from latest checkpoint")
parser.add_argument("--episodes", type=int, default=NUM_EPISODES, help="Total training episodes")
parser.add_argument("--seed", type=int, default=SEED, help="Random seed")
parser.add_argument("--device", type=str, default=DEVICE, help="Compute device (cuda or cpu)")
parser.add_argument("--object", type=str, default="", help="Specific object to balance for training/eval/viewer")
parser.add_argument("--category", type=str, default="all", choices=["all", "standard", "harder", "cookie"],
help="Filter object category (all, standard, harder, cookie)")
args = parser.parse_args()
SEED = args.seed
DEVICE = args.device
NUM_EPISODES = args.episodes
# Filter active objects
if args.object:
if args.object not in OBJECT_INFO:
raise ValueError(f"Unknown object '{args.object}'. Available: {', '.join(OBJECT_TYPES)}")
active_objects = [args.object]
elif args.category == "standard":
active_objects = STANDARD_OBJECT_TYPES
elif args.category == "harder":
active_objects = HARDER_OBJECT_TYPES
elif args.category == "cookie":
active_objects = MULTI_BODY_OBJECT_TYPES
else:
active_objects = OBJECT_TYPES
set_seed(SEED)
print_config(active_objects)
root = Path(OUTPUT_DIR) / EXPERIMENT_NAME
root.mkdir(parents=True, exist_ok=True)
trainer = PPOTrainer(device=DEVICE)
# Determine checkpoint loading
ckpt_path: Optional[Path] = None
if args.checkpoint:
ckpt_path = Path(args.checkpoint)
elif args.resume or RESUME:
ckpt_path = Path(RESUME_PATH) if RESUME_PATH else latest_checkpoint(root / "checkpoints")
if ckpt_path is not None and ckpt_path.exists():
trainer.load(ckpt_path)
# Export ONNX mode
if args.export_onnx:
export_onnx(trainer.policy, Path(args.export_onnx))
return
# Interactive Human Viewer mode
if args.human_view or HUMAN_VIEW:
human_view(trainer.policy, specific_object=args.object or None)
return
# Benchmark Evaluation mode
if args.eval:
evaluate_policy(trainer.policy, eval_objects=active_objects)
return
# Record Video mode
if args.record_video:
vid_path = root / "eval_videos" / "eval_demo.mp4"
record_video(trainer.policy, vid_path, seed=SEED, specific_object=args.object or None)
return
# Standard Training Setup
envs = [
PlateBalanceEnv(SEED + 1000 * i + trainer.global_steps, render=False, active_objects=active_objects)
for i in range(NUM_ENVS)
]
current_obs = np.stack([env.reset() for env in envs], axis=0)
start_time = time.time()
running_rewards: List[float] = []
running_lengths: List[int] = []
running_success: List[float] = []
print(f"\n[train] Starting PPO training loop with {NUM_ENVS} parallel environments...\n")
try:
while trainer.episodes < NUM_EPISODES:
batch, current_obs, completed, rollout_start = collect_rollout(
trainer.policy, envs, current_obs, trainer
)
if batch["obs"].shape[0] >= 2:
metrics = trainer.update(batch, total_steps_target=NUM_EPISODES * 500)
else:
metrics = {
"policy_loss": 0.0, "value_loss": 0.0, "entropy": 0.0,
"approx_kl": 0.0, "clipfrac": 0.0, "explained_var": 0.0,
}
running_rewards.extend(ep.reward for ep in completed)
running_lengths.extend(ep.length for ep in completed)
running_success.extend(1.0 if ep.success else 0.0 for ep in completed)
# Periodic console logging
if trainer.global_steps - trainer.last_log_step >= LOGGING_STEPS:
trainer.last_log_step = trainer.global_steps
elapsed = max(time.time() - start_time, 1e-9)
sps = (trainer.global_steps / elapsed)
mean_reward = float(np.mean(running_rewards[-50:])) if running_rewards else 0.0
mean_len = float(np.mean(running_lengths[-50:])) if running_lengths else 0.0
succ_rate = (float(np.mean(running_success[-50:])) * 100.0) if running_success else 0.0
print(
f"step={trainer.global_steps:>9,} | ep={trainer.episodes:>7,} | "
f"sps={sps:>6.0f} | r50={mean_reward:>8.2f} | len50={mean_len:>6.1f} | "
f"succ={succ_rate:>5.1f}% | pi={metrics.get('policy_loss', 0.0):+.4f} | "
f"vf={metrics.get('value_loss', 0.0):.4f} | kl={metrics.get('approx_kl', 0.0):.5f}"
)
# Periodic checkpoint saving
if trainer.global_steps - trainer.last_save_step >= SAVE_STEPS:
trainer.last_save_step = trainer.global_steps
ckpt = root / "checkpoints" / f"step_{trainer.global_steps:08d}"
trainer.save(ckpt, latest=True, extra={
"mean_reward_50": float(np.mean(running_rewards[-50:])) if running_rewards else 0.0,
"mean_length_50": float(np.mean(running_lengths[-50:])) if running_lengths else 0.0,
})
print(f"[save] Checkpoint saved: {ckpt}")
# Periodic visualization video recording
if VISUALIZE and trainer.global_steps - trainer.last_video_step >= VIDEO_EVERY_STEPS:
trainer.last_video_step = trainer.global_steps
ckpt = root / "checkpoints" / f"step_{trainer.global_steps:08d}"
ckpt.mkdir(parents=True, exist_ok=True)
video_path = ckpt / f"balance_step_{trainer.global_steps:08d}.mp4"
print(f"[video] Recording rollout video to: {video_path}")
record_video(trainer.policy, video_path, seed=SEED + trainer.global_steps, max_steps=VIDEO_LENGTH_STEPS)
trainer.save(ckpt, latest=True, extra={"video": str(video_path.name)})
except KeyboardInterrupt:
print("\n[interrupt] Training interrupted by user. Saving emergency checkpoint...")
ckpt = root / "checkpoints" / f"step_{trainer.global_steps:08d}_interrupt"
trainer.save(ckpt, latest=True, extra={"interrupted": True})
print(f"[interrupt] Saved emergency checkpoint: {ckpt}")
finally:
for env in envs:
env.close()
final_dir = root / "checkpoints" / f"step_{trainer.global_steps:08d}_final"
trainer.save(final_dir, latest=True, extra={"finished": True})
print(f"\n[done] Training completed! Total episodes={trainer.episodes:,} steps={trainer.global_steps:,}")
print(f"[done] Final model and state saved to: {final_dir}")
if __name__ == "__main__":
main()