Files
Mujoco_WASM/training_server/rl/scripts/evaluate_obstacle.py
T
chenlin 438e56bcc8
web-platform-ci / TypeScript, lint, unit, build (push) Has been cancelled
web-platform-ci / Playwright E2E (push) Has been cancelled
feat(training): release V0.9.1 避障训练与基础策略迁移
2026-09-08 10:50:13 +08:00

242 lines
10 KiB
Python

"""Actual checkpoint inference, capturing first-terminal physics before mjlab resets."""
import hashlib
import json
import shutil
import subprocess
import sys
import tempfile
from dataclasses import asdict
from pathlib import Path
from statistics import fmean
from types import SimpleNamespace
# Also executable as a fresh-interpreter seed worker (never fork a CUDA context).
for source in (Path(__file__).resolve().parents[1], Path(__file__).resolve().parents[2]):
if str(source) not in sys.path:
sys.path.insert(0, str(source))
import torch
from mjlab.envs import ManagerBasedRlEnv
from mjlab.rl import MjlabOnPolicyRunner, RslRlVecEnvWrapper
from mjlab.sensor import ContactMatch, ContactSensorCfg
from mjlab.tasks.registry import load_env_cfg, load_rl_cfg, load_runner_cls
from scripts.train import _load_reward_config, _load_task_config
from src.tasks.obstacle_avoidance.env_cfg import apply_obstacle_configuration
from task_config import OBSTACLE_TASK
from tuning.obstacle_scoring import METRICS, STEPS, protocol, score_trajectory
from tuning.schema import apply_reward_configuration
class FirstEpisodeRecorder:
"""The termination hook runs post-physics, before _reset_idx (including first terminal)."""
def __init__(self, env, layout):
self.env = env
self.layout = layout
self.samples = [[] for _ in range(env.num_envs)]
self.previous = torch.zeros((env.num_envs, 12), device=env.device)
self.calls = 0
self.compute = env.termination_manager.compute
env.termination_manager.compute = self.capture
def capture(self):
env = self.env
terminal = self.compute()
robot = env.scene["robot"].data
# Like mjlab termination itself, derived pose is one physics substep old.
position = robot.root_link_pos_w
local_xy = position[:, :2] - env.scene.env_origins[:, :2]
xy = local_xy + torch.tensor(self.layout["spawn"][:2], device=env.device)
boxes = self.layout["boxes"][1:] # Never include the support floor.
clearance = torch.full((env.num_envs,), 0.5, device=env.device)
if boxes:
centers = torch.tensor([b["pos"][:2] for b in boxes], device=env.device)
sizes = torch.tensor([b["size"][:2] for b in boxes], device=env.device)
clearance = (
((xy[:, None, :] - centers).abs() - sizes).clamp(min=0).norm(dim=2).amin(dim=1)
)
clearance = (clearance - 0.3).clamp(min=0) # Conservative body footprint radius.
forces = env.scene["evaluation_obstacles"].data.force_history
collision = forces.norm(dim=-1).flatten(1).amax(dim=1) > 1.0
else:
collision = torch.zeros(env.num_envs, dtype=torch.bool, device=env.device)
nonfoot = env.scene["nonfoot_ground_touch"].data.force_history
collision |= nonfoot.norm(dim=-1).flatten(1).amax(dim=1) > 10.0
fall = (robot.projected_gravity_b[:, 2] > -0.3420201433) | (position[:, 2] < 0.12)
actions = env.action_manager.action
delta = (actions - self.previous).square().mean(dim=1)
self.previous.copy_(actions)
rays = env.scene["forward_scan"].data.distances
hits = (
((rays >= 0) & (rays <= env.scene["forward_scan"].cfg.max_distance)).float().mean(dim=1)
)
distance = env.command_manager.get_term("twist").errors()[1]
values = torch.stack((distance, clearance, delta, hits, collision, fall, terminal), dim=1)
if not torch.isfinite(values).all():
raise ValueError("Nonfinite rollout measurements")
keys = ("distance", "clearance", "action_delta", "ray_hit", "collision", "fall", "terminal")
for samples, row in zip(self.samples, values.cpu().tolist(), strict=True):
if not samples or not samples[-1]["terminal"]:
samples.append(dict(zip(keys, row, strict=True)))
self.calls += 1
return terminal
def metrics(self, horizon=STEPS):
if self.calls != horizon:
raise ValueError("Incomplete rollout horizon")
metrics = [score_trajectory(samples, horizon) for samples in self.samples]
return {key: fmean(item[key] for item in metrics) for key in METRICS}
def configure_seed(task_id, cfg, scenario, reward):
env_cfg = load_env_cfg(task_id, play=False)
env_cfg.seed = scenario["seed"]
env_cfg.scene.num_envs = cfg.num_envs
# Keep the benchmark's declared pair fixed; training itself randomizes every reset.
apply_obstacle_configuration(env_cfg, scenario["taskConfig"], randomize_navigation=False)
apply_reward_configuration(env_cfg, reward, task_id)
env_cfg.curriculum = {}
env_cfg.observations["actor"].enable_corruption = False
env_cfg.events.pop("push_robot", None)
# Primary geom names are literal compiled static terrain geoms, not a regex.
names = tuple(f"terrain_{i}" for i in range(1, len(scenario["terrain"]["boxes"])))
if names:
env_cfg.scene.sensors += (
ContactSensorCfg(
name="evaluation_obstacles",
primary=ContactMatch(mode="geom", pattern=names),
secondary=ContactMatch(mode="subtree", pattern="base_link", entity="robot"),
fields=("force",),
reduce="maxforce",
history_length=env_cfg.decimation,
),
)
return env_cfg
def evaluate_seed(task_id, cfg, scenario, reward):
torch.manual_seed(scenario["seed"])
env_cfg = configure_seed(task_id, cfg, scenario, reward)
agent_cfg = load_rl_cfg(task_id)
env = ManagerBasedRlEnv(env_cfg, device=cfg.device)
wrapped = RslRlVecEnvWrapper(env, clip_actions=agent_cfg.clip_actions)
try:
runner_cls = load_runner_cls(task_id) or MjlabOnPolicyRunner
runner = runner_cls(wrapped, asdict(agent_cfg), log_dir=None, device=wrapped.device)
# rsl_rl actor state_dict includes observation_normalizer buffers. Strict load
# rejects incompatible architecture/statistics; no fixture or zero-action fallback.
runner.load(
str(Path(cfg.checkpoint).resolve(strict=True)),
load_cfg={"actor": True},
strict=True,
map_location=str(wrapped.device),
)
policy = runner.get_inference_policy(device=str(wrapped.device))
obs, _ = env.reset(seed=scenario["seed"])
recorder = FirstEpisodeRecorder(env, scenario["terrain"])
with torch.inference_mode():
for _ in range(STEPS):
actions = policy(obs)
if not torch.isfinite(actions).all():
raise ValueError("Nonfinite policy actions")
obs, _, _, _ = wrapped.step(actions)
return {
"seed": scenario["seed"],
"metrics": recorder.metrics(),
"episodes": cfg.num_envs,
"rolloutSteps": STEPS,
}
finally:
wrapped.close()
def run_obstacle_evaluation(task_id, cfg):
from tuning.obstacle_scoring import SEEDS
if tuple(cfg.seeds) != SEEDS or cfg.steps_per_seed != STEPS:
raise ValueError("Obstacle evaluation seeds/horizon are fixed")
if not cfg.task_config or not cfg.reward_config:
raise ValueError("Obstacle evaluation requires task and parameter configurations")
raw = json.loads(Path(cfg.task_config).read_text())
custom = _load_task_config(task_id, cfg.task_config, raw["seed"])
reward = _load_reward_config(cfg.reward_config, None, OBSTACLE_TASK)
fixed = protocol(custom, cfg.num_envs)
seeds = evaluate_isolated_seeds(task_id, cfg, fixed, reward)
return {
"protocol": fixed,
"seedMetrics": seeds,
"metrics": {key: fmean(item["metrics"][key] for item in seeds) for key in METRICS},
}
def file_hash(path):
digest = hashlib.sha256()
with Path(path).open("rb") as stream:
for chunk in iter(lambda: stream.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def evaluate_isolated_seeds(task_id, cfg, fixed, reward):
from tuning.obstacle_scoring import validate_evaluation
seeds = []
# Inherit the evaluation parent's process group: manager cancellation kills
# both parent and the active seed. subprocess.run kills/reaps on timeout too.
with tempfile.TemporaryDirectory(prefix="go2-obstacle-eval-") as directory:
root = Path(directory)
checkpoint = root / "checkpoint.pt"
source_hash = file_hash(cfg.checkpoint)
shutil.copyfile(cfg.checkpoint, checkpoint)
if file_hash(checkpoint) != source_hash:
raise ValueError("Checkpoint changed during snapshot")
for scenario in fixed["scenarios"]:
request = root / f"request-{scenario['seed']}.json"
output = root / f"result-{scenario['seed']}.json"
request.write_text(
json.dumps(
{
"taskId": task_id,
"config": {**asdict(cfg), "checkpoint": str(checkpoint)},
"scenario": scenario,
"reward": reward,
"checkpointSha256": source_hash,
},
allow_nan=False,
)
)
subprocess.run(
[sys.executable, str(Path(__file__).resolve()), str(request), str(output)],
check=True,
timeout=1200,
shell=False,
)
result = json.loads(output.read_text())
if (
result.pop("checkpointSha256", None) != source_hash
or file_hash(checkpoint) != source_hash
):
raise ValueError("Seed checkpoint identity mismatch")
seeds.append(result)
candidate = {
"protocol": fixed,
"seedMetrics": seeds,
"metrics": {key: fmean(item["metrics"][key] for item in seeds) for key in METRICS},
}
validate_evaluation(candidate, fixed)
return seeds
if __name__ == "__main__":
import src.tasks # noqa: F401
request_path, output_path = map(Path, sys.argv[1:])
request = json.loads(request_path.read_text())
config = SimpleNamespace(**request["config"])
if file_hash(config.checkpoint) != request["checkpointSha256"]:
raise ValueError("Worker checkpoint identity mismatch")
result = evaluate_seed(request["taskId"], config, request["scenario"], request["reward"])
result["checkpointSha256"] = file_hash(config.checkpoint)
output_path.write_text(json.dumps(result, allow_nan=False))