mjlab/scripts/benchmarks/measure_throughput.py
Upstream Snapshot 32a241c28f
Some checks failed
nightly / Test against latest dependencies (py3.10) (push) Has been cancelled
nightly / Test against latest dependencies (py3.13) (push) Has been cancelled
tests / tests (3.13, locked) (push) Has been cancelled
tests / tests (3.13, unlocked) (push) Has been cancelled
tests / pyright (3.10) (push) Has been cancelled
tests / lint-format (push) Has been cancelled
tests / tests (3.10, locked) (push) Has been cancelled
tests / tests (3.11, locked) (push) Has been cancelled
tests / tests (3.12, locked) (push) Has been cancelled
tests / pyright (3.11) (push) Has been cancelled
tests / pyright (3.12) (push) Has been cancelled
tests / pyright (3.13) (push) Has been cancelled
tests / ty-check (3.10) (push) Has been cancelled
tests / ty-check (3.11) (push) Has been cancelled
tests / ty-check (3.12) (push) Has been cancelled
tests / ty-check (3.13) (push) Has been cancelled
tests / stubs (push) Has been cancelled
tests / smoke-test (push) Has been cancelled
Docker / check_paths (push) Has been cancelled
docs / build (push) Has been cancelled
Docker / build (push) Has been cancelled
Import upstream snapshot c19f713c415a699a79d71cd96aa13c3104a05047
Upstream: https://github.com/michaelgillett/mjlab
Upstream-Commit: c19f713c415a699a79d71cd96aa13c3104a05047
Upstream-Branch: main
2026-08-28 15:42:17 +08:00

