#!/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""" """ # ============================================================================= # 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()