| |
| """ |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| |
| SEED = 42 |
| EXPERIMENT_NAME = "plate_balance_v1" |
| OUTPUT_DIR = "runs" |
| RESUME = True |
| RESUME_PATH = r"runs\plate_balance_v1\checkpoints\step_03825064_final" |
| DEVICE = "cuda" if torch.cuda.is_available() else "cpu" |
|
|
| |
| OBSERVATION_SIZE = 64 |
| HIDDEN_SIZE = 128 |
| INTERMEDIATE_SIZE = 128 |
| BOTTLENECK_SIZE = 64 |
| ACTION_SIZE = 2 |
|
|
| |
| ACTOR_LOG_STD_INIT = -0.75 |
| ORTHOGONAL_INIT = True |
|
|
| |
| 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 |
|
|
| |
| NUM_EPISODES = 30_000 |
| MAX_EPISODE_STEPS = 10_000 |
| LOGGING_STEPS = 32240 |
| SAVE_STEPS = 1_000_000 |
|
|
| |
| VISUALIZE = True |
| VIDEO_EVERY_STEPS = 1_000_000 |
| VIDEO_LENGTH_STEPS = 3000 |
| VIDEO_FPS = 60 |
| RENDER_WIDTH = 640 |
| RENDER_HEIGHT = 480 |
| HUMAN_VIEW = False |
|
|
| |
| PHYSICS_TIMESTEP = 0.0025 |
| CONTROL_DECIMATION = 4 |
| GRAVITY = 9.81 |
| PLATE_HALF_SIZE = 0.90 |
| PLATE_THICKNESS = 0.08 |
| PLATE_FRICTION = (1.0, 0.02, 0.002) |
| PLATE_DAMPING = 0.08 |
| MAX_PLATE_ANGLE = math.radians(15.0) |
| MAX_PLATE_TORQUE = 15.0 |
|
|
| OBJECT_START_HEIGHT = 0.02 |
| OBJECT_START_POS_RANGE = 0.35 |
| OBJECT_START_VELOCITY = 0.15 |
| OBJECT_START_ANGULAR_VELOCITY = 0.20 |
|
|
| |
| RANDOMIZE_MASS = True |
| RANDOMIZE_FRICTION = True |
| RANDOMIZE_OBJECT_SIZE = True |
| MASS_LOG_RANGE = (0.50, 2.00) |
| FRICTION_RANGE = (0.45, 1.25) |
| SIZE_RANGE = (0.85, 1.15) |
| RANDOMIZE_PLATE_ANGLE = True |
| INITIAL_PLATE_ANGLE_RANGE = math.radians(3.0) |
|
|
| |
| CENTER_REWARD_SCALE = 2.0 |
| VELOCITY_REWARD_SCALE = 0.40 |
| ANGLE_PENALTY_SCALE = 0.20 |
| ACTION_PENALTY_SCALE = 0.005 |
| ACTION_RATE_PENALTY_SCALE = 0.01 |
| EDGE_PENALTY_SCALE = 0.50 |
| FALL_PENALTY = -10.0 |
| SURVIVAL_REWARD = 0.05 |
|
|
| |
| FALL_MARGIN = 0.10 |
| FALL_Z_THRESHOLD = 0.90 |
|
|
| |
| 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", |
| ] |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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()) |
|
|
|
|
| |
| |
| |
|
|
| OBJECT_INFO: Dict[str, Dict[str, float]] = { |
| |
| "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}, |
|
|
| |
| "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}, |
|
|
| |
| "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) |
|
|
|
|
| |
| TRIANGLE_MESH = """ |
| <mesh name="triangle_mesh" |
| vertex="-0.30 -0.24 -0.08 0.30 -0.24 -0.08 0.00 0.30 -0.08 |
| -0.30 -0.24 0.08 0.30 -0.24 0.08 0.00 0.30 0.08" |
| face="0 1 2 3 5 4 0 3 4 0 4 1 1 4 5 1 5 2 2 5 3 2 3 0" /> |
| """ |
|
|
| CONE_MESH = """ |
| <mesh name="cone_mesh" |
| vertex=" 0.00 0.00 0.18 |
| 0.16 0.00 -0.14 |
| 0.11 0.11 -0.14 |
| 0.00 0.16 -0.14 |
| -0.11 0.11 -0.14 |
| -0.16 0.00 -0.14 |
| -0.11 -0.11 -0.14 |
| 0.00 -0.16 -0.14 |
| 0.11 -0.11 -0.14 |
| 0.00 0.00 -0.14" |
| face="0 1 2 0 2 3 0 3 4 0 4 5 0 5 6 0 6 7 0 7 8 0 8 1 |
| 9 2 1 9 3 2 9 4 3 9 5 4 9 6 5 9 7 6 9 8 7 9 1 8" /> |
| """ |
|
|
| WEDGE_MESH = """ |
| <mesh name="wedge_mesh" |
| vertex="-0.18 -0.15 -0.10 |
| 0.18 -0.15 -0.10 |
| 0.18 0.15 -0.10 |
| -0.18 0.15 -0.10 |
| -0.18 -0.15 0.10 |
| 0.18 -0.15 0.10" |
| face="0 1 2 0 2 3 |
| 0 4 5 0 5 1 |
| 1 5 2 |
| 0 3 4 |
| 3 2 5 3 5 4" /> |
| """ |
|
|
| TETRA_MESH = """ |
| <mesh name="tetra_mesh" |
| vertex=" 0.00 0.00 0.18 |
| 0.18 -0.12 -0.12 |
| -0.18 -0.12 -0.12 |
| 0.00 0.20 -0.12" |
| face="0 1 2 0 2 3 0 3 1 1 3 2" /> |
| """ |
|
|
|
|
| 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 model="plate_balance"> |
| <compiler angle="radian" coordinate="local" inertiafromgeom="auto" /> |
| <option timestep="{PHYSICS_TIMESTEP}" gravity="0 0 -{GRAVITY}" integrator="implicitfast" /> |
| |
| <default> |
| <joint damping="{PLATE_DAMPING}" armature="0.01" /> |
| <geom solref="0.008 1" solimp="0.90 0.95 0.01" /> |
| <motor ctrllimited="true" ctrlrange="-1.0 1.0" /> |
| </default> |
| |
| <asset> |
| <texture type="skybox" builtin="gradient" rgb1="0.85 0.90 0.98" rgb2="0.65 0.75 0.90" width="512" height="512" /> |
| <texture name="grid" type="2d" builtin="checker" width="512" height="512" rgb1="0.92 0.92 0.92" rgb2="0.80 0.80 0.80" /> |
| <material name="grid_mat" texture="grid" texrepeat="5 5" reflectance="0.1" /> |
| <texture name="cookie_tex" type="2d" builtin="checker" width="64" height="64" rgb1="0.82 0.58 0.32" rgb2="0.68 0.42 0.22" /> |
| <material name="cookie_mat" texture="cookie_tex" reflectance="0.1" /> |
| <material name="plate_mat" rgba="0.16 0.45 0.92 1" reflectance="0.3" specular="0.5" /> |
| <material name="object_mat" rgba="0.95 0.45 0.16 1" reflectance="0.2" specular="0.4" /> |
| <material name="heavy_mat" rgba="0.32 0.35 0.42 1" reflectance="0.6" specular="0.8" /> |
| {TRIANGLE_MESH} |
| {CONE_MESH} |
| {WEDGE_MESH} |
| {TETRA_MESH} |
| </asset> |
| |
| <worldbody> |
| <light directional="true" pos="0 -3 5" dir="0 0.5 -1" diffuse="0.8 0.8 0.8" specular="0.3 0.3 0.3" /> |
| <light directional="true" pos="3 3 5" dir="-0.5 -0.5 -1" diffuse="0.4 0.4 0.4" /> |
| <geom name="floor" type="plane" size="6 6 0.1" pos="0 0 0" material="grid_mat" contype="1" conaffinity="1" /> |
| |
| <!-- Cinematic cameras for visualization and video recording --> |
| <camera name="track" pos="0 -2.4 2.2" xyaxes="1 0 0 0 0.6 0.8" /> |
| <camera name="isometric" pos="2.0 -2.0 2.3" xyaxes="0.707 0.707 0 -0.408 0.408 0.816" /> |
| <camera name="top_down" pos="0 0 3.3" xyaxes="1 0 0 0 1 0" /> |
| <camera name="side" pos="-2.8 0 1.5" xyaxes="0 -1 0 0.2 0 0.98" /> |
| |
| <!-- 2-Axis Gimbal Tilting Plate Mechanism --> |
| <body name="plate_x" pos="0 0 1.05"> |
| <joint name="plate_roll" type="hinge" axis="1 0 0" range="-{MAX_PLATE_ANGLE} {MAX_PLATE_ANGLE}" limited="true" /> |
| <inertial pos="0 0 0" mass="4.0" diaginertia="1.0 1.0 1.0" /> |
| |
| <body name="plate_y"> |
| <joint name="plate_pitch" type="hinge" axis="0 1 0" range="-{MAX_PLATE_ANGLE} {MAX_PLATE_ANGLE}" limited="true" /> |
| <inertial pos="0 0 0" mass="4.0" diaginertia="1.0 1.0 1.0" /> |
| <geom name="plate" type="box" size="{PLATE_HALF_SIZE} {PLATE_HALF_SIZE} {PLATE_THICKNESS}" |
| material="plate_mat" friction="{PLATE_FRICTION[0]} {PLATE_FRICTION[1]} {PLATE_FRICTION[2]}" |
| contype="1" conaffinity="1" mass="4.0" /> |
| <site name="plate_center" pos="0 0 {PLATE_THICKNESS + 0.005}" size="0.015" rgba="1 1 1 0.6" /> |
| </body> |
| </body> |
| |
| <!-- Primary candidate object body (for single rigid objects) --> |
| <body name="object" pos="0 0 {spawn_z}"> |
| <freejoint name="object_free" /> |
| |
| <!-- Baseline 10 Shapes --> |
| <geom name="g_sphere" type="sphere" size="0.16" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_disk" type="cylinder" size="0.22 0.0275" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_egg" type="ellipsoid" size="0.135 0.135 0.19" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| |
| <!-- Hollow cup: base + 4 walls --> |
| <geom name="g_cup_bottom" type="cylinder" size="0.18 0.015" pos="0 0 -0.11" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_cup_w1" type="box" size="0.015 0.18 0.125" pos="0.165 0 0" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_cup_w2" type="box" size="0.015 0.18 0.125" pos="-0.165 0 0" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_cup_w3" type="box" size="0.15 0.015 0.125" pos="0 0.165 0" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_cup_w4" type="box" size="0.15 0.015 0.125" pos="0 -0.165 0" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| |
| <geom name="g_coin" type="cylinder" size="0.18 0.0125" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_stick" type="capsule" size="0.035 0.24" fromto="0 0 -0.24 0 0 0.24" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_tall" type="box" size="0.045 0.045 0.26" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_triangle" type="mesh" mesh="triangle_mesh" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_block" type="box" size="0.14 0.14 0.14" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_puck" type="cylinder" size="0.25 0.04" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| |
| <!-- Harder 10 Shapes --> |
| <geom name="g_cone" type="mesh" mesh="cone_mesh" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_capsule" type="capsule" size="0.065 0.18" fromto="-0.18 0 0 0.18 0 0" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_wedge" type="mesh" mesh="wedge_mesh" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_tetra" type="mesh" mesh="tetra_mesh" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_flat_bar" type="box" size="0.38 0.05 0.02" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_cross_1" type="box" size="0.26 0.055 0.04" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_cross_2" type="box" size="0.055 0.26 0.04" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_lshape_1" type="box" size="0.055 0.16 0.04" pos="0 -0.08 0" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_lshape_2" type="box" size="0.14 0.055 0.04" pos="0.085 0.08 0" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_wide_block" type="box" size="0.32 0.32 0.02" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| <geom name="g_heavy_ball" type="sphere" size="0.16" material="heavy_mat" contype="1" conaffinity="1" /> |
| <geom name="g_offcenter_block" type="box" size="0.18 0.18 0.10" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" /> |
| </body> |
| |
| <!-- 5 Crumb bodies for the Crumbling Cookie challenge --> |
| <body name="crumb_0" pos="0 0 1.15"> |
| <freejoint name="crumb_free_0" /> |
| <geom name="g_crumb_0" type="cylinder" size="0.08 0.02" material="cookie_mat" contype="1" conaffinity="1" mass="0.15" /> |
| </body> |
| <body name="crumb_1" pos="0.09 0 1.15"> |
| <freejoint name="crumb_free_1" /> |
| <geom name="g_crumb_1" type="box" size="0.035 0.035 0.02" material="cookie_mat" contype="1" conaffinity="1" mass="0.08" /> |
| </body> |
| <body name="crumb_2" pos="-0.09 0 1.15"> |
| <freejoint name="crumb_free_2" /> |
| <geom name="g_crumb_2" type="cylinder" size="0.035 0.02" material="cookie_mat" contype="1" conaffinity="1" mass="0.08" /> |
| </body> |
| <body name="crumb_3" pos="0 0.09 1.15"> |
| <freejoint name="crumb_free_3" /> |
| <geom name="g_crumb_3" type="box" size="0.04 0.03 0.02" material="cookie_mat" contype="1" conaffinity="1" mass="0.08" /> |
| </body> |
| <body name="crumb_4" pos="0 -0.09 1.15"> |
| <freejoint name="crumb_free_4" /> |
| <geom name="g_crumb_4" type="sphere" size="0.03" material="cookie_mat" contype="1" conaffinity="1" mass="0.06" /> |
| </body> |
| </worldbody> |
| |
| <actuator> |
| <motor name="roll_motor" joint="plate_roll" gear="{MAX_PLATE_TORQUE}" /> |
| <motor name="pitch_motor" joint="plate_pitch" gear="{MAX_PLATE_TORQUE}" /> |
| </actuator> |
| </mujoco> |
| """ |
|
|
|
|
| |
| |
| |
|
|
| class PlateBalanceEnv: |
| """ |
| Robust single-instance MuJoCo environment supporting 21 single & multi-body |
| objects with accurate physical scaling and multi-crumb tracking. |
| """ |
|
|
| OBJECT_GEOM_NAMES = { |
| |
| "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"], |
| |
| "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) |
|
|
| |
| 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") |
|
|
| |
| 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") |
|
|
| |
| 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)] |
|
|
| |
| 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() |
| } |
|
|
| |
| 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.""" |
| |
| 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 |
|
|
| |
| 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": |
| |
| 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 |
|
|
| |
| 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 |
|
|
| |
| 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]) |
|
|
| |
| 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: |
| 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) |
|
|
| |
| 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": |
| |
| 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 |
|
|
| |
| 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: |
| |
| 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 |
|
|
| |
| 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": |
| |
| 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_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, |
| rel_xy, |
| object_linvel, |
| object_angvel, |
| object_quat, |
| plate_angles, |
| plate_vel, |
| self.data.qfrc_actuator[:2], |
| self.last_action, |
| object_descriptor, |
| multi_body_feats, |
| ]).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": |
| |
| 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 = ( |
| 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 |
| ) |
|
|
| |
| fallen = bool(crumbs_on < 5) |
| dist_xy = mean_dist |
| vel_xy = mean_vel |
| else: |
| |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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) |
| |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| @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) |
|
|
| |
| adv_std = advantages.std() |
| if adv_std > 1e-8: |
| advantages = (advantages - advantages.mean()) / (adv_std + 1e-8) |
|
|
| |
| 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] |
|
|
| |
| 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 = 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 |
|
|
|
|
| |
| |
| |
|
|
| 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)), |
| )) |
| |
| 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() |
|
|
| |
| 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] |
|
|
| |
| 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 |
|
|
|
|
| |
| |
| |
|
|
| 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) |
|
|
| |
| 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() |
|
|
|
|
| |
| |
| |
|
|
| 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 |
|
|
| |
| 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) |
|
|
| |
| 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) |
|
|
| |
| if args.export_onnx: |
| export_onnx(trainer.policy, Path(args.export_onnx)) |
| return |
|
|
| |
| if args.human_view or HUMAN_VIEW: |
| human_view(trainer.policy, specific_object=args.object or None) |
| return |
|
|
| |
| if args.eval: |
| evaluate_policy(trainer.policy, eval_objects=active_objects) |
| return |
|
|
| |
| 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 |
|
|
| |
| 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) |
|
|
| |
| 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}" |
| ) |
|
|
| |
| 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}") |
|
|
| |
| 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() |
|
|