Files
Mujoco_WASM/training_server/tuning/obstacle_scoring.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

131 lines
5.1 KiB
Python
Raw 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.
"""Fixed obstacle protocol; objective measurements never consume training rewards."""
import math
from statistics import fmean
from task_config import build_terrain_layout, validate_task_config
from .schema import OBSTACLE_TASK
from .scoring import EvaluationError
SEEDS = (101, 202, 303)
STEPS = 1000
WEIGHTS = {"success": 0.4, "time": 0.2, "clearance": 0.2, "smooth": 0.1, "no_fall": 0.1}
METRICS = (*WEIGHTS, "collision_rate", "fall_rate", "arrival_rate", "ray_hit_rate")
def evaluation_scenarios(custom):
scenarios = []
for seed in SEEDS:
value = validate_task_config(OBSTACLE_TASK, custom, seed)
scenarios.append(
{"seed": seed, "taskConfig": value, "terrain": build_terrain_layout(value)}
)
return scenarios
def protocol(custom, num_envs):
return {
"protocolVersion": "obstacle-v1",
"seeds": list(SEEDS),
"stepsPerSeed": STEPS,
"numEnvs": num_envs,
"objectiveWeights": WEIGHTS,
"sceneMode": "fixed-custom-map"
if custom["terrainPreset"] == "custom_boxes"
else "three-seed-layouts",
"scenarios": evaluation_scenarios(custom),
"episodePolicy": "first-episode-only; terminal snapshot before auto-reset; fixed horizon",
}
def score_trajectory(samples, horizon=STEPS):
"""One first-episode trajectory, one sample per policy step including first terminal."""
if not samples or len(samples) > horizon:
raise EvaluationError("轨迹样本数不完整")
keys = {"distance", "clearance", "action_delta", "ray_hit", "collision", "fall", "terminal"}
for sample in samples:
if set(sample) != keys:
raise EvaluationError("轨迹指标缺失/未知")
for key in keys:
value = sample[key]
if not isinstance(value, (int, float)) or not math.isfinite(value) or value < 0:
raise EvaluationError(f"无效轨迹指标 {key}")
if key in {"ray_hit", "collision", "fall", "terminal"} and value > 1:
raise EvaluationError(f"无效标志 {key}")
if any(s["terminal"] for s in samples[:-1]) or (
len(samples) != horizon and not samples[-1]["terminal"]
):
raise EvaluationError("首episode轨迹不完整")
collision = any(s["collision"] for s in samples)
fall = any(s["fall"] for s in samples)
arrival = next((i + 1 for i, s in enumerate(samples) if s["distance"] < 0.5), None)
success = arrival is not None and not collision and not fall
# Missing steps after an early terminal earn zero clearance/smoothness, not a bonus.
return {
"success": float(success),
"time": 1 - arrival / horizon if success else 0.0,
"clearance": 0.0 if fall else sum(min(s["clearance"] / 0.5, 1) for s in samples) / horizon,
"smooth": sum(1 - min(s["action_delta"], 1) for s in samples) / horizon,
"no_fall": float(not fall),
"collision_rate": float(collision),
"fall_rate": float(fall),
"arrival_rate": float(arrival is not None),
"ray_hit_rate": fmean(s["ray_hit"] for s in samples),
}
def validate_metrics(value):
if not isinstance(value, dict) or set(value) != set(METRICS):
raise EvaluationError("避障评估指标缺失/未知")
for key, number in value.items():
if (
isinstance(number, bool)
or not isinstance(number, (int, float))
or not math.isfinite(number)
or not 0 <= number <= 1
):
raise EvaluationError(f"避障指标 {key} 必须在01且有限")
if not math.isclose(value["no_fall"] + value["fall_rate"], 1, abs_tol=1e-8):
raise EvaluationError("跌倒指标不一致")
return value
def validate_evaluation(value, expected_protocol):
if value.get("protocol") != expected_protocol:
raise EvaluationError("避障评估协议/场景不一致")
seeds = value.get("seedMetrics", [])
if len(seeds) != 3:
raise EvaluationError("避障评估必须完整包含3seed")
for item, seed in zip(seeds, SEEDS, strict=True):
if (
item.get("seed") != seed
or item.get("episodes") != expected_protocol["numEnvs"]
or item.get("rolloutSteps") != STEPS
):
raise EvaluationError("避障评估样本数不完整")
validate_metrics(item.get("metrics"))
metrics = validate_metrics(value.get("metrics"))
if any(
not math.isclose(metrics[k], fmean(s["metrics"][k] for s in seeds), abs_tol=1e-8)
for k in METRICS
):
raise EvaluationError("避障评估聚合不一致")
return metrics
def score_evaluation(current, baseline=None):
current = validate_metrics(current)
baseline = validate_metrics(baseline) if baseline is not None else current
eligible = (
current["fall_rate"] <= baseline["fall_rate"] + 0.02 + 1e-12
and current["success"] >= baseline["success"] - 0.02 - 1e-12
)
total = sum(WEIGHTS[k] * current[k] for k in WEIGHTS)
return {
"eligible": eligible,
"score": total if eligible else -1.0,
"components": {k: current[k] for k in WEIGHTS},
"metrics": current,
}