247 lines
6.6 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""Measure environment throughput for regression tracking.
This script measures physics and environment step throughput across canonical tasks
to catch performance regressions in the manager-based API.
"""
from __future__ import annotations
import json
import subprocess
import time
from dataclasses import asdict, dataclass, field
from datetime import datetime, timezone
from pathlib import Path
import torch
import tyro
import wandb
import mjlab
import mjlab.tasks # noqa: F401 - registers tasks
from mjlab.envs import ManagerBasedRlEnv
from mjlab.tasks.registry import load_env_cfg
from mjlab.tasks.tracking.mdp.commands import MotionCommandCfg
@dataclass
class BenchmarkResult:
"""Results from a single benchmark run."""
task: str
num_envs: int
num_steps: int
decimation: int
physics_sps: float
env_sps: float
overhead_pct: float
def __str__(self) -> str:
return (
f"{self.task} (dec={self.decimation}):\n"
f" Physics SPS: {self.physics_sps:,.0f}\n"
f" Env SPS: {self.env_sps:,.0f}\n"
f" Overhead: {self.overhead_pct:.1f}%"
)
def to_dict(self) -> dict:
return asdict(self)
@dataclass
class ThroughputConfig:
"""Configuration for throughput benchmarking."""
num_envs: int = 4096
"""Number of parallel environments."""
num_steps: int = 200
"""Number of steps to measure (after warmup)."""
warmup_steps: int = 50
"""Number of warmup steps before measuring."""
device: str = "cuda:0"
"""Device to run on."""
tasks: list[str] = field(
default_factory=lambda: [
"Mjlab-Velocity-Flat-Unitree-Go1",
"Mjlab-Tracking-Flat-Unitree-G1",
"Mjlab-Lift-Cube-Yam",
]
)
"""Tasks to benchmark."""
tracking_motion: str = "rll_humanoid/wandb-registry-Motions/lafan_cartwheel:latest"
"""W&B artifact path for tracking task motion (entity/project/name:alias)."""
output_dir: Path | None = None
"""Output directory for JSON results. If None, results are only printed."""
def measure_physics_sps(env: ManagerBasedRlEnv, num_steps: int) -> float:
"""Measure raw physics stepping in env steps per second.
Runs num_steps worth of physics (i.e., num_steps * decimation sim.step calls)
and reports throughput in env steps/sec for direct comparison with env.step().
"""
decimation = env.cfg.decimation
total_physics_steps = num_steps * decimation
torch.cuda.synchronize()
start = time.perf_counter()
for _ in range(total_physics_steps):
env.sim.step()
torch.cuda.synchronize()
elapsed = time.perf_counter() - start
# Report in env steps/sec (not physics steps/sec) for fair comparison.
return (num_steps * env.num_envs) / elapsed
def measure_env_sps(env: ManagerBasedRlEnv, num_steps: int) -> float:
"""Measure full environment step throughput in env steps per second."""
action_dim = sum(env.action_manager.action_term_dim)
action = torch.zeros(env.num_envs, action_dim, device=env.device)
torch.cuda.synchronize()
start = time.perf_counter()
for _ in range(num_steps):
env.step(action)
torch.cuda.synchronize()
elapsed = time.perf_counter() - start
return (num_steps * env.num_envs) / elapsed
def benchmark_task(task: str, cfg: ThroughputConfig) -> BenchmarkResult:
"""Benchmark a single task."""
print(f"\nBenchmarking {task}...")
env_cfg = load_env_cfg(task)
env_cfg.scene.num_envs = cfg.num_envs
# Handle tracking task motion file.
if len(env_cfg.commands) > 0:
motion_cmd = env_cfg.commands.get("motion")
if isinstance(motion_cmd, MotionCommandCfg):
api = wandb.Api()
artifact = api.artifact(cfg.tracking_motion)
motion_dir = artifact.download()
motion_cmd.motion_file = str(Path(motion_dir) / "motion.npz")
env = ManagerBasedRlEnv(cfg=env_cfg, device=cfg.device)
env.reset()
# Warmup.
action_dim = sum(env.action_manager.action_term_dim)
action = torch.zeros(env.num_envs, action_dim, device=env.device)
for _ in range(cfg.warmup_steps):
env.step(action)
torch.cuda.synchronize()
decimation = env.cfg.decimation
physics_sps = measure_physics_sps(env, cfg.num_steps)
env.reset()
torch.cuda.synchronize()
env_sps = measure_env_sps(env, cfg.num_steps)
overhead_pct = 100 * (1 - env_sps / physics_sps)
env.close()
return BenchmarkResult(
task=task,
num_envs=cfg.num_envs,
num_steps=cfg.num_steps,
decimation=decimation,
physics_sps=physics_sps,
env_sps=env_sps,
overhead_pct=overhead_pct,
)
def get_git_commit() -> str:
"""Get current git commit SHA."""
try:
result = subprocess.run(
["git", "rev-parse", "HEAD"],
capture_output=True,
text=True,
check=True,
)
return result.stdout.strip()[:7]
except subprocess.CalledProcessError:
return "unknown"
def save_results(results: list[BenchmarkResult], output_dir: Path) -> None:
"""Save benchmark results to JSON, appending to existing data."""
output_dir.mkdir(parents=True, exist_ok=True)
data_file = output_dir / "throughput_data.json"
# Load existing data.
existing: list[dict] = []
if data_file.exists():
with open(data_file) as f:
existing = json.load(f)
# Create new run entry.
run_entry = {
"created_at": datetime.now(timezone.utc).isoformat(),
"commit": get_git_commit(),
"results": [r.to_dict() for r in results],
}
existing.append(run_entry)
with open(data_file, "w") as f:
json.dump(existing, f, indent=2)
print(f"\nResults saved to {data_file}")
def main(cfg: ThroughputConfig) -> list[BenchmarkResult]:
"""Run throughput benchmarks on all configured tasks."""
print("Throughput Benchmark")
print(f" Envs: {cfg.num_envs}")
print(f" Steps: {cfg.num_steps} (+ {cfg.warmup_steps} warmup)")
print(f" Device: {cfg.device}")
results = []
for task in cfg.tasks:
result = benchmark_task(task, cfg)
results.append(result)
print(result)
print("\n" + "=" * 74)
print("Summary (all values in env steps per second):")
print(" Physics SPS: sim.step() only (×decimation per env step)")
print(" Env SPS: full env.step() including managers")
print(" Overhead: time spent on non-physics work (observations, rewards, etc.)")
print("=" * 74)
print(f"{'Task':<35} {'Dec':>4} {'Physics SPS':>12} {'Env SPS':>12} {'Overhead':>8}")
print("-" * 74)
for r in results:
task_short = r.task.replace("Mjlab-", "").replace("-Unitree-", "-")
print(
f"{task_short:<35} {r.decimation:>4} {r.physics_sps:>12,.0f} {r.env_sps:>12,.0f} {r.overhead_pct:>7.1f}%"
)
if cfg.output_dir:
save_results(results, cfg.output_dir)
return results
if __name__ == "__main__":
cfg = tyro.cli(ThroughputConfig, config=mjlab.TYRO_FLAGS)
main(cfg)