Harley-ml commited on
Commit
e217950
·
verified ·
1 Parent(s): 8796ba9

Upload 2 files

Browse files
Files changed (2) hide show
  1. balance_plate_rl.py +1666 -0
  2. inference.py +474 -0
balance_plate_rl.py ADDED
@@ -0,0 +1,1666 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ PlateBalance RL: Advanced MuJoCo + PyTorch PPO Object Balancing System
4
+
5
+ A high-performance MuJoCo + PyTorch PPO reinforcement learning system where
6
+ a 2-axis tilting plate learns to balance 21 diverse rigid and multi-body objects:
7
+ - Standard Shapes: sphere, disk, egg, hollow cup, coin, stick, tall block,
8
+ triangular prism, cube block, puck.
9
+ - Harder Challenge Shapes: cone, horizontal capsule, ramp wedge, tetrahedron,
10
+ long flat bar, cross/plus, asymmetric L-shape, wide tile block, heavy bowling ball,
11
+ and off-center mass block.
12
+ - Multi-Body Crumbling Cookie: A breakable cookie with 5 independent physical crumb
13
+ fragments that scatter and slide independently, challenging the agent to keep
14
+ EVERY individual crumb on the plate!
15
+
16
+ Installation:
17
+ pip install -U torch numpy mujoco imageio imageio-ffmpeg
18
+
19
+ Usage:
20
+ # Train policy on all objects (default)
21
+ python balance_plate_rl.py
22
+
23
+ # Train or evaluate on a specific object (e.g., crumbling cookie or cone)
24
+ python balance_plate_rl.py --object cookie
25
+ python balance_plate_rl.py --object cone
26
+
27
+ # Resume training from latest checkpoint
28
+ python balance_plate_rl.py --resume
29
+
30
+ # Run full benchmark evaluation across all 21 objects
31
+ python balance_plate_rl.py --eval --checkpoint runs/plate_balance_v1/checkpoints/latest
32
+
33
+ # Interactive real-time 3D viewer (watch the cookie crumble and balance in real-time)
34
+ python balance_plate_rl.py --human-view --object cookie --checkpoint runs/plate_balance_v1/checkpoints/latest
35
+
36
+ # Record evaluation video
37
+ python balance_plate_rl.py --record-video --checkpoint runs/plate_balance_v1/checkpoints/latest
38
+ """
39
+
40
+ from __future__ import annotations
41
+
42
+ import argparse
43
+ import json
44
+ import math
45
+ import os
46
+ import random
47
+ import re
48
+ import sys
49
+ import time
50
+ from dataclasses import dataclass
51
+ from pathlib import Path
52
+ from typing import Any, Dict, List, Optional, Tuple
53
+
54
+ import numpy as np
55
+ import torch
56
+ import torch.nn as nn
57
+ from torch.distributions import Normal
58
+
59
+ try:
60
+ import mujoco
61
+ except ImportError as exc:
62
+ raise SystemExit("MuJoCo is required. Install with: pip install -U mujoco") from exc
63
+
64
+ try:
65
+ import imageio.v2 as imageio
66
+ except ImportError:
67
+ try:
68
+ import imageio
69
+ except ImportError:
70
+ imageio = None # Optional for headless training/eval
71
+
72
+
73
+ # =============================================================================
74
+ # CONFIGURATION — DEFAULT VALUES
75
+ # =============================================================================
76
+
77
+ # ----------------------------- Experiment -----------------------------------
78
+ SEED = 42
79
+ EXPERIMENT_NAME = "plate_balance_v1"
80
+ OUTPUT_DIR = "runs"
81
+ RESUME = True
82
+ RESUME_PATH = r"runs\plate_balance_v1\checkpoints\step_03825064_final" # Empty = automatically use latest checkpoint.
83
+ DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
84
+
85
+ # ----------------------------- Architecture ---------------------------------
86
+ OBSERVATION_SIZE = 64
87
+ HIDDEN_SIZE = 128 # First/second hidden layer width.
88
+ INTERMEDIATE_SIZE = 128 # Third hidden layer width.
89
+ BOTTLENECK_SIZE = 64 # Fourth hidden layer width.
90
+ ACTION_SIZE = 2 # Roll torque, Pitch torque
91
+
92
+ # Architecture: 64 -> 192 -> 192 -> 64 -> {actor: 2, critic: 1}
93
+ ACTOR_LOG_STD_INIT = -0.75
94
+ ORTHOGONAL_INIT = True
95
+
96
+ # ----------------------------- PPO ------------------------------------------
97
+ NUM_ENVS = 8
98
+ ROLLOUT_STEPS = 512
99
+ PPO_EPOCHS = 6
100
+ MINIBATCH_SIZE = 2048
101
+ GAMMA = 0.99
102
+ GAE_LAMBDA = 0.95
103
+ CLIP_COEF = 0.20
104
+ VALUE_COEF = 0.50
105
+ ENTROPY_COEF = 0.005
106
+ MAX_GRAD_NORM = 0.50
107
+ LEARNING_RATE = 3e-4
108
+ ADAM_EPS = 1e-5
109
+ ANNEAL_LR = True # Linearly anneal learning rate to 0.
110
+
111
+ # ----------------------------- Training -------------------------------------
112
+ NUM_EPISODES = 30_000
113
+ MAX_EPISODE_STEPS = 10_000
114
+ LOGGING_STEPS = 32240
115
+ SAVE_STEPS = 1_000_000
116
+
117
+ # ----------------------------- Video / Visualization ------------------------
118
+ VISUALIZE = True # True = save rollout videos periodically during training.
119
+ VIDEO_EVERY_STEPS = 1_000_000
120
+ VIDEO_LENGTH_STEPS = 3000
121
+ VIDEO_FPS = 60
122
+ RENDER_WIDTH = 640
123
+ RENDER_HEIGHT = 480
124
+ HUMAN_VIEW = False # Optional interactive MuJoCo viewer.
125
+
126
+ # ----------------------------- Physics & Environment ------------------------
127
+ PHYSICS_TIMESTEP = 0.0025 # 400 Hz physics simulation.
128
+ CONTROL_DECIMATION = 4 # 4 physics sub-steps per control step => 100 Hz RL control.
129
+ GRAVITY = 9.81
130
+ PLATE_HALF_SIZE = 0.90 # Half-width and half-length of square plate (meters).
131
+ PLATE_THICKNESS = 0.08 # Half-thickness in z (box geom size z = 0.08).
132
+ PLATE_FRICTION = (1.0, 0.02, 0.002) # Sliding, torsional, and rolling friction.
133
+ PLATE_DAMPING = 0.08
134
+ MAX_PLATE_ANGLE = math.radians(15.0) # Maximum tilt angle in radians (~0.2618 rad).
135
+ MAX_PLATE_TORQUE = 15.0 # Maximum plate torque in N*m.
136
+
137
+ OBJECT_START_HEIGHT = 0.02 # Gentle initial clearance above plate when spawning (2cm).
138
+ OBJECT_START_POS_RANGE = 0.35 # Spawn xy radius from plate center.
139
+ OBJECT_START_VELOCITY = 0.15 # Initial linear velocity standard deviation.
140
+ OBJECT_START_ANGULAR_VELOCITY = 0.20 # Initial angular velocity standard deviation.
141
+
142
+ # Curriculum & Domain Randomization
143
+ RANDOMIZE_MASS = True
144
+ RANDOMIZE_FRICTION = True
145
+ RANDOMIZE_OBJECT_SIZE = True
146
+ MASS_LOG_RANGE = (0.50, 2.00) # Multiplicative log-uniform mass range.
147
+ FRICTION_RANGE = (0.45, 1.25) # Multiplicative friction coefficient range.
148
+ SIZE_RANGE = (0.85, 1.15) # Multiplicative linear geometry scale range.
149
+ RANDOMIZE_PLATE_ANGLE = True
150
+ INITIAL_PLATE_ANGLE_RANGE = math.radians(3.0)
151
+
152
+ # Reward Shaping
153
+ CENTER_REWARD_SCALE = 2.0 # Reward for centering object on plate.
154
+ VELOCITY_REWARD_SCALE = 0.40 # Reward for damping object linear velocity.
155
+ ANGLE_PENALTY_SCALE = 0.20 # Penalty for excessive plate tilt.
156
+ ACTION_PENALTY_SCALE = 0.005 # Penalty for excessive torque commands.
157
+ ACTION_RATE_PENALTY_SCALE = 0.01 # Penalty for rapid actuator jerk (smooth control).
158
+ EDGE_PENALTY_SCALE = 0.50 # Progressive penalty near plate perimeter.
159
+ FALL_PENALTY = -10.0 # Terminal penalty when object falls off.
160
+ SURVIVAL_REWARD = 0.05 # Step reward for keeping object on plate.
161
+
162
+ # Fall detection margins
163
+ FALL_MARGIN = 0.10 # Beyond plate half-size + margin triggers fall.
164
+ FALL_Z_THRESHOLD = 0.90 # Height below which object is deemed fallen.
165
+
166
+ # ----------------------------- Object Classification -----------------------
167
+ STANDARD_OBJECT_TYPES = [
168
+ "sphere",
169
+ "disk",
170
+ "egg",
171
+ "cup",
172
+ "coin",
173
+ "stick",
174
+ "tall",
175
+ "triangle",
176
+ "block",
177
+ "puck",
178
+ ]
179
+
180
+ HARDER_OBJECT_TYPES = [
181
+ "cone",
182
+ "capsule",
183
+ "wedge",
184
+ "tetra",
185
+ "flat_bar",
186
+ "cross",
187
+ "L_shape",
188
+ "wide_block",
189
+ "heavy_ball",
190
+ "offcenter_block",
191
+ ]
192
+
193
+ MULTI_BODY_OBJECT_TYPES = [
194
+ "cookie", # Multi-body crumbling cookie with 5 independent crumb fragments
195
+ ]
196
+ OBJECT_TYPES = STANDARD_OBJECT_TYPES + HARDER_OBJECT_TYPES + MULTI_BODY_OBJECT_TYPES
197
+ EVAL_OBJECT_TYPES = OBJECT_TYPES.copy()
198
+
199
+ PRINT_OBJECT_COUNTS = True
200
+ PRINT_HYPERPARAMS = True
201
+
202
+
203
+ # =============================================================================
204
+ # UTILITIES & REPRODUCIBILITY
205
+ # =============================================================================
206
+
207
+ def set_seed(seed: int) -> None:
208
+ """Set random seeds across Python, NumPy, and PyTorch for reproducibility."""
209
+ random.seed(seed)
210
+ np.random.seed(seed)
211
+ torch.manual_seed(seed)
212
+ if torch.cuda.is_available():
213
+ torch.cuda.manual_seed_all(seed)
214
+ torch.backends.cudnn.deterministic = True
215
+ torch.backends.cudnn.benchmark = False
216
+
217
+
218
+ def extract_step_number(path: Path) -> int:
219
+ """Safely extract the integer step number from checkpoint folder names."""
220
+ matches = re.findall(r"\d+", path.name)
221
+ if matches:
222
+ return int(matches[0])
223
+ return -1
224
+
225
+
226
+ def latest_checkpoint(root: Path) -> Optional[Path]:
227
+ """Find the latest checkpoint directory in a robust, crash-free manner."""
228
+ if not root.exists():
229
+ return None
230
+ latest_alias = root / "latest"
231
+ if latest_alias.exists() and (latest_alias / "model.pt").exists():
232
+ return latest_alias
233
+
234
+ candidates = [p for p in root.glob("step_*") if p.is_dir() and (p / "model.pt").exists()]
235
+ if not candidates:
236
+ return None
237
+ candidates.sort(key=extract_step_number)
238
+ return candidates[-1]
239
+
240
+
241
+ def safe_float(x: Any) -> float:
242
+ return float(np.asarray(x).item())
243
+
244
+
245
+ # =============================================================================
246
+ # INERTIA CALCULATIONS & PHYSICAL DEFINITIONS
247
+ # =============================================================================
248
+
249
+ OBJECT_INFO: Dict[str, Dict[str, float]] = {
250
+ # Baseline 10 Shapes
251
+ "sphere": {"radius": 0.16, "mass": 0.55},
252
+ "disk": {"radius": 0.22, "height": 0.055, "mass": 0.60},
253
+ "egg": {"radius": 0.17, "height": 0.34, "mass": 0.58},
254
+ "cup": {"radius": 0.18, "height": 0.25, "mass": 0.48},
255
+ "coin": {"radius": 0.18, "height": 0.025, "mass": 0.25},
256
+ "stick": {"radius": 0.035, "length": 0.48, "mass": 0.22},
257
+ "tall": {"radius": 0.045, "height": 0.52, "mass": 0.35},
258
+ "triangle": {"size": 0.34, "height": 0.16, "mass": 0.50},
259
+ "block": {"size": 0.28, "mass": 0.70},
260
+ "puck": {"radius": 0.25, "height": 0.08, "mass": 0.70},
261
+
262
+ # Harder 10 Shapes
263
+ "cone": {"radius": 0.16, "height": 0.32, "mass": 0.50},
264
+ "capsule": {"radius": 0.065, "length": 0.36, "mass": 0.45},
265
+ "wedge": {"size": 0.36, "height": 0.20, "mass": 0.55},
266
+ "tetra": {"size": 0.36, "height": 0.30, "mass": 0.40},
267
+ "flat_bar": {"length": 0.76, "height": 0.04, "mass": 0.60},
268
+ "cross": {"size": 0.52, "height": 0.08, "mass": 0.65},
269
+ "L_shape": {"size": 0.32, "height": 0.08, "mass": 0.60},
270
+ "wide_block": {"size": 0.64, "height": 0.04, "mass": 0.75},
271
+ "heavy_ball": {"radius": 0.16, "mass": 3.50},
272
+ "offcenter_block": {"size": 0.36, "height": 0.20, "mass": 0.75},
273
+
274
+ # Multi-Body Crumbling Cookie Challenge (5 Crumb Fragments)
275
+ "cookie": {"radius": 0.22, "height": 0.04, "mass": 0.45},
276
+ }
277
+
278
+
279
+ def cube_inertia(m: float, sx: float, sy: float, sz: float) -> np.ndarray:
280
+ return np.array([
281
+ m * (sy * sy + sz * sz) / 12.0,
282
+ m * (sx * sx + sz * sz) / 12.0,
283
+ m * (sx * sx + sy * sy) / 12.0,
284
+ ], dtype=np.float64)
285
+
286
+
287
+ def sphere_inertia(m: float, r: float) -> np.ndarray:
288
+ i = 0.4 * m * r * r
289
+ return np.array([i, i, i], dtype=np.float64)
290
+
291
+
292
+ def cylinder_inertia(m: float, r: float, h: float) -> np.ndarray:
293
+ axial = 0.5 * m * r * r
294
+ transverse = m * (3.0 * r * r + h * h) / 12.0
295
+ return np.array([transverse, transverse, axial], dtype=np.float64)
296
+
297
+
298
+ def capsule_inertia(m: float, r: float, length: float) -> np.ndarray:
299
+ rod = max(0.55 * m, 1e-6)
300
+ cyl = max(m - rod, 1e-6)
301
+ i_trans = rod * (length * length) / 12.0 + cyl * (3.0 * r * r + length * length) / 12.0
302
+ i_axial = 0.5 * cyl * r * r + rod * r * r / 2.0
303
+ return np.array([i_trans, i_trans, i_axial], dtype=np.float64)
304
+
305
+
306
+ # Custom 3D Surface Meshes
307
+ TRIANGLE_MESH = """
308
+ <mesh name="triangle_mesh"
309
+ vertex="-0.30 -0.24 -0.08 0.30 -0.24 -0.08 0.00 0.30 -0.08
310
+ -0.30 -0.24 0.08 0.30 -0.24 0.08 0.00 0.30 0.08"
311
+ 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" />
312
+ """
313
+
314
+ CONE_MESH = """
315
+ <mesh name="cone_mesh"
316
+ vertex=" 0.00 0.00 0.18
317
+ 0.16 0.00 -0.14
318
+ 0.11 0.11 -0.14
319
+ 0.00 0.16 -0.14
320
+ -0.11 0.11 -0.14
321
+ -0.16 0.00 -0.14
322
+ -0.11 -0.11 -0.14
323
+ 0.00 -0.16 -0.14
324
+ 0.11 -0.11 -0.14
325
+ 0.00 0.00 -0.14"
326
+ 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
327
+ 9 2 1 9 3 2 9 4 3 9 5 4 9 6 5 9 7 6 9 8 7 9 1 8" />
328
+ """
329
+
330
+ WEDGE_MESH = """
331
+ <mesh name="wedge_mesh"
332
+ vertex="-0.18 -0.15 -0.10
333
+ 0.18 -0.15 -0.10
334
+ 0.18 0.15 -0.10
335
+ -0.18 0.15 -0.10
336
+ -0.18 -0.15 0.10
337
+ 0.18 -0.15 0.10"
338
+ face="0 1 2 0 2 3
339
+ 0 4 5 0 5 1
340
+ 1 5 2
341
+ 0 3 4
342
+ 3 2 5 3 5 4" />
343
+ """
344
+
345
+ TETRA_MESH = """
346
+ <mesh name="tetra_mesh"
347
+ vertex=" 0.00 0.00 0.18
348
+ 0.18 -0.12 -0.12
349
+ -0.18 -0.12 -0.12
350
+ 0.00 0.20 -0.12"
351
+ face="0 1 2 0 2 3 0 3 1 1 3 2" />
352
+ """
353
+
354
+
355
+ def build_model_xml() -> str:
356
+ """Build and return the comprehensive MuJoCo XML containing all 21 objects and cookie crumbs."""
357
+ spawn_z = 1.05 + PLATE_THICKNESS + OBJECT_START_HEIGHT
358
+ return f"""
359
+ <mujoco model="plate_balance">
360
+ <compiler angle="radian" coordinate="local" inertiafromgeom="auto" />
361
+ <option timestep="{PHYSICS_TIMESTEP}" gravity="0 0 -{GRAVITY}" integrator="implicitfast" />
362
+
363
+ <default>
364
+ <joint damping="{PLATE_DAMPING}" armature="0.01" />
365
+ <geom solref="0.008 1" solimp="0.90 0.95 0.01" />
366
+ <motor ctrllimited="true" ctrlrange="-1.0 1.0" />
367
+ </default>
368
+
369
+ <asset>
370
+ <texture type="skybox" builtin="gradient" rgb1="0.85 0.90 0.98" rgb2="0.65 0.75 0.90" width="512" height="512" />
371
+ <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" />
372
+ <material name="grid_mat" texture="grid" texrepeat="5 5" reflectance="0.1" />
373
+ <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" />
374
+ <material name="cookie_mat" texture="cookie_tex" reflectance="0.1" />
375
+ <material name="plate_mat" rgba="0.16 0.45 0.92 1" reflectance="0.3" specular="0.5" />
376
+ <material name="object_mat" rgba="0.95 0.45 0.16 1" reflectance="0.2" specular="0.4" />
377
+ <material name="heavy_mat" rgba="0.32 0.35 0.42 1" reflectance="0.6" specular="0.8" />
378
+ {TRIANGLE_MESH}
379
+ {CONE_MESH}
380
+ {WEDGE_MESH}
381
+ {TETRA_MESH}
382
+ </asset>
383
+
384
+ <worldbody>
385
+ <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" />
386
+ <light directional="true" pos="3 3 5" dir="-0.5 -0.5 -1" diffuse="0.4 0.4 0.4" />
387
+ <geom name="floor" type="plane" size="6 6 0.1" pos="0 0 0" material="grid_mat" contype="1" conaffinity="1" />
388
+
389
+ <!-- Cinematic cameras for visualization and video recording -->
390
+ <camera name="track" pos="0 -2.4 2.2" xyaxes="1 0 0 0 0.6 0.8" />
391
+ <camera name="isometric" pos="2.0 -2.0 2.3" xyaxes="0.707 0.707 0 -0.408 0.408 0.816" />
392
+ <camera name="top_down" pos="0 0 3.3" xyaxes="1 0 0 0 1 0" />
393
+ <camera name="side" pos="-2.8 0 1.5" xyaxes="0 -1 0 0.2 0 0.98" />
394
+
395
+ <!-- 2-Axis Gimbal Tilting Plate Mechanism -->
396
+ <body name="plate_x" pos="0 0 1.05">
397
+ <joint name="plate_roll" type="hinge" axis="1 0 0" range="-{MAX_PLATE_ANGLE} {MAX_PLATE_ANGLE}" limited="true" />
398
+ <inertial pos="0 0 0" mass="4.0" diaginertia="1.0 1.0 1.0" />
399
+
400
+ <body name="plate_y">
401
+ <joint name="plate_pitch" type="hinge" axis="0 1 0" range="-{MAX_PLATE_ANGLE} {MAX_PLATE_ANGLE}" limited="true" />
402
+ <inertial pos="0 0 0" mass="4.0" diaginertia="1.0 1.0 1.0" />
403
+ <geom name="plate" type="box" size="{PLATE_HALF_SIZE} {PLATE_HALF_SIZE} {PLATE_THICKNESS}"
404
+ material="plate_mat" friction="{PLATE_FRICTION[0]} {PLATE_FRICTION[1]} {PLATE_FRICTION[2]}"
405
+ contype="1" conaffinity="1" mass="4.0" />
406
+ <site name="plate_center" pos="0 0 {PLATE_THICKNESS + 0.005}" size="0.015" rgba="1 1 1 0.6" />
407
+ </body>
408
+ </body>
409
+
410
+ <!-- Primary candidate object body (for single rigid objects) -->
411
+ <body name="object" pos="0 0 {spawn_z}">
412
+ <freejoint name="object_free" />
413
+
414
+ <!-- Baseline 10 Shapes -->
415
+ <geom name="g_sphere" type="sphere" size="0.16" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" />
416
+ <geom name="g_disk" type="cylinder" size="0.22 0.0275" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" />
417
+ <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" />
418
+
419
+ <!-- Hollow cup: base + 4 walls -->
420
+ <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" />
421
+ <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" />
422
+ <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" />
423
+ <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" />
424
+ <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" />
425
+
426
+ <geom name="g_coin" type="cylinder" size="0.18 0.0125" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" />
427
+ <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" />
428
+ <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" />
429
+ <geom name="g_triangle" type="mesh" mesh="triangle_mesh" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" />
430
+ <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" />
431
+ <geom name="g_puck" type="cylinder" size="0.25 0.04" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" />
432
+
433
+ <!-- Harder 10 Shapes -->
434
+ <geom name="g_cone" type="mesh" mesh="cone_mesh" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" />
435
+ <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" />
436
+ <geom name="g_wedge" type="mesh" mesh="wedge_mesh" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" />
437
+ <geom name="g_tetra" type="mesh" mesh="tetra_mesh" rgba="0.95 0.45 0.16 1" contype="1" conaffinity="1" />
438
+ <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" />
439
+ <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" />
440
+ <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" />
441
+ <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" />
442
+ <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" />
443
+ <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" />
444
+ <geom name="g_heavy_ball" type="sphere" size="0.16" material="heavy_mat" contype="1" conaffinity="1" />
445
+ <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" />
446
+ </body>
447
+
448
+ <!-- 5 Crumb bodies for the Crumbling Cookie challenge -->
449
+ <body name="crumb_0" pos="0 0 1.15">
450
+ <freejoint name="crumb_free_0" />
451
+ <geom name="g_crumb_0" type="cylinder" size="0.08 0.02" material="cookie_mat" contype="1" conaffinity="1" mass="0.15" />
452
+ </body>
453
+ <body name="crumb_1" pos="0.09 0 1.15">
454
+ <freejoint name="crumb_free_1" />
455
+ <geom name="g_crumb_1" type="box" size="0.035 0.035 0.02" material="cookie_mat" contype="1" conaffinity="1" mass="0.08" />
456
+ </body>
457
+ <body name="crumb_2" pos="-0.09 0 1.15">
458
+ <freejoint name="crumb_free_2" />
459
+ <geom name="g_crumb_2" type="cylinder" size="0.035 0.02" material="cookie_mat" contype="1" conaffinity="1" mass="0.08" />
460
+ </body>
461
+ <body name="crumb_3" pos="0 0.09 1.15">
462
+ <freejoint name="crumb_free_3" />
463
+ <geom name="g_crumb_3" type="box" size="0.04 0.03 0.02" material="cookie_mat" contype="1" conaffinity="1" mass="0.08" />
464
+ </body>
465
+ <body name="crumb_4" pos="0 -0.09 1.15">
466
+ <freejoint name="crumb_free_4" />
467
+ <geom name="g_crumb_4" type="sphere" size="0.03" material="cookie_mat" contype="1" conaffinity="1" mass="0.06" />
468
+ </body>
469
+ </worldbody>
470
+
471
+ <actuator>
472
+ <motor name="roll_motor" joint="plate_roll" gear="{MAX_PLATE_TORQUE}" />
473
+ <motor name="pitch_motor" joint="plate_pitch" gear="{MAX_PLATE_TORQUE}" />
474
+ </actuator>
475
+ </mujoco>
476
+ """
477
+
478
+
479
+ # =============================================================================
480
+ # MUJOCO ENVIRONMENT (SINGLE & MULTI-BODY CRUMBLING COOKIE)
481
+ # =============================================================================
482
+
483
+ class PlateBalanceEnv:
484
+ """
485
+ Robust single-instance MuJoCo environment supporting 21 single & multi-body
486
+ objects with accurate physical scaling and multi-crumb tracking.
487
+ """
488
+
489
+ OBJECT_GEOM_NAMES = {
490
+ # Baseline 10
491
+ "sphere": ["g_sphere"],
492
+ "disk": ["g_disk"],
493
+ "egg": ["g_egg"],
494
+ "cup": ["g_cup_bottom", "g_cup_w1", "g_cup_w2", "g_cup_w3", "g_cup_w4"],
495
+ "coin": ["g_coin"],
496
+ "stick": ["g_stick"],
497
+ "tall": ["g_tall"],
498
+ "triangle": ["g_triangle"],
499
+ "block": ["g_block"],
500
+ "puck": ["g_puck"],
501
+ # Harder 10
502
+ "cone": ["g_cone"],
503
+ "capsule": ["g_capsule"],
504
+ "wedge": ["g_wedge"],
505
+ "tetra": ["g_tetra"],
506
+ "flat_bar": ["g_flat_bar"],
507
+ "cross": ["g_cross_1", "g_cross_2"],
508
+ "L_shape": ["g_lshape_1", "g_lshape_2"],
509
+ "wide_block": ["g_wide_block"],
510
+ "heavy_ball": ["g_heavy_ball"],
511
+ "offcenter_block": ["g_offcenter_block"],
512
+ }
513
+
514
+ def __init__(self, seed: int, render: bool = False, active_objects: Optional[List[str]] = None):
515
+ self.rng = np.random.default_rng(seed)
516
+ self.model = mujoco.MjModel.from_xml_string(build_model_xml())
517
+ self.data = mujoco.MjData(self.model)
518
+ self.model.opt.timestep = PHYSICS_TIMESTEP
519
+ self.render_enabled = render
520
+ self.renderer: Optional[mujoco.Renderer] = None
521
+ self.active_objects = list(active_objects) if active_objects else list(OBJECT_TYPES)
522
+
523
+ if render:
524
+ self.renderer = mujoco.Renderer(self.model, height=RENDER_HEIGHT, width=RENDER_WIDTH)
525
+
526
+ # Primary Single Object
527
+ self.qpos_object = mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_JOINT, "object_free")
528
+ self.object_qpos_addr = self.model.jnt_qposadr[self.qpos_object]
529
+ self.object_dof_addr = self.model.jnt_dofadr[self.qpos_object]
530
+ self.object_body_id = mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_BODY, "object")
531
+
532
+ # Plate Joints and Sites
533
+ self.plate_roll_joint = mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_JOINT, "plate_roll")
534
+ self.plate_pitch_joint = mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_JOINT, "plate_pitch")
535
+ self.plate_roll_qpos = self.model.jnt_qposadr[self.plate_roll_joint]
536
+ self.plate_pitch_qpos = self.model.jnt_qposadr[self.plate_pitch_joint]
537
+ self.plate_roll_dof = self.model.jnt_dofadr[self.plate_roll_joint]
538
+ self.plate_pitch_dof = self.model.jnt_dofadr[self.plate_pitch_joint]
539
+ self.plate_body_id = mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_BODY, "plate_y")
540
+ self.plate_center_site = mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_SITE, "plate_center")
541
+ self.plate_geom_id = mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_GEOM, "plate")
542
+
543
+ # Cookie Crumb Bodies (5 fragments)
544
+ self.crumb_body_ids = [mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_BODY, f"crumb_{i}") for i in range(5)]
545
+ 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)]
546
+ 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)]
547
+ self.crumb_geom_ids = [mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_GEOM, f"g_crumb_{i}") for i in range(5)]
548
+
549
+ # Map candidate geoms
550
+ self.object_geom_ids = {
551
+ name: [mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_GEOM, g) for g in geoms]
552
+ for name, geoms in self.OBJECT_GEOM_NAMES.items()
553
+ }
554
+
555
+ # Cache baseline sizes and positions for physical scaling
556
+ self.base_geom_sizes: Dict[int, np.ndarray] = {}
557
+ self.base_geom_positions: Dict[int, np.ndarray] = {}
558
+ for ids in self.object_geom_ids.values():
559
+ for gid in ids:
560
+ self.base_geom_sizes[gid] = self.model.geom_size[gid].copy()
561
+ self.base_geom_positions[gid] = self.model.geom_pos[gid].copy()
562
+
563
+ for gid in self.crumb_geom_ids:
564
+ self.base_geom_sizes[gid] = self.model.geom_size[gid].copy()
565
+ self.base_geom_positions[gid] = self.model.geom_pos[gid].copy()
566
+
567
+ self.steps = 0
568
+ self.episode_reward = 0.0
569
+ self.current_object = "sphere"
570
+ self.object_mass = 0.55
571
+ self.object_scale = 1.0
572
+ self.base_friction = 0.90
573
+ self.last_action = np.zeros(ACTION_SIZE, dtype=np.float32)
574
+
575
+ def _set_object_geometry(self, object_name: str) -> None:
576
+ """Activate the selected single or multi-body object and scale its physics."""
577
+ # Deactivate all single object geoms
578
+ for ids in self.object_geom_ids.values():
579
+ for gid in ids:
580
+ self.model.geom_size[gid, :] = 1e-6
581
+ self.model.geom_pos[gid, :] = np.array([0.0, 0.0, -999.0])
582
+ self.model.geom_rgba[gid, 3] = 0.0
583
+
584
+ # Deactivate all crumb geoms by default
585
+ for gid in self.crumb_geom_ids:
586
+ self.model.geom_size[gid, :] = 1e-6
587
+ self.model.geom_pos[gid, :] = np.array([0.0, 0.0, -999.0])
588
+ self.model.geom_rgba[gid, 3] = 0.0
589
+
590
+ s = self.object_scale
591
+ info = OBJECT_INFO[object_name]
592
+ m = self.object_mass
593
+
594
+ if object_name == "cookie":
595
+ # Activate 5 cookie crumb fragments
596
+ for i, gid in enumerate(self.crumb_geom_ids):
597
+ self.model.geom_friction[gid, 0] = self.base_friction
598
+ self.model.geom_rgba[gid, :] = np.array([0.82, 0.58, 0.32, 1.0])
599
+ self.model.geom_size[gid, :] = self.base_geom_sizes[gid] * s
600
+ self.model.geom_pos[gid, :] = self.base_geom_positions[gid] * s
601
+ self.model.body_mass[self.crumb_body_ids[i]] = (m / 5.0)
602
+ return
603
+
604
+ # Activate single candidate object
605
+ active = self.object_geom_ids[object_name]
606
+ 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])
607
+ for gid in active:
608
+ self.model.geom_friction[gid, 0] = self.base_friction
609
+ self.model.geom_rgba[gid, :] = rgba
610
+ self.model.geom_size[gid, :] = self.base_geom_sizes[gid] * s
611
+ self.model.geom_pos[gid, :] = self.base_geom_positions[gid] * s
612
+
613
+ # Reset body center of mass position (offset for offcenter_block)
614
+ if object_name == "offcenter_block":
615
+ self.model.body_ipos[self.object_body_id] = np.array([0.08 * s, 0.06 * s, -0.02 * s])
616
+ else:
617
+ self.model.body_ipos[self.object_body_id] = np.array([0.0, 0.0, 0.0])
618
+
619
+ # Compute accurate 3D moment of inertia tensor
620
+ if object_name in ("sphere", "heavy_ball"):
621
+ inertia = sphere_inertia(m, info["radius"] * s)
622
+ elif object_name in ("disk", "coin", "puck"):
623
+ inertia = cylinder_inertia(m, info["radius"] * s, info["height"] * s)
624
+ elif object_name == "egg":
625
+ rx = info["radius"] * 0.85 * s
626
+ rz = info["height"] * 0.55 * s
627
+ inertia = np.array([
628
+ m * (rx * rx + rz * rz) / 5.0,
629
+ m * (rx * rx + rz * rz) / 5.0,
630
+ m * (2.0 * rx * rx) / 5.0,
631
+ ])
632
+ elif object_name == "cup":
633
+ inertia = cylinder_inertia(m, info["radius"] * s, info["height"] * s)
634
+ elif object_name in ("stick", "capsule"):
635
+ inertia = capsule_inertia(m, info["radius"] * s, info["length"] * s)
636
+ elif object_name == "tall":
637
+ side = 0.09 * s
638
+ height = info["height"] * s
639
+ inertia = cube_inertia(m, side, side, height)
640
+ elif object_name == "cone":
641
+ r = info["radius"] * s
642
+ h = info["height"] * s
643
+ i_trans = 0.6 * m * (0.25 * r * r + h * h)
644
+ i_axial = 0.3 * m * r * r
645
+ inertia = np.array([i_trans, i_trans, i_axial])
646
+ elif object_name in ("triangle", "wedge"):
647
+ side = info["size"] * s
648
+ height = info["height"] * s
649
+ inertia = cube_inertia(m, side, side, height) * 1.10
650
+ elif object_name == "tetra":
651
+ side = info["size"] * s
652
+ i_val = (m * side * side) / 20.0
653
+ inertia = np.array([i_val, i_val, i_val])
654
+ elif object_name == "flat_bar":
655
+ inertia = cube_inertia(m, info["length"] * s, 0.10 * s, info["height"] * s)
656
+ elif object_name == "cross":
657
+ side = info["size"] * s
658
+ h = info["height"] * s
659
+ inertia = cube_inertia(m, side, side, h) * 0.70
660
+ elif object_name == "L_shape":
661
+ side = info["size"] * s
662
+ h = info["height"] * s
663
+ inertia = cube_inertia(m, side, side, h) * 0.85
664
+ elif object_name == "wide_block":
665
+ side = info["size"] * s
666
+ h = info["height"] * s
667
+ inertia = cube_inertia(m, side, side, h)
668
+ elif object_name == "offcenter_block":
669
+ side = info["size"] * s
670
+ h = info["height"] * s
671
+ inertia = cube_inertia(m, side, side, h) * 1.20
672
+ else: # block
673
+ side = info["size"] * s
674
+ inertia = cube_inertia(m, side, side, side)
675
+
676
+ self.model.body_mass[self.object_body_id] = m
677
+ self.model.body_inertia[self.object_body_id, :] = np.maximum(inertia, 1e-6)
678
+
679
+ def reset(self, specific_object: Optional[str] = None) -> np.ndarray:
680
+ """Reset the environment state for a new episode."""
681
+ if specific_object and specific_object in OBJECT_INFO:
682
+ self.current_object = specific_object
683
+ else:
684
+ self.current_object = str(self.rng.choice(self.active_objects))
685
+
686
+ info = OBJECT_INFO[self.current_object]
687
+ self.object_mass = info["mass"]
688
+ if RANDOMIZE_MASS:
689
+ mult = float(np.exp(self.rng.uniform(math.log(MASS_LOG_RANGE[0]), math.log(MASS_LOG_RANGE[1]))))
690
+ self.object_mass *= mult
691
+
692
+ self.object_scale = float(self.rng.uniform(*SIZE_RANGE)) if RANDOMIZE_OBJECT_SIZE else 1.0
693
+ self.base_friction = float(self.rng.uniform(*FRICTION_RANGE)) if RANDOMIZE_FRICTION else 0.90
694
+ self._set_object_geometry(self.current_object)
695
+
696
+ mujoco.mj_resetData(self.model, self.data)
697
+
698
+ # Initial plate tilt randomization
699
+ if RANDOMIZE_PLATE_ANGLE:
700
+ self.data.qpos[self.plate_roll_qpos] = self.rng.uniform(-INITIAL_PLATE_ANGLE_RANGE, INITIAL_PLATE_ANGLE_RANGE)
701
+ self.data.qpos[self.plate_pitch_qpos] = self.rng.uniform(-INITIAL_PLATE_ANGLE_RANGE, INITIAL_PLATE_ANGLE_RANGE)
702
+ else:
703
+ self.data.qpos[self.plate_roll_qpos] = 0.0
704
+ self.data.qpos[self.plate_pitch_qpos] = 0.0
705
+
706
+ x = float(self.rng.uniform(-OBJECT_START_POS_RANGE, OBJECT_START_POS_RANGE))
707
+ y = float(self.rng.uniform(-OBJECT_START_POS_RANGE, OBJECT_START_POS_RANGE))
708
+
709
+ if self.current_object == "cookie":
710
+ # Park single object far away
711
+ base_s = self.object_qpos_addr
712
+ self.data.qpos[base_s:base_s + 7] = np.array([0.0, 0.0, -999.0, 1.0, 0.0, 0.0, 0.0])
713
+ self.data.qvel[self.object_dof_addr:self.object_dof_addr + 6] = 0.0
714
+
715
+ # Spawn 5 cookie crumb fragments on the plate with scatter offsets & velocities
716
+ crumb_offsets = [
717
+ (0.0, 0.0),
718
+ (0.08 * self.object_scale, 0.0),
719
+ (-0.08 * self.object_scale, 0.0),
720
+ (0.0, 0.08 * self.object_scale),
721
+ (0.0, -0.08 * self.object_scale),
722
+ ]
723
+ for i in range(5):
724
+ dx, dy = crumb_offsets[i]
725
+ dx += float(self.rng.uniform(-0.02, 0.02))
726
+ dy += float(self.rng.uniform(-0.02, 0.02))
727
+ c_addr = self.crumb_qpos_addrs[i]
728
+ c_dof = self.crumb_dof_addrs[i]
729
+ z_c = 1.05 + PLATE_THICKNESS + 0.02 + OBJECT_START_HEIGHT
730
+ self.data.qpos[c_addr:c_addr + 7] = np.array([x + dx, y + dy, z_c, 1.0, 0.0, 0.0, 0.0])
731
+ self.data.qvel[c_dof:c_dof + 6] = np.concatenate([
732
+ self.rng.normal(0.0, OBJECT_START_VELOCITY, 3),
733
+ self.rng.normal(0.0, OBJECT_START_ANGULAR_VELOCITY, 3),
734
+ ])
735
+ else:
736
+ # Park all crumb fragments far away
737
+ for i in range(5):
738
+ c_addr = self.crumb_qpos_addrs[i]
739
+ c_dof = self.crumb_dof_addrs[i]
740
+ self.data.qpos[c_addr:c_addr + 7] = np.array([0.0, 0.0, -999.0, 1.0, 0.0, 0.0, 0.0])
741
+ self.data.qvel[c_dof:c_dof + 6] = 0.0
742
+
743
+ # Spawn single object with safe clearance
744
+ base = self.object_qpos_addr
745
+ if self.current_object in ("sphere", "heavy_ball"):
746
+ half_height = info["radius"] * self.object_scale
747
+ elif self.current_object in ("stick", "capsule", "flat_bar"):
748
+ half_height = (info["length"] / 2.0) * self.object_scale if "length" in info else (info["height"] / 2.0) * self.object_scale
749
+ elif self.current_object in ("block", "cross", "L_shape", "wide_block", "wedge", "tetra", "offcenter_block"):
750
+ half_height = (info.get("height", info.get("size", 0.20)) / 2.0) * self.object_scale
751
+ else:
752
+ half_height = (info["height"] / 2.0) * self.object_scale
753
+
754
+ z = 1.05 + PLATE_THICKNESS + half_height + OBJECT_START_HEIGHT
755
+ self.data.qpos[base:base + 7] = np.array([x, y, z, 1.0, 0.0, 0.0, 0.0])
756
+ self.data.qvel[self.object_dof_addr:self.object_dof_addr + 6] = np.concatenate([
757
+ self.rng.normal(0.0, OBJECT_START_VELOCITY, 3),
758
+ self.rng.normal(0.0, OBJECT_START_ANGULAR_VELOCITY, 3),
759
+ ])
760
+
761
+ self.data.ctrl[:] = 0.0
762
+ mujoco.mj_forward(self.model, self.data)
763
+ self.steps = 0
764
+ self.episode_reward = 0.0
765
+ self.last_action = np.zeros(ACTION_SIZE, dtype=np.float32)
766
+ return self._observation()
767
+
768
+ def _observation(self) -> np.ndarray:
769
+ """Construct the normalized 64-dimensional observation vector."""
770
+ plate_angles = np.array([self.data.qpos[self.plate_roll_qpos], self.data.qpos[self.plate_pitch_qpos]], dtype=np.float32)
771
+ plate_vel = np.array([self.data.qvel[self.plate_roll_dof], self.data.qvel[self.plate_pitch_dof]], dtype=np.float32)
772
+ plate_surface_center = self.data.site_xpos[self.plate_center_site]
773
+
774
+ info = OBJECT_INFO[self.current_object]
775
+ object_descriptor = np.array([
776
+ self.object_mass,
777
+ self.object_scale,
778
+ self.base_friction,
779
+ info.get("radius", info.get("size", 0.10)),
780
+ info.get("height", info.get("length", 0.10)),
781
+ ], dtype=np.float32)
782
+
783
+ if self.current_object == "cookie":
784
+ # Multi-body crumb tracking
785
+ crumb_xys = [self.data.xpos[bid, :2] for bid in self.crumb_body_ids]
786
+ crumb_zs = [self.data.qpos[addr + 2] for addr in self.crumb_qpos_addrs]
787
+ on_plate_crumbs = [
788
+ p for p, z in zip(crumb_xys, crumb_zs)
789
+ if max(abs(p[0]), abs(p[1])) <= (PLATE_HALF_SIZE + FALL_MARGIN) and z >= FALL_Z_THRESHOLD
790
+ ]
791
+ if on_plate_crumbs:
792
+ centroid_xy = np.mean(on_plate_crumbs, axis=0)
793
+ mean_z = float(np.mean([z for z in crumb_zs if z >= FALL_Z_THRESHOLD]))
794
+ spread = float(np.max([np.linalg.norm(p - centroid_xy) for p in on_plate_crumbs]))
795
+ else:
796
+ centroid_xy = np.zeros(2, dtype=np.float32)
797
+ mean_z = 1.15
798
+ spread = 0.0
799
+
800
+ object_pos = np.array([centroid_xy[0], centroid_xy[1], mean_z], dtype=np.float32)
801
+ rel_xy = (centroid_xy - plate_surface_center[:2]).astype(np.float32)
802
+ object_linvel = np.mean([self.data.qvel[dof:dof+3] for dof in self.crumb_dof_addrs], axis=0).astype(np.float32)
803
+ object_angvel = np.zeros(3, dtype=np.float32)
804
+ object_quat = np.array([1.0, 0.0, 0.0, 0.0], dtype=np.float32)
805
+
806
+ # Crumb specific relative features
807
+ crumb_ratio = float(len(on_plate_crumbs) / 5.0)
808
+ crumb_rel_1 = (crumb_xys[1] - centroid_xy).astype(np.float32)
809
+ crumb_rel_2 = (crumb_xys[2] - centroid_xy).astype(np.float32)
810
+ multi_body_feats = np.concatenate([
811
+ np.array([1.0, spread, crumb_ratio], dtype=np.float32),
812
+ crumb_rel_1,
813
+ crumb_rel_2,
814
+ ])
815
+ else:
816
+ base_q = self.object_qpos_addr
817
+ base_v = self.object_dof_addr
818
+ object_pos = self.data.qpos[base_q:base_q + 3].astype(np.float32)
819
+ object_quat = self.data.qpos[base_q + 3:base_q + 7].astype(np.float32)
820
+ object_linvel = self.data.qvel[base_v:base_v + 3].astype(np.float32)
821
+ object_angvel = self.data.qvel[base_v + 3:base_v + 6].astype(np.float32)
822
+ rel_pos = object_pos - plate_surface_center
823
+ rel_xy = rel_pos[:2].astype(np.float32)
824
+ multi_body_feats = np.zeros(7, dtype=np.float32)
825
+
826
+ state = np.concatenate([
827
+ object_pos, # 3
828
+ rel_xy, # 2
829
+ object_linvel, # 3
830
+ object_angvel, # 3
831
+ object_quat, # 4
832
+ plate_angles, # 2
833
+ plate_vel, # 2
834
+ self.data.qfrc_actuator[:2], # 2
835
+ self.last_action, # 2
836
+ object_descriptor, # 5
837
+ multi_body_feats, # 7
838
+ ]).astype(np.float32)
839
+
840
+ out = np.zeros(OBSERVATION_SIZE, dtype=np.float32)
841
+ n = min(len(state), OBSERVATION_SIZE)
842
+ out[:n] = state[:n]
843
+ out = np.clip(out, -10.0, 10.0)
844
+ return out
845
+
846
+ def step(self, action: np.ndarray) -> Tuple[np.ndarray, float, bool, bool, Dict[str, Any]]:
847
+ """Advance simulation and compute reward for single or multi-crumb cookie objects."""
848
+ action = np.asarray(action, dtype=np.float64)
849
+ action = np.clip(action, -1.0, 1.0)
850
+ self.data.ctrl[0] = action[0]
851
+ self.data.ctrl[1] = action[1]
852
+
853
+ for _ in range(CONTROL_DECIMATION):
854
+ mujoco.mj_step(self.model, self.data)
855
+
856
+ self.steps += 1
857
+ plate_ang = np.array([self.data.qpos[self.plate_roll_qpos], self.data.qpos[self.plate_pitch_qpos]])
858
+ plate_ang_mag = float(np.linalg.norm(plate_ang))
859
+ action_rate_penalty = float(np.mean((action - self.last_action) ** 2))
860
+ self.last_action = action.astype(np.float32).copy()
861
+
862
+ if self.current_object == "cookie":
863
+ # Multi-body Crumbling Cookie Evaluation
864
+ crumb_xys = [self.data.xpos[bid, :2] for bid in self.crumb_body_ids]
865
+ crumb_zs = [self.data.qpos[addr + 2] for addr in self.crumb_qpos_addrs]
866
+
867
+ on_plate_flags = [
868
+ bool(max(abs(p[0]), abs(p[1])) <= (PLATE_HALF_SIZE + FALL_MARGIN) and z >= FALL_Z_THRESHOLD)
869
+ for p, z in zip(crumb_xys, crumb_zs)
870
+ ]
871
+ crumbs_on = sum(on_plate_flags)
872
+ crumb_ratio = crumbs_on / 5.0
873
+
874
+ dists = [float(np.linalg.norm(p)) for p in crumb_xys]
875
+ max_dist = max(dists)
876
+ mean_dist = float(np.mean(dists))
877
+ box_dist = float(max(max(abs(p[0]), abs(p[1])) for p in crumb_xys))
878
+
879
+ vels = [float(np.linalg.norm(self.data.qvel[dof:dof+2])) for dof in self.crumb_dof_addrs]
880
+ mean_vel = float(np.mean(vels))
881
+
882
+ center_score = 0.5 * math.exp(-4.0 * mean_dist * mean_dist) + 0.5 * math.exp(-4.0 * max_dist * max_dist)
883
+ velocity_score = math.exp(-1.8 * mean_vel * mean_vel)
884
+ edge_fraction = np.clip(box_dist / PLATE_HALF_SIZE, 0.0, 1.5)
885
+ edge_penalty = max(0.0, float(edge_fraction) - 0.65) ** 2
886
+
887
+ # Reward scaled by the fraction of cookie crumbs kept on the plate
888
+ reward = (
889
+ SURVIVAL_REWARD * crumb_ratio
890
+ + CENTER_REWARD_SCALE * center_score
891
+ + VELOCITY_REWARD_SCALE * velocity_score
892
+ - ANGLE_PENALTY_SCALE * (plate_ang_mag / MAX_PLATE_ANGLE) ** 2
893
+ - ACTION_PENALTY_SCALE * float(np.mean(action ** 2))
894
+ - ACTION_RATE_PENALTY_SCALE * action_rate_penalty
895
+ - EDGE_PENALTY_SCALE * edge_penalty
896
+ )
897
+
898
+ # Failure occurs when any crumb falls off (or partial penalty for lost crumbs)
899
+ fallen = bool(crumbs_on < 5)
900
+ dist_xy = mean_dist
901
+ vel_xy = mean_vel
902
+ else:
903
+ # Single Rigid Object Evaluation
904
+ pos = self.data.xpos[self.object_body_id, :2]
905
+ dist_xy = float(np.linalg.norm(pos))
906
+ box_dist = float(max(abs(pos[0]), abs(pos[1])))
907
+
908
+ obj_vel = self.data.qvel[self.object_dof_addr:self.object_dof_addr + 3]
909
+ vel_xy = float(np.linalg.norm(obj_vel[:2]))
910
+
911
+ center_score = math.exp(-4.0 * dist_xy * dist_xy)
912
+ velocity_score = math.exp(-1.8 * vel_xy * vel_xy)
913
+ edge_fraction = np.clip(box_dist / PLATE_HALF_SIZE, 0.0, 1.5)
914
+ edge_penalty = max(0.0, float(edge_fraction) - 0.65) ** 2
915
+
916
+ reward = (
917
+ SURVIVAL_REWARD
918
+ + CENTER_REWARD_SCALE * center_score
919
+ + VELOCITY_REWARD_SCALE * velocity_score
920
+ - ANGLE_PENALTY_SCALE * (plate_ang_mag / MAX_PLATE_ANGLE) ** 2
921
+ - ACTION_PENALTY_SCALE * float(np.mean(action ** 2))
922
+ - ACTION_RATE_PENALTY_SCALE * action_rate_penalty
923
+ - EDGE_PENALTY_SCALE * edge_penalty
924
+ )
925
+
926
+ object_z = self.data.qpos[self.object_qpos_addr + 2]
927
+ fallen = bool(box_dist > (PLATE_HALF_SIZE + FALL_MARGIN) or object_z < FALL_Z_THRESHOLD)
928
+ crumbs_on = 1 if not fallen else 0
929
+
930
+ timeout = bool(self.steps >= MAX_EPISODE_STEPS)
931
+ terminated = fallen
932
+ truncated = timeout and not fallen
933
+
934
+ if fallen:
935
+ reward += FALL_PENALTY
936
+
937
+ self.episode_reward += reward
938
+ info = {
939
+ "object": self.current_object,
940
+ "distance": dist_xy,
941
+ "box_distance": box_dist,
942
+ "object_velocity": vel_xy,
943
+ "fallen": fallen,
944
+ "timeout": timeout,
945
+ "crumbs_on_plate": crumbs_on,
946
+ "episode_reward": self.episode_reward,
947
+ "episode_length": self.steps,
948
+ }
949
+ return self._observation(), float(reward), terminated, truncated, info
950
+
951
+ def render(self, camera: str = "track") -> np.ndarray:
952
+ """Render RGB visual frame from the specified camera."""
953
+ if not self.render_enabled:
954
+ raise RuntimeError("Environment initialized with render=False")
955
+ assert self.renderer is not None
956
+ self.renderer.update_scene(self.data, camera=camera)
957
+ return self.renderer.render().copy()
958
+
959
+ def close(self) -> None:
960
+ """Cleanly release rendering and physics resources."""
961
+ if self.renderer is not None:
962
+ self.renderer.close()
963
+ self.renderer = None
964
+
965
+
966
+ # =============================================================================
967
+ # ACTOR / CRITIC NEURAL NETWORK
968
+ # =============================================================================
969
+
970
+ class ActorCritic(nn.Module):
971
+ """
972
+ Continuous-action Actor-Critic MLP with Tanh-squashed Gaussian policy
973
+ and exact, numerically stable log-probability computation.
974
+ """
975
+
976
+ def __init__(self):
977
+ super().__init__()
978
+ self.backbone = nn.Sequential(
979
+ nn.Linear(OBSERVATION_SIZE, HIDDEN_SIZE),
980
+ nn.SiLU(),
981
+ nn.Linear(HIDDEN_SIZE, INTERMEDIATE_SIZE),
982
+ nn.SiLU(),
983
+ nn.Linear(INTERMEDIATE_SIZE, BOTTLENECK_SIZE),
984
+ nn.SiLU(),
985
+ )
986
+ self.actor = nn.Linear(BOTTLENECK_SIZE, ACTION_SIZE)
987
+ self.critic = nn.Linear(BOTTLENECK_SIZE, 1)
988
+ self.log_std = nn.Parameter(torch.full((ACTION_SIZE,), ACTOR_LOG_STD_INIT))
989
+ self._init_weights()
990
+
991
+ def _init_weights(self) -> None:
992
+ if not ORTHOGONAL_INIT:
993
+ return
994
+ for layer in self.backbone:
995
+ if isinstance(layer, nn.Linear):
996
+ nn.init.orthogonal_(layer.weight, gain=math.sqrt(2.0))
997
+ nn.init.zeros_(layer.bias)
998
+ nn.init.orthogonal_(self.actor.weight, gain=0.01)
999
+ nn.init.zeros_(self.actor.bias)
1000
+ nn.init.orthogonal_(self.critic.weight, gain=1.0)
1001
+ nn.init.zeros_(self.critic.bias)
1002
+
1003
+ def forward(self, obs: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
1004
+ h = self.backbone(obs)
1005
+ mean = self.actor(h)
1006
+ value = self.critic(h).squeeze(-1)
1007
+ std = self.log_std.exp().expand_as(mean)
1008
+ return mean, std, value
1009
+
1010
+ def get_value(self, obs: torch.Tensor) -> torch.Tensor:
1011
+ return self.critic(self.backbone(obs)).squeeze(-1)
1012
+
1013
+ def get_action_and_value(
1014
+ self, obs: torch.Tensor, raw_action: Optional[torch.Tensor] = None, deterministic: bool = False
1015
+ ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
1016
+ mean, std, value = self(obs)
1017
+ dist = Normal(mean, std)
1018
+
1019
+ if raw_action is None:
1020
+ raw_action = mean if deterministic else dist.sample()
1021
+
1022
+ squashed = torch.tanh(raw_action)
1023
+ # Numerically stable exact log-prob correction for Tanh transformation
1024
+ log_prob = dist.log_prob(raw_action).sum(-1)
1025
+ log_prob -= torch.log(torch.clamp(1.0 - squashed.pow(2), min=1e-6)).sum(-1)
1026
+ entropy = dist.entropy().sum(-1)
1027
+ return squashed, log_prob, entropy, value, raw_action
1028
+
1029
+
1030
+ # =============================================================================
1031
+ # PPO TRAINER
1032
+ # =============================================================================
1033
+
1034
+ @dataclass
1035
+ class EpisodeStats:
1036
+ reward: float = 0.0
1037
+ length: int = 0
1038
+ object_name: str = ""
1039
+ success: bool = False
1040
+ crumbs_on: int = 1
1041
+
1042
+
1043
+ class PPOTrainer:
1044
+ """High-throughput PPO Trainer with robust GAE, metric logging, and checkpointing."""
1045
+
1046
+ def __init__(self, device: Optional[str] = None):
1047
+ self.device = torch.device(device or DEVICE)
1048
+ self.policy = ActorCritic().to(self.device)
1049
+ self.optimizer = torch.optim.Adam(self.policy.parameters(), lr=LEARNING_RATE, eps=ADAM_EPS)
1050
+ self.global_steps = 0
1051
+ self.episodes = 0
1052
+ self.updates = 0
1053
+ self.last_log_step = 0
1054
+ self.last_save_step = 0
1055
+ self.last_video_step = 0
1056
+ self.object_counts: Dict[str, int] = {name: 0 for name in OBJECT_TYPES}
1057
+
1058
+ def save(self, directory: Path, latest: bool = True, extra: Optional[Dict[str, Any]] = None) -> Path:
1059
+ """Save a complete resumable checkpoint."""
1060
+ directory.mkdir(parents=True, exist_ok=True)
1061
+ model_path = directory / "model.pt"
1062
+ state_path = directory / "trainer_state.pt"
1063
+ torch.save(self.policy.state_dict(), model_path)
1064
+ trainer_state = {
1065
+ "optimizer": self.optimizer.state_dict(),
1066
+ "global_steps": self.global_steps,
1067
+ "episodes": self.episodes,
1068
+ "updates": self.updates,
1069
+ "last_log_step": self.last_log_step,
1070
+ "last_save_step": self.last_save_step,
1071
+ "last_video_step": self.last_video_step,
1072
+ "object_counts": self.object_counts,
1073
+ "torch_rng_state": torch.get_rng_state(),
1074
+ "numpy_rng_state": np.random.get_state(),
1075
+ "python_rng_state": random.getstate(),
1076
+ "config": {
1077
+ k: v for k, v in globals().items()
1078
+ if k.isupper() and isinstance(v, (int, float, str, bool, tuple, list))
1079
+ },
1080
+ "extra": extra or {},
1081
+ }
1082
+ if torch.cuda.is_available():
1083
+ trainer_state["cuda_rng_state"] = torch.cuda.get_rng_state_all()
1084
+ torch.save(trainer_state, state_path)
1085
+
1086
+ if latest:
1087
+ latest_dir = directory.parent / "latest"
1088
+ latest_dir.mkdir(parents=True, exist_ok=True)
1089
+ torch.save(self.policy.state_dict(), latest_dir / "model.pt")
1090
+ torch.save(trainer_state, latest_dir / "trainer_state.pt")
1091
+ with open(latest_dir / "config.json", "w", encoding="utf-8") as f:
1092
+ json.dump(trainer_state["config"], f, indent=2, default=str)
1093
+ return directory
1094
+
1095
+ def load(self, directory: Path) -> None:
1096
+ """Restore policy weights, optimizer state, and RNGs from checkpoint."""
1097
+ model_path = directory / "model.pt"
1098
+ state_path = directory / "trainer_state.pt"
1099
+ if not model_path.exists() or not state_path.exists():
1100
+ raise FileNotFoundError(f"Checkpoint directory missing model.pt or trainer_state.pt: {directory}")
1101
+
1102
+ self.policy.load_state_dict(torch.load(model_path, map_location=self.device, weights_only=True))
1103
+ state = torch.load(state_path, map_location="cpu", weights_only=False)
1104
+ self.optimizer.load_state_dict(state["optimizer"])
1105
+ self.global_steps = int(state.get("global_steps", 0))
1106
+ self.episodes = int(state.get("episodes", 0))
1107
+ self.updates = int(state.get("updates", 0))
1108
+ self.last_log_step = int(state.get("last_log_step", 0))
1109
+ self.last_save_step = int(state.get("last_save_step", 0))
1110
+ self.last_video_step = int(state.get("last_video_step", 0))
1111
+ self.object_counts.update(state.get("object_counts", {}))
1112
+
1113
+ try:
1114
+ torch.set_rng_state(state["torch_rng_state"])
1115
+ np.random.set_state(state["numpy_rng_state"])
1116
+ random.setstate(state["python_rng_state"])
1117
+ if torch.cuda.is_available() and "cuda_rng_state" in state:
1118
+ torch.cuda.set_rng_state_all(state["cuda_rng_state"])
1119
+ except Exception as exc:
1120
+ print(f"[resume] Notice: Could not restore full RNG state: {exc}")
1121
+
1122
+ print(f"[resume] Loaded checkpoint {directory} | global_steps={self.global_steps:,} episodes={self.episodes:,}")
1123
+
1124
+ def update(self, batch: Dict[str, np.ndarray], total_steps_target: int = NUM_EPISODES * 100) -> Dict[str, float]:
1125
+ """Perform PPO mini-batch updates over the collected rollout."""
1126
+ obs = torch.as_tensor(batch["obs"], dtype=torch.float32, device=self.device)
1127
+ actions = torch.as_tensor(batch["raw_actions"], dtype=torch.float32, device=self.device)
1128
+ old_logprobs = torch.as_tensor(batch["logprobs"], dtype=torch.float32, device=self.device)
1129
+ advantages = torch.as_tensor(batch["advantages"], dtype=torch.float32, device=self.device)
1130
+ returns = torch.as_tensor(batch["returns"], dtype=torch.float32, device=self.device)
1131
+ old_values = torch.as_tensor(batch["values"], dtype=torch.float32, device=self.device)
1132
+
1133
+ # Standard advantage normalization
1134
+ adv_std = advantages.std()
1135
+ if adv_std > 1e-8:
1136
+ advantages = (advantages - advantages.mean()) / (adv_std + 1e-8)
1137
+
1138
+ # Optional learning rate schedule annealing
1139
+ if ANNEAL_LR:
1140
+ frac = 1.0 - (self.global_steps / max(1, total_steps_target))
1141
+ lr_now = max(1e-6, frac * LEARNING_RATE)
1142
+ for param_group in self.optimizer.param_groups:
1143
+ param_group["lr"] = lr_now
1144
+
1145
+ n = obs.shape[0]
1146
+ minibatch = min(MINIBATCH_SIZE, n)
1147
+ indices = np.arange(n)
1148
+ metrics: Dict[str, List[float]] = {
1149
+ "policy_loss": [],
1150
+ "value_loss": [],
1151
+ "entropy": [],
1152
+ "approx_kl": [],
1153
+ "clipfrac": [],
1154
+ "explained_var": [],
1155
+ }
1156
+
1157
+ for _ in range(PPO_EPOCHS):
1158
+ np.random.shuffle(indices)
1159
+ for start in range(0, n, minibatch):
1160
+ mb = indices[start:start + minibatch]
1161
+ _, new_logprob, entropy, new_value, _ = self.policy.get_action_and_value(obs[mb], actions[mb])
1162
+
1163
+ logratio = new_logprob - old_logprobs[mb]
1164
+ ratio = logratio.exp()
1165
+ mb_adv = advantages[mb]
1166
+
1167
+ # Policy loss (clipped surrogate)
1168
+ pg_loss1 = -mb_adv * ratio
1169
+ pg_loss2 = -mb_adv * torch.clamp(ratio, 1.0 - CLIP_COEF, 1.0 + CLIP_COEF)
1170
+ policy_loss = torch.max(pg_loss1, pg_loss2).mean()
1171
+
1172
+ # Value loss (smooth MSE)
1173
+ value_loss = 0.5 * ((new_value - returns[mb]) ** 2).mean()
1174
+ entropy_loss = entropy.mean()
1175
+
1176
+ loss = policy_loss + VALUE_COEF * value_loss - ENTROPY_COEF * entropy_loss
1177
+
1178
+ self.optimizer.zero_grad(set_to_none=True)
1179
+ loss.backward()
1180
+ nn.utils.clip_grad_norm_(self.policy.parameters(), MAX_GRAD_NORM)
1181
+ self.optimizer.step()
1182
+
1183
+ approx_kl = ((ratio - 1.0) - logratio).mean().detach().cpu().item()
1184
+ clipfrac = ((ratio - 1.0).abs() > CLIP_COEF).float().mean().detach().cpu().item()
1185
+
1186
+ metrics["policy_loss"].append(policy_loss.detach().cpu().item())
1187
+ metrics["value_loss"].append(value_loss.detach().cpu().item())
1188
+ metrics["entropy"].append(entropy_loss.detach().cpu().item())
1189
+ metrics["approx_kl"].append(approx_kl)
1190
+ metrics["clipfrac"].append(clipfrac)
1191
+
1192
+ self.updates += 1
1193
+ y_true = returns.cpu().numpy()
1194
+ y_pred = old_values.cpu().numpy()
1195
+ var_y = np.var(y_true)
1196
+ explained_var = float(1.0 - np.var(y_true - y_pred) / (var_y + 1e-8)) if var_y > 1e-8 else 0.0
1197
+
1198
+ out_metrics = {k: float(np.mean(v)) for k, v in metrics.items() if len(v) > 0}
1199
+ out_metrics["explained_var"] = explained_var
1200
+ return out_metrics
1201
+
1202
+
1203
+ # =============================================================================
1204
+ # ROLLOUT COLLECTION & GAE
1205
+ # =============================================================================
1206
+
1207
+ def collect_rollout(
1208
+ policy: ActorCritic,
1209
+ envs: List[PlateBalanceEnv],
1210
+ current_obs: np.ndarray,
1211
+ trainer: PPOTrainer,
1212
+ ) -> Tuple[Dict[str, np.ndarray], np.ndarray, List[EpisodeStats], int]:
1213
+ """Collect vector rollouts and compute generalized advantage estimations."""
1214
+ n_envs = len(envs)
1215
+ T = ROLLOUT_STEPS
1216
+
1217
+ obs_buf = np.zeros((T, n_envs, OBSERVATION_SIZE), dtype=np.float32)
1218
+ actions_buf = np.zeros((T, n_envs, ACTION_SIZE), dtype=np.float32)
1219
+ raw_actions_buf = np.zeros((T, n_envs, ACTION_SIZE), dtype=np.float32)
1220
+ logprob_buf = np.zeros((T, n_envs), dtype=np.float32)
1221
+ rewards_buf = np.zeros((T, n_envs), dtype=np.float32)
1222
+ terminated_buf = np.zeros((T, n_envs), dtype=np.float32)
1223
+ truncated_buf = np.zeros((T, n_envs), dtype=np.float32)
1224
+ values_buf = np.zeros((T, n_envs), dtype=np.float32)
1225
+
1226
+ completed_episodes: List[EpisodeStats] = []
1227
+ start_steps = trainer.global_steps
1228
+
1229
+ policy.eval()
1230
+ for t in range(T):
1231
+ obs_buf[t] = current_obs
1232
+ obs_t = torch.as_tensor(current_obs, dtype=torch.float32, device=trainer.device)
1233
+ with torch.no_grad():
1234
+ action_t, logprob_t, _, value_t, raw_t = policy.get_action_and_value(obs_t)
1235
+
1236
+ actions = action_t.cpu().numpy()
1237
+ raw_actions = raw_t.cpu().numpy()
1238
+
1239
+ actions_buf[t] = actions
1240
+ raw_actions_buf[t] = raw_actions
1241
+ logprob_buf[t] = logprob_t.cpu().numpy()
1242
+ values_buf[t] = value_t.cpu().numpy()
1243
+
1244
+ next_obs = np.empty_like(current_obs)
1245
+ for e, env in enumerate(envs):
1246
+ next_obs[e], reward, terminated, truncated, info = env.step(actions[e])
1247
+ rewards_buf[t, e] = reward
1248
+ terminated_buf[t, e] = float(terminated)
1249
+ truncated_buf[t, e] = float(truncated)
1250
+ trainer.global_steps += 1
1251
+
1252
+ if terminated or truncated:
1253
+ trainer.episodes += 1
1254
+ trainer.object_counts[info["object"]] = trainer.object_counts.get(info["object"], 0) + 1
1255
+ completed_episodes.append(EpisodeStats(
1256
+ reward=float(info["episode_reward"]),
1257
+ length=int(info["episode_length"]),
1258
+ object_name=str(info["object"]),
1259
+ success=bool(not info["fallen"] and info["episode_length"] >= MAX_EPISODE_STEPS),
1260
+ crumbs_on=int(info.get("crumbs_on_plate", 1)),
1261
+ ))
1262
+ # Reset environment immediately upon termination or truncation
1263
+ next_obs[e] = env.reset()
1264
+
1265
+ current_obs = next_obs
1266
+ if trainer.episodes >= NUM_EPISODES:
1267
+ break
1268
+
1269
+ actual_T = t + 1
1270
+ obs_buf = obs_buf[:actual_T]
1271
+ actions_buf = actions_buf[:actual_T]
1272
+ raw_actions_buf = raw_actions_buf[:actual_T]
1273
+ logprob_buf = logprob_buf[:actual_T]
1274
+ rewards_buf = rewards_buf[:actual_T]
1275
+ terminated_buf = terminated_buf[:actual_T]
1276
+ truncated_buf = truncated_buf[:actual_T]
1277
+ values_buf = values_buf[:actual_T]
1278
+
1279
+ with torch.no_grad():
1280
+ next_obs_t = torch.as_tensor(current_obs, dtype=torch.float32, device=trainer.device)
1281
+ next_value = policy.get_value(next_obs_t).cpu().numpy()
1282
+
1283
+ # GAE Computation with proper termination vs. truncation bootstrapping
1284
+ advantages = np.zeros_like(rewards_buf)
1285
+ lastgaelam = np.zeros(n_envs, dtype=np.float32)
1286
+ for t2 in reversed(range(actual_T)):
1287
+ if t2 == actual_T - 1:
1288
+ next_vals = next_value
1289
+ else:
1290
+ next_vals = values_buf[t2 + 1]
1291
+
1292
+ # Only true terminations (falling) zero out future value; timeouts bootstrap value
1293
+ nonterminal = 1.0 - terminated_buf[t2]
1294
+ delta = rewards_buf[t2] + GAMMA * next_vals * nonterminal - values_buf[t2]
1295
+ lastgaelam = delta + GAMMA * GAE_LAMBDA * nonterminal * lastgaelam
1296
+ advantages[t2] = lastgaelam
1297
+
1298
+ returns = advantages + values_buf
1299
+
1300
+ batch = {
1301
+ "obs": obs_buf.reshape(-1, OBSERVATION_SIZE),
1302
+ "actions": actions_buf.reshape(-1, ACTION_SIZE),
1303
+ "raw_actions": raw_actions_buf.reshape(-1, ACTION_SIZE),
1304
+ "logprobs": logprob_buf.reshape(-1),
1305
+ "rewards": rewards_buf.reshape(-1),
1306
+ "terminated": terminated_buf.reshape(-1),
1307
+ "values": values_buf.reshape(-1),
1308
+ "advantages": advantages.reshape(-1),
1309
+ "returns": returns.reshape(-1),
1310
+ }
1311
+ policy.train()
1312
+ return batch, current_obs, completed_episodes, start_steps
1313
+
1314
+
1315
+ # =============================================================================
1316
+ # EVALUATION & VIDEO VISUALIZATION
1317
+ # =============================================================================
1318
+
1319
+ def record_video(
1320
+ policy: ActorCritic,
1321
+ out_path: Path,
1322
+ seed: int = 1234,
1323
+ max_steps: int = VIDEO_LENGTH_STEPS,
1324
+ camera: str = "track",
1325
+ specific_object: Optional[str] = None,
1326
+ ) -> None:
1327
+ """Record and save an MP4 demonstration video of the policy."""
1328
+ if imageio is None:
1329
+ print("[video] imageio / imageio-ffmpeg not installed. Skipping video recording.")
1330
+ return
1331
+
1332
+ env = PlateBalanceEnv(seed=seed, render=True)
1333
+ obs = env.reset(specific_object=specific_object)
1334
+ frames: List[np.ndarray] = []
1335
+ policy.eval()
1336
+
1337
+ try:
1338
+ with torch.no_grad():
1339
+ for _ in range(max_steps):
1340
+ frames.append(env.render(camera=camera))
1341
+ obs_t = torch.as_tensor(obs, dtype=torch.float32, device=DEVICE).unsqueeze(0)
1342
+ action, _, _, _, _ = policy.get_action_and_value(obs_t, deterministic=True)
1343
+ obs, _, terminated, truncated, _ = env.step(action[0].cpu().numpy())
1344
+ if terminated or truncated:
1345
+ obs = env.reset()
1346
+
1347
+ out_path.parent.mkdir(parents=True, exist_ok=True)
1348
+ imageio.mimsave(out_path, frames, fps=VIDEO_FPS, codec="libx264", quality=7)
1349
+ print(f"[video] Saved {len(frames)} frames to: {out_path}")
1350
+ finally:
1351
+ env.close()
1352
+ policy.train()
1353
+
1354
+
1355
+ def evaluate_policy(
1356
+ policy: ActorCritic,
1357
+ episodes_per_object: int = 5,
1358
+ seed: int = 777,
1359
+ deterministic: bool = True,
1360
+ eval_objects: Optional[List[str]] = None,
1361
+ ) -> Dict[str, Any]:
1362
+ """Benchmark the policy across all 21 supported object geometries."""
1363
+ results: Dict[str, Dict[str, float]] = {}
1364
+ policy.eval()
1365
+
1366
+ active_eval_list = eval_objects or OBJECT_TYPES
1367
+
1368
+ print("\n" + "=" * 88)
1369
+ print("EVALUATION BENCHMARK ACROSS ALL 21 OBJECT GEOMETRIES (INCL. CRUMBLING COOKIE)")
1370
+ print("=" * 88)
1371
+ print(f"{'Category':<14} | {'Object':<16} | {'Reward Mean':<12} | {'Len Mean':<10} | {'Survival %':<12} | {'Tracking Err':<12}")
1372
+ print("-" * 88)
1373
+
1374
+ for obj_name in active_eval_list:
1375
+ category = "Multi-Body" if obj_name == "cookie" else ("Harder" if obj_name in HARDER_OBJECT_TYPES else "Standard")
1376
+ env = PlateBalanceEnv(seed=seed, render=False, active_objects=[obj_name])
1377
+ rewards: List[float] = []
1378
+ lengths: List[int] = []
1379
+ distances: List[float] = []
1380
+ survived = 0
1381
+
1382
+ for ep in range(episodes_per_object):
1383
+ obs = env.reset(specific_object=obj_name)
1384
+ done = False
1385
+ ep_dist = []
1386
+ while not done:
1387
+ with torch.no_grad():
1388
+ obs_t = torch.as_tensor(obs, dtype=torch.float32, device=DEVICE).unsqueeze(0)
1389
+ action, _, _, _, _ = policy.get_action_and_value(obs_t, deterministic=deterministic)
1390
+ obs, reward, terminated, truncated, info = env.step(action[0].cpu().numpy())
1391
+ ep_dist.append(info["distance"])
1392
+ if terminated or truncated:
1393
+ rewards.append(info["episode_reward"])
1394
+ lengths.append(info["episode_length"])
1395
+ if not info["fallen"]:
1396
+ survived += 1
1397
+ done = True
1398
+
1399
+ distances.append(float(np.mean(ep_dist)))
1400
+
1401
+ env.close()
1402
+ mean_r = float(np.mean(rewards)) if rewards else 0.0
1403
+ mean_l = float(np.mean(lengths)) if lengths else 0.0
1404
+ surv_rate = (survived / episodes_per_object) * 100.0
1405
+ mean_d = float(np.mean(distances)) if distances else 0.0
1406
+
1407
+ results[obj_name] = {
1408
+ "mean_reward": mean_r,
1409
+ "mean_length": mean_l,
1410
+ "survival_rate": surv_rate,
1411
+ "tracking_error": mean_d,
1412
+ }
1413
+ print(f"{category:<14} | {obj_name:<16} | {mean_r:>12.2f} | {mean_l:>10.1f} | {surv_rate:>11.1f}% | {mean_d:>12.4f}m")
1414
+
1415
+ print("=" * 88 + "\n")
1416
+ policy.train()
1417
+ return results
1418
+
1419
+
1420
+ def human_view(policy: ActorCritic, specific_object: Optional[str] = None) -> None:
1421
+ """Interactive real-time 3D MuJoCo viewer."""
1422
+ try:
1423
+ import mujoco.viewer
1424
+ except ImportError:
1425
+ print("[viewer] mujoco.viewer is unavailable on this system.")
1426
+ return
1427
+
1428
+ env = PlateBalanceEnv(seed=SEED + 999, render=False)
1429
+ obs = env.reset(specific_object=specific_object)
1430
+ control_dt = PHYSICS_TIMESTEP * CONTROL_DECIMATION
1431
+
1432
+ print(f"\n[viewer] Launching interactive 3D viewer for: {specific_object or 'Random Objects'}")
1433
+ print("[viewer] Controls: Space to pause/resume, Esc/close window to exit.")
1434
+ with mujoco.viewer.launch_passive(env.model, env.data) as viewer:
1435
+ policy.eval()
1436
+ while viewer.is_running():
1437
+ step_start = time.time()
1438
+ with torch.no_grad():
1439
+ obs_t = torch.as_tensor(obs, dtype=torch.float32, device=DEVICE).unsqueeze(0)
1440
+ action, _, _, _, _ = policy.get_action_and_value(obs_t, deterministic=True)
1441
+ obs, _, terminated, truncated, info = env.step(action[0].cpu().numpy())
1442
+ viewer.sync()
1443
+
1444
+ if terminated or truncated:
1445
+ status = "TIMEOUT" if truncated else "FALL"
1446
+ crumbs_info = f" | Crumbs on plate: {info.get('crumbs_on_plate', 1)}/5" if info["object"] == "cookie" else ""
1447
+ print(f"[viewer] Episode end: {status} | Object: {info['object']} | Length: {info['episode_length']}{crumbs_info}")
1448
+ obs = env.reset(specific_object=specific_object)
1449
+
1450
+ # Precise real-time rate pacing
1451
+ elapsed = time.time() - step_start
1452
+ if elapsed < control_dt:
1453
+ time.sleep(control_dt - elapsed)
1454
+
1455
+ policy.train()
1456
+ env.close()
1457
+
1458
+
1459
+ def export_onnx(policy: ActorCritic, out_path: Path) -> None:
1460
+ """Export the trained policy backbone and actor head to standard ONNX format."""
1461
+ policy.eval()
1462
+ dummy_input = torch.zeros(1, OBSERVATION_SIZE, dtype=torch.float32, device=DEVICE)
1463
+
1464
+ class ExportWrapper(nn.Module):
1465
+ def __init__(self, p: ActorCritic):
1466
+ super().__init__()
1467
+ self.policy = p
1468
+
1469
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
1470
+ mean, _, _ = self.policy(x)
1471
+ return torch.tanh(mean)
1472
+
1473
+ wrapper = ExportWrapper(policy)
1474
+ out_path.parent.mkdir(parents=True, exist_ok=True)
1475
+ torch.onnx.export(
1476
+ wrapper,
1477
+ dummy_input,
1478
+ str(out_path),
1479
+ input_names=["observation"],
1480
+ output_names=["action"],
1481
+ dynamic_axes={"observation": {0: "batch_size"}, "action": {0: "batch_size"}},
1482
+ opset_version=14,
1483
+ )
1484
+ print(f"[export] Successfully exported ONNX model to: {out_path}")
1485
+ policy.train()
1486
+
1487
+
1488
+ # =============================================================================
1489
+ # MAIN ENTRY POINT & TRAINING LOOP
1490
+ # =============================================================================
1491
+
1492
+ def print_config(active_objects: List[str]) -> None:
1493
+ if not PRINT_HYPERPARAMS:
1494
+ return
1495
+ print("=" * 88)
1496
+ print("PlateBalance RL - Advanced MuJoCo + PyTorch PPO System (21 Objects & Cookie Crumble)")
1497
+ print("=" * 88)
1498
+ print(f"Device: {DEVICE} | Master Seed: {SEED}")
1499
+ print(f"Policy: Obs({OBSERVATION_SIZE}) -> {HIDDEN_SIZE} -> {INTERMEDIATE_SIZE} -> {BOTTLENECK_SIZE} -> Act({ACTION_SIZE})")
1500
+ print(f"PPO: Envs={NUM_ENVS} | Rollout={ROLLOUT_STEPS} | Epochs={PPO_EPOCHS} | Batch={MINIBATCH_SIZE} | LR={LEARNING_RATE}")
1501
+ print(f"Target: Episodes={NUM_EPISODES:,} | MaxSteps/Ep={MAX_EPISODE_STEPS:,}")
1502
+ print(f"Active Objects ({len(active_objects)}): {', '.join(active_objects)}")
1503
+ print(f"Domain Randomization: Mass={RANDOMIZE_MASS}, Friction={RANDOMIZE_FRICTION}, Geometry={RANDOMIZE_OBJECT_SIZE}")
1504
+ print(f"Actuator Torque: {MAX_PLATE_TORQUE} N*m | Control Rate: {1.0/(PHYSICS_TIMESTEP*CONTROL_DECIMATION):.0f} Hz")
1505
+ print("=" * 88)
1506
+
1507
+
1508
+ def main() -> None:
1509
+ global SEED, DEVICE, NUM_EPISODES
1510
+ parser = argparse.ArgumentParser(description="PlateBalance RL: MuJoCo PPO Object Balancing")
1511
+ parser.add_argument("--eval", action="store_true", help="Run benchmark evaluation across all 21 object shapes")
1512
+ parser.add_argument("--human-view", action="store_true", help="Launch interactive 3D viewer")
1513
+ parser.add_argument("--record-video", action="store_true", help="Record evaluation video")
1514
+ parser.add_argument("--export-onnx", type=str, default="", help="Export model to ONNX file path")
1515
+ parser.add_argument("--checkpoint", type=str, default="", help="Path to checkpoint folder to load")
1516
+ parser.add_argument("--resume", action="store_true", help="Auto-resume training from latest checkpoint")
1517
+ parser.add_argument("--episodes", type=int, default=NUM_EPISODES, help="Total training episodes")
1518
+ parser.add_argument("--seed", type=int, default=SEED, help="Random seed")
1519
+ parser.add_argument("--device", type=str, default=DEVICE, help="Compute device (cuda or cpu)")
1520
+ parser.add_argument("--object", type=str, default="", help="Specific object to balance for training/eval/viewer")
1521
+ parser.add_argument("--category", type=str, default="all", choices=["all", "standard", "harder", "cookie"],
1522
+ help="Filter object category (all, standard, harder, cookie)")
1523
+ args = parser.parse_args()
1524
+
1525
+ SEED = args.seed
1526
+ DEVICE = args.device
1527
+ NUM_EPISODES = args.episodes
1528
+
1529
+ # Filter active objects
1530
+ if args.object:
1531
+ if args.object not in OBJECT_INFO:
1532
+ raise ValueError(f"Unknown object '{args.object}'. Available: {', '.join(OBJECT_TYPES)}")
1533
+ active_objects = [args.object]
1534
+ elif args.category == "standard":
1535
+ active_objects = STANDARD_OBJECT_TYPES
1536
+ elif args.category == "harder":
1537
+ active_objects = HARDER_OBJECT_TYPES
1538
+ elif args.category == "cookie":
1539
+ active_objects = MULTI_BODY_OBJECT_TYPES
1540
+ else:
1541
+ active_objects = OBJECT_TYPES
1542
+
1543
+ set_seed(SEED)
1544
+ print_config(active_objects)
1545
+
1546
+ root = Path(OUTPUT_DIR) / EXPERIMENT_NAME
1547
+ root.mkdir(parents=True, exist_ok=True)
1548
+
1549
+ trainer = PPOTrainer(device=DEVICE)
1550
+
1551
+ # Determine checkpoint loading
1552
+ ckpt_path: Optional[Path] = None
1553
+ if args.checkpoint:
1554
+ ckpt_path = Path(args.checkpoint)
1555
+ elif args.resume or RESUME:
1556
+ ckpt_path = Path(RESUME_PATH) if RESUME_PATH else latest_checkpoint(root / "checkpoints")
1557
+
1558
+ if ckpt_path is not None and ckpt_path.exists():
1559
+ trainer.load(ckpt_path)
1560
+
1561
+ # Export ONNX mode
1562
+ if args.export_onnx:
1563
+ export_onnx(trainer.policy, Path(args.export_onnx))
1564
+ return
1565
+
1566
+ # Interactive Human Viewer mode
1567
+ if args.human_view or HUMAN_VIEW:
1568
+ human_view(trainer.policy, specific_object=args.object or None)
1569
+ return
1570
+
1571
+ # Benchmark Evaluation mode
1572
+ if args.eval:
1573
+ evaluate_policy(trainer.policy, eval_objects=active_objects)
1574
+ return
1575
+
1576
+ # Record Video mode
1577
+ if args.record_video:
1578
+ vid_path = root / "eval_videos" / "eval_demo.mp4"
1579
+ record_video(trainer.policy, vid_path, seed=SEED, specific_object=args.object or None)
1580
+ return
1581
+
1582
+ # Standard Training Setup
1583
+ envs = [
1584
+ PlateBalanceEnv(SEED + 1000 * i + trainer.global_steps, render=False, active_objects=active_objects)
1585
+ for i in range(NUM_ENVS)
1586
+ ]
1587
+ current_obs = np.stack([env.reset() for env in envs], axis=0)
1588
+
1589
+ start_time = time.time()
1590
+ running_rewards: List[float] = []
1591
+ running_lengths: List[int] = []
1592
+ running_success: List[float] = []
1593
+
1594
+ print(f"\n[train] Starting PPO training loop with {NUM_ENVS} parallel environments...\n")
1595
+
1596
+ try:
1597
+ while trainer.episodes < NUM_EPISODES:
1598
+ batch, current_obs, completed, rollout_start = collect_rollout(
1599
+ trainer.policy, envs, current_obs, trainer
1600
+ )
1601
+
1602
+ if batch["obs"].shape[0] >= 2:
1603
+ metrics = trainer.update(batch, total_steps_target=NUM_EPISODES * 500)
1604
+ else:
1605
+ metrics = {
1606
+ "policy_loss": 0.0, "value_loss": 0.0, "entropy": 0.0,
1607
+ "approx_kl": 0.0, "clipfrac": 0.0, "explained_var": 0.0,
1608
+ }
1609
+
1610
+ running_rewards.extend(ep.reward for ep in completed)
1611
+ running_lengths.extend(ep.length for ep in completed)
1612
+ running_success.extend(1.0 if ep.success else 0.0 for ep in completed)
1613
+
1614
+ # Periodic console logging
1615
+ if trainer.global_steps - trainer.last_log_step >= LOGGING_STEPS:
1616
+ trainer.last_log_step = trainer.global_steps
1617
+ elapsed = max(time.time() - start_time, 1e-9)
1618
+ sps = (trainer.global_steps / elapsed)
1619
+ mean_reward = float(np.mean(running_rewards[-50:])) if running_rewards else 0.0
1620
+ mean_len = float(np.mean(running_lengths[-50:])) if running_lengths else 0.0
1621
+ succ_rate = (float(np.mean(running_success[-50:])) * 100.0) if running_success else 0.0
1622
+
1623
+ print(
1624
+ f"step={trainer.global_steps:>9,} | ep={trainer.episodes:>7,} | "
1625
+ f"sps={sps:>6.0f} | r50={mean_reward:>8.2f} | len50={mean_len:>6.1f} | "
1626
+ f"succ={succ_rate:>5.1f}% | pi={metrics.get('policy_loss', 0.0):+.4f} | "
1627
+ f"vf={metrics.get('value_loss', 0.0):.4f} | kl={metrics.get('approx_kl', 0.0):.5f}"
1628
+ )
1629
+
1630
+ # Periodic checkpoint saving
1631
+ if trainer.global_steps - trainer.last_save_step >= SAVE_STEPS:
1632
+ trainer.last_save_step = trainer.global_steps
1633
+ ckpt = root / "checkpoints" / f"step_{trainer.global_steps:08d}"
1634
+ trainer.save(ckpt, latest=True, extra={
1635
+ "mean_reward_50": float(np.mean(running_rewards[-50:])) if running_rewards else 0.0,
1636
+ "mean_length_50": float(np.mean(running_lengths[-50:])) if running_lengths else 0.0,
1637
+ })
1638
+ print(f"[save] Checkpoint saved: {ckpt}")
1639
+
1640
+ # Periodic visualization video recording
1641
+ if VISUALIZE and trainer.global_steps - trainer.last_video_step >= VIDEO_EVERY_STEPS:
1642
+ trainer.last_video_step = trainer.global_steps
1643
+ ckpt = root / "checkpoints" / f"step_{trainer.global_steps:08d}"
1644
+ ckpt.mkdir(parents=True, exist_ok=True)
1645
+ video_path = ckpt / f"balance_step_{trainer.global_steps:08d}.mp4"
1646
+ print(f"[video] Recording rollout video to: {video_path}")
1647
+ record_video(trainer.policy, video_path, seed=SEED + trainer.global_steps, max_steps=VIDEO_LENGTH_STEPS)
1648
+ trainer.save(ckpt, latest=True, extra={"video": str(video_path.name)})
1649
+
1650
+ except KeyboardInterrupt:
1651
+ print("\n[interrupt] Training interrupted by user. Saving emergency checkpoint...")
1652
+ ckpt = root / "checkpoints" / f"step_{trainer.global_steps:08d}_interrupt"
1653
+ trainer.save(ckpt, latest=True, extra={"interrupted": True})
1654
+ print(f"[interrupt] Saved emergency checkpoint: {ckpt}")
1655
+ finally:
1656
+ for env in envs:
1657
+ env.close()
1658
+
1659
+ final_dir = root / "checkpoints" / f"step_{trainer.global_steps:08d}_final"
1660
+ trainer.save(final_dir, latest=True, extra={"finished": True})
1661
+ print(f"\n[done] Training completed! Total episodes={trainer.episodes:,} steps={trainer.global_steps:,}")
1662
+ print(f"[done] Final model and state saved to: {final_dir}")
1663
+
1664
+
1665
+ if __name__ == "__main__":
1666
+ main()
inference.py ADDED
@@ -0,0 +1,474 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ Test MrBalance directly from Hugging Face Hub.
4
+
5
+ Requirements:
6
+ pip install -U torch transformers safetensors mujoco
7
+
8
+ Files:
9
+ balance_plate_rl.py
10
+ test_mrbalance_hf.py
11
+
12
+ Usage:
13
+ python test_mrbalance_hf.py
14
+ python test_mrbalance_hf.py --object sphere
15
+ python test_mrbalance_hf.py --object egg
16
+ python test_mrbalance_hf.py --object heavy_ball
17
+
18
+ The script:
19
+ 1. Downloads MrBalance from Hugging Face.
20
+ 2. Loads it through AutoModel with trust_remote_code=True.
21
+ 3. Creates the original MuJoCo environment.
22
+ 4. Uses the Hugging Face policy to control the plate.
23
+ 5. Prints episode statistics.
24
+ 6. Optionally launches the MuJoCo viewer.
25
+ """
26
+
27
+ from __future__ import annotations
28
+
29
+ import argparse
30
+ import time
31
+
32
+ import numpy as np
33
+ import torch
34
+ from transformers import AutoModel
35
+
36
+ from balance_plate_rl import (
37
+ DEVICE,
38
+ MAX_EPISODE_STEPS,
39
+ OBJECT_TYPES,
40
+ PlateBalanceEnv,
41
+ )
42
+
43
+
44
+ # ============================================================================
45
+ # CONFIG
46
+ # ============================================================================
47
+
48
+ MODEL_ID = "fromziro/MrBalance"
49
+
50
+ DEFAULT_EPISODES = 5
51
+ DEFAULT_OBJECT = "sphere"
52
+
53
+ DEVICE_TO_USE = torch.device(
54
+ DEVICE
55
+ )
56
+
57
+
58
+ # ============================================================================
59
+ # MODEL LOADING
60
+ # ============================================================================
61
+
62
+ def load_model():
63
+ print("=" * 72)
64
+ print("Loading MrBalance from Hugging Face")
65
+ print("=" * 72)
66
+
67
+ print(f"Model: {MODEL_ID}")
68
+ print(f"Device: {DEVICE_TO_USE}")
69
+ print()
70
+
71
+ model = AutoModel.from_pretrained(
72
+ MODEL_ID,
73
+ trust_remote_code=True,
74
+ )
75
+
76
+ model = model.to(DEVICE_TO_USE)
77
+ model.eval()
78
+
79
+ print("[load] Model loaded successfully.")
80
+ print()
81
+
82
+ # Print basic architecture information.
83
+ print(
84
+ f"[model] Observation size: "
85
+ f"{model.config.observation_size}"
86
+ )
87
+
88
+ print(
89
+ f"[model] Hidden size: "
90
+ f"{model.config.hidden_size}"
91
+ )
92
+
93
+ print(
94
+ f"[model] Intermediate size: "
95
+ f"{model.config.intermediate_size}"
96
+ )
97
+
98
+ print(
99
+ f"[model] Bottleneck size: "
100
+ f"{model.config.bottleneck_size}"
101
+ )
102
+
103
+ print(
104
+ f"[model] Action size: "
105
+ f"{model.config.action_size}"
106
+ )
107
+
108
+ print()
109
+
110
+ return model
111
+
112
+
113
+ # ============================================================================
114
+ # SINGLE EPISODE
115
+ # ============================================================================
116
+
117
+ def run_episode(
118
+ model,
119
+ object_name: str,
120
+ seed: int,
121
+ render: bool = False,
122
+ ):
123
+ env = PlateBalanceEnv(
124
+ seed=seed,
125
+ render=render,
126
+ active_objects=[object_name],
127
+ )
128
+
129
+ obs = env.reset(
130
+ specific_object=object_name
131
+ )
132
+
133
+ total_reward = 0.0
134
+ episode_length = 0
135
+
136
+ if render:
137
+ try:
138
+ import mujoco.viewer
139
+
140
+ viewer_context = (
141
+ mujoco.viewer.launch_passive(
142
+ env.model,
143
+ env.data,
144
+ )
145
+ )
146
+ except Exception as exc:
147
+ env.close()
148
+ raise RuntimeError(
149
+ f"Could not launch MuJoCo viewer: {exc}"
150
+ ) from exc
151
+ else:
152
+ viewer_context = None
153
+
154
+ try:
155
+ if viewer_context is not None:
156
+ with viewer_context as viewer:
157
+
158
+ while viewer.is_running():
159
+
160
+ step_start = time.time()
161
+
162
+ obs_tensor = torch.as_tensor(
163
+ obs,
164
+ dtype=torch.float32,
165
+ device=DEVICE_TO_USE,
166
+ ).unsqueeze(0)
167
+
168
+ with torch.no_grad():
169
+ output = model(
170
+ obs_tensor,
171
+ deterministic=True,
172
+ )
173
+
174
+ action = (
175
+ output.action[0]
176
+ .detach()
177
+ .cpu()
178
+ .numpy()
179
+ .astype(np.float64)
180
+ )
181
+
182
+ obs, reward, terminated, truncated, info = (
183
+ env.step(action)
184
+ )
185
+
186
+ total_reward += reward
187
+ episode_length += 1
188
+
189
+ viewer.sync()
190
+
191
+ if terminated or truncated:
192
+ break
193
+
194
+ # Match roughly the environment's 100 Hz control rate.
195
+ target_dt = 0.01
196
+ elapsed = time.time() - step_start
197
+
198
+ if elapsed < target_dt:
199
+ time.sleep(
200
+ target_dt - elapsed
201
+ )
202
+
203
+ else:
204
+ while True:
205
+
206
+ obs_tensor = torch.as_tensor(
207
+ obs,
208
+ dtype=torch.float32,
209
+ device=DEVICE_TO_USE,
210
+ ).unsqueeze(0)
211
+
212
+ with torch.no_grad():
213
+ output = model(
214
+ obs_tensor,
215
+ deterministic=True,
216
+ )
217
+
218
+ action = (
219
+ output.action[0]
220
+ .detach()
221
+ .cpu()
222
+ .numpy()
223
+ .astype(np.float64)
224
+ )
225
+
226
+ obs, reward, terminated, truncated, info = (
227
+ env.step(action)
228
+ )
229
+
230
+ total_reward += reward
231
+ episode_length += 1
232
+
233
+ if terminated or truncated:
234
+ break
235
+
236
+ finally:
237
+ env.close()
238
+
239
+ success = (
240
+ not info["fallen"]
241
+ and episode_length >= MAX_EPISODE_STEPS
242
+ )
243
+
244
+ return {
245
+ "object": object_name,
246
+ "reward": float(total_reward),
247
+ "length": int(episode_length),
248
+ "success": bool(success),
249
+ "fallen": bool(info["fallen"]),
250
+ "distance": float(info["distance"]),
251
+ "velocity": float(info["object_velocity"]),
252
+ }
253
+
254
+
255
+ # ============================================================================
256
+ # RANDOM OBSERVATION SANITY CHECK
257
+ # ============================================================================
258
+
259
+ def sanity_check(model):
260
+ """
261
+ Verify that the Hugging Face model can actually execute inference
262
+ independently of MuJoCo.
263
+ """
264
+
265
+ print("=" * 72)
266
+ print("Running model sanity check")
267
+ print("=" * 72)
268
+
269
+ x = torch.randn(
270
+ 4,
271
+ 64,
272
+ dtype=torch.float32,
273
+ device=DEVICE_TO_USE,
274
+ )
275
+
276
+ with torch.no_grad():
277
+ output = model(
278
+ x,
279
+ deterministic=True,
280
+ )
281
+
282
+ print(
283
+ "[sanity] action shape:",
284
+ tuple(output.action.shape),
285
+ )
286
+
287
+ print(
288
+ "[sanity] value shape:",
289
+ tuple(output.value.shape),
290
+ )
291
+
292
+ print(
293
+ "[sanity] action range:",
294
+ float(output.action.min()),
295
+ "to",
296
+ float(output.action.max()),
297
+ )
298
+
299
+ print(
300
+ "[sanity] example action:",
301
+ output.action[0].detach().cpu().numpy(),
302
+ )
303
+
304
+ print("[sanity] PASS")
305
+ print()
306
+
307
+
308
+ # ============================================================================
309
+ # MAIN
310
+ # ============================================================================
311
+
312
+ def main():
313
+ parser = argparse.ArgumentParser(
314
+ description="Test MrBalance from Hugging Face."
315
+ )
316
+
317
+ parser.add_argument(
318
+ "--object",
319
+ type=str,
320
+ default=DEFAULT_OBJECT,
321
+ choices=OBJECT_TYPES,
322
+ help="Object to balance.",
323
+ )
324
+
325
+ parser.add_argument(
326
+ "--episodes",
327
+ type=int,
328
+ default=DEFAULT_EPISODES,
329
+ help="Number of evaluation episodes.",
330
+ )
331
+
332
+ parser.add_argument(
333
+ "--render",
334
+ action="store_true",
335
+ help="Launch interactive MuJoCo viewer.",
336
+ )
337
+
338
+ parser.add_argument(
339
+ "--seed",
340
+ type=int,
341
+ default=12345,
342
+ help="Evaluation seed.",
343
+ )
344
+
345
+ args = parser.parse_args()
346
+
347
+ model = load_model()
348
+
349
+ sanity_check(model)
350
+
351
+ print("=" * 72)
352
+ print(
353
+ f"Testing object: {args.object}"
354
+ )
355
+ print(
356
+ f"Episodes: {args.episodes}"
357
+ )
358
+ print("=" * 72)
359
+ print()
360
+
361
+ results = []
362
+
363
+ for episode in range(args.episodes):
364
+
365
+ print(
366
+ f"[episode {episode + 1}/{args.episodes}] "
367
+ f"Running..."
368
+ )
369
+
370
+ result = run_episode(
371
+ model=model,
372
+ object_name=args.object,
373
+ seed=args.seed + episode,
374
+ render=args.render,
375
+ )
376
+
377
+ results.append(result)
378
+
379
+ status = (
380
+ "SUCCESS"
381
+ if result["success"]
382
+ else (
383
+ "FALL"
384
+ if result["fallen"]
385
+ else "TIMEOUT"
386
+ )
387
+ )
388
+
389
+ print(
390
+ f" status: {status}"
391
+ )
392
+
393
+ print(
394
+ f" reward: {result['reward']:.2f}"
395
+ )
396
+
397
+ print(
398
+ f" length: {result['length']}"
399
+ )
400
+
401
+ print(
402
+ f" distance: {result['distance']:.4f} m"
403
+ )
404
+
405
+ print(
406
+ f" velocity: {result['velocity']:.4f} m/s"
407
+ )
408
+
409
+ print()
410
+
411
+ # ------------------------------------------------------------------------
412
+ # Summary
413
+ # ------------------------------------------------------------------------
414
+
415
+ rewards = [
416
+ r["reward"]
417
+ for r in results
418
+ ]
419
+
420
+ lengths = [
421
+ r["length"]
422
+ for r in results
423
+ ]
424
+
425
+ distances = [
426
+ r["distance"]
427
+ for r in results
428
+ ]
429
+
430
+ successes = sum(
431
+ r["success"]
432
+ for r in results
433
+ )
434
+
435
+ print("=" * 72)
436
+ print("RESULTS")
437
+ print("=" * 72)
438
+
439
+ print(
440
+ f"Object: {args.object}"
441
+ )
442
+
443
+ print(
444
+ f"Reward mean: {np.mean(rewards):.2f}"
445
+ )
446
+
447
+ print(
448
+ f"Reward std: {np.std(rewards):.2f}"
449
+ )
450
+
451
+ print(
452
+ f"Length mean: {np.mean(lengths):.1f}"
453
+ )
454
+
455
+ print(
456
+ f"Survival rate: "
457
+ f"{100.0 * sum(not r['fallen'] for r in results) / len(results):.1f}%"
458
+ )
459
+
460
+ print(
461
+ f"Success rate: "
462
+ f"{100.0 * successes / len(results):.1f}%"
463
+ )
464
+
465
+ print(
466
+ f"Tracking error: "
467
+ f"{np.mean(distances):.4f} m"
468
+ )
469
+
470
+ print("=" * 72)
471
+
472
+
473
+ if __name__ == "__main__":
474
+ main()