"""Pure-Python reward tuning schema shared by the service and trainer.""" from __future__ import annotations import math from collections.abc import Mapping from copy import deepcopy from dataclasses import dataclass from typing import Any MAX_PROPOSAL_CHANGES = 4 MIN_CHANGE_RATIO = 0.5 MAX_CHANGE_RATIO = 2.0 class RewardConfigError(ValueError): """A reward configuration or proposal violated the allowlist.""" @dataclass(frozen=True) class NumericSpec: minimum: float maximum: float default: float allow_zero: bool = True WEIGHT_SPECS: dict[str, NumericSpec] = { "track_linear_velocity": NumericSpec(0.5, 3.0, 1.0, False), "track_angular_velocity": NumericSpec(0.25, 2.0, 1.0, False), "body_orientation_l2": NumericSpec(-3.0, -0.1, -1.0, False), "pose": NumericSpec(0.0, 2.5, 1.0), "body_ang_vel": NumericSpec(-0.2, 0.0, -0.05), "angular_momentum": NumericSpec(-0.1, 0.0, -0.025), "is_terminated": NumericSpec(-400.0, -50.0, -200.0, False), "joint_acc_l2": NumericSpec(-2.0e-6, 0.0, -2.5e-7), "joint_pos_limits": NumericSpec(-30.0, -2.0, -10.0, False), "action_rate_l2": NumericSpec(-0.2, -0.005, -0.05, False), "foot_gait": NumericSpec(0.0, 1.5, 0.5), "foot_clearance": NumericSpec(-3.0, 0.0, -1.0), "foot_slip": NumericSpec(-1.0, 0.0, -0.25), "soft_landing": NumericSpec(-5.0e-3, 0.0, -1.0e-3), "stand_still": NumericSpec(-3.0, 0.0, -1.0), "electrical_power": NumericSpec(-5.0e-3, 0.0, 0.0), } PARAMETER_SPECS: dict[str, NumericSpec] = { "track_linear_velocity.std": NumericSpec(0.25, 1.0, math.sqrt(0.25), False), "track_angular_velocity.std": NumericSpec(0.35, 1.2, math.sqrt(0.5), False), "pose.std_standing_scale": NumericSpec(0.5, 2.0, 1.0, False), "pose.std_walking_scale": NumericSpec(0.5, 2.0, 1.0, False), "pose.std_running_scale": NumericSpec(0.5, 2.0, 1.0, False), "pose.walking_threshold": NumericSpec(0.05, 0.5, 0.1, False), "pose.running_threshold": NumericSpec(1.0, 2.5, 1.5, False), "foot_gait.period": NumericSpec(0.4, 0.8, 0.6, False), "foot_gait.threshold": NumericSpec(0.45, 0.65, 0.56, False), "foot_gait.command_threshold": NumericSpec(0.02, 0.3, 0.1, False), "foot_clearance.target_height": NumericSpec(0.06, 0.16, 0.1, False), "foot_clearance.command_threshold": NumericSpec(0.02, 0.3, 0.1, False), "foot_slip.command_threshold": NumericSpec(0.02, 0.3, 0.1, False), "soft_landing.command_threshold": NumericSpec(0.02, 0.3, 0.1, False), "stand_still.command_threshold": NumericSpec(0.02, 0.3, 0.1, False), } BASE_REWARD_CONFIGURATION: dict[str, dict[str, float]] = { "weights": {name: spec.default for name, spec in WEIGHT_SPECS.items()}, "params": {name: spec.default for name, spec in PARAMETER_SPECS.items()}, } FLAT_TASK = "Unitree-Go2-Flat" OBSTACLE_TASK = "Unitree-Go2-ObstacleAvoidance" OBSTACLE_WEIGHT_SPECS = { "avoidance_weight": NumericSpec(0.5, 5.0, 2.0, False), "collision_penalty": NumericSpec(-10.0, -0.5, -5.0, False), "action_smoothness": NumericSpec(-0.05, -0.001, -0.05, False), } OBSTACLE_PARAMETER_SPECS = {"target_velocity": NumericSpec(0.3, 1.2, 0.6, False)} def task_specs(task_id=FLAT_TASK): if task_id == FLAT_TASK: return WEIGHT_SPECS, PARAMETER_SPECS if task_id == OBSTACLE_TASK: return OBSTACLE_WEIGHT_SPECS, OBSTACLE_PARAMETER_SPECS raise RewardConfigError(f"只支持Flat或Obstacle调参任务:{task_id}") def base_configuration(task_id=FLAT_TASK): weights, params = task_specs(task_id) return { "weights": {k: v.default for k, v in weights.items()}, "params": {k: v.default for k, v in params.items()}, } def _number(name: str, value: Any, spec: NumericSpec) -> float: if isinstance(value, bool) or not isinstance(value, (int, float)): raise RewardConfigError(f"{name} 必须是数值") result = float(value) if not math.isfinite(result): raise RewardConfigError(f"{name} 必须是有限数值") if result == 0.0 and not spec.allow_zero: raise RewardConfigError(f"{name} 不允许关闭") if result < spec.minimum or result > spec.maximum: raise RewardConfigError(f"{name} 必须在 {spec.minimum}–{spec.maximum} 之间") return result def _mapping(value: Any, name: str) -> Mapping[str, Any]: if not isinstance(value, Mapping): raise RewardConfigError(f"{name} 必须是对象") return value def _cross_validate(config: Mapping[str, Mapping[str, float]], task_id=FLAT_TASK) -> None: if task_id == OBSTACLE_TASK: return params = config["params"] if params["pose.walking_threshold"] >= params["pose.running_threshold"]: raise RewardConfigError("pose.walking_threshold 必须小于 pose.running_threshold") def _path_spec(path: str, task_id=FLAT_TASK) -> tuple[str, str, NumericSpec]: weights, params = task_specs(task_id) if path.startswith("weights."): section, name = "weights", path.removeprefix("weights.") spec = weights.get(name) elif path.startswith("params."): section, name = "params", path.removeprefix("params.") spec = params.get(name) else: section, name, spec = "", "", None if spec is None: raise RewardConfigError(f"未知参数约束:{path}") return section, name, spec def validate_constraints(value: Any, task_id=FLAT_TASK) -> dict[str, dict[str, float | str]]: """Validate sparse per-session range/fixed safety constraints.""" root = _mapping(value, "constraints") if len(root) > sum(len(specs) for specs in task_specs(task_id)): raise RewardConfigError("constraints 数量超过白名单参数总数") result: dict[str, dict[str, float | str]] = {} for raw_path, raw_constraint in root.items(): if not isinstance(raw_path, str): raise RewardConfigError("constraint path 必须是字符串") _, _, spec = _path_spec(raw_path, task_id) constraint = _mapping(raw_constraint, raw_path) kind = constraint.get("kind") if kind == "fixed": if set(constraint) != {"kind", "value"}: raise RewardConfigError(f"{raw_path} fixed 约束只能包含 kind/value") fixed = _number(f"{raw_path}.value", constraint["value"], spec) result[raw_path] = {"kind": "fixed", "value": fixed} elif kind == "range": if set(constraint) != {"kind", "min", "max"}: raise RewardConfigError(f"{raw_path} range 约束只能包含 kind/min/max") minimum = _number(f"{raw_path}.min", constraint["min"], spec) maximum = _number(f"{raw_path}.max", constraint["max"], spec) if minimum > maximum: raise RewardConfigError(f"{raw_path} 下限不能大于上限") result[raw_path] = {"kind": "range", "min": minimum, "max": maximum} else: raise RewardConfigError(f"{raw_path}.kind 必须是 range 或 fixed") return result def validate_configuration_constraints(value: Any, constraints: Any, task_id=FLAT_TASK) -> None: """Ensure a complete reward configuration satisfies every session constraint.""" config = validate_configuration(value, task_id) checked = validate_constraints(constraints, task_id) for path, constraint in checked.items(): section, name, _ = _path_spec(path, task_id) current = config[section][name] if constraint["kind"] == "fixed": if current != constraint["value"]: raise RewardConfigError( f"{path} 已固定为 {constraint['value']},不能设为 {current}" ) elif current < constraint["min"] or current > constraint["max"]: raise RewardConfigError( f"{path}={current} 超出工程锁定范围 {constraint['min']}–{constraint['max']}" ) def validate_configuration(value: Any, task_id=FLAT_TASK) -> dict[str, dict[str, float]]: """Validate a complete configuration and reject missing/unknown fields.""" root = _mapping(value, "rewardConfig") if set(root) != {"weights", "params"}: raise RewardConfigError("rewardConfig 只能包含 weights 和 params") weights, params = task_specs(task_id) raw_weights = _mapping(root["weights"], "weights") raw_params = _mapping(root["params"], "params") if set(raw_weights) != set(weights): raise RewardConfigError("weights 必须完整且不能包含未知奖励项") if set(raw_params) != set(params): raise RewardConfigError("params 必须完整且不能包含未知参数") config = { "weights": { name: _number(f"weights.{name}", raw_weights[name], spec) for name, spec in weights.items() }, "params": { name: _number(f"params.{name}", raw_params[name], spec) for name, spec in params.items() }, } _cross_validate(config, task_id) return config def validate_proposal( value: Any, previous: Any, constraints: Any | None = None, task_id=FLAT_TASK ) -> dict[str, dict[str, float]]: """Validate a sparse Agent patch relative to a complete previous config and guardrails.""" weights, params = task_specs(task_id) current = validate_configuration(previous, task_id) root = _mapping(value, "proposal") if not set(root).issubset({"weights", "params"}): raise RewardConfigError("proposal 只能包含 weights 和 params") raw_weights = _mapping(root.get("weights", {}), "weights") raw_params = _mapping(root.get("params", {}), "params") if len(raw_weights) + len(raw_params) == 0: raise RewardConfigError("proposal 至少需要一项修改") if len(raw_weights) + len(raw_params) > MAX_PROPOSAL_CHANGES: raise RewardConfigError(f"proposal 每轮最多修改 {MAX_PROPOSAL_CHANGES} 项") unknown_weights = set(raw_weights) - set(weights) unknown_params = set(raw_params) - set(params) if unknown_weights: raise RewardConfigError(f"未知奖励项:{', '.join(sorted(unknown_weights))}") if unknown_params: raise RewardConfigError(f"未知奖励参数:{', '.join(sorted(unknown_params))}") patch: dict[str, dict[str, float]] = {"weights": {}, "params": {}} for name, raw in raw_weights.items(): value_number = _number(f"weights.{name}", raw, weights[name]) old = current["weights"][name] if old != 0.0 and value_number != 0.0: ratio = abs(value_number / old) if ratio < MIN_CHANGE_RATIO or ratio > MAX_CHANGE_RATIO: raise RewardConfigError( f"weights.{name} 单轮变化必须在旧值幅度的 " f"{MIN_CHANGE_RATIO}×–{MAX_CHANGE_RATIO}×" ) if value_number == old: raise RewardConfigError(f"weights.{name} 没有发生变化") patch["weights"][name] = value_number for name, raw in raw_params.items(): value_number = _number(f"params.{name}", raw, params[name]) old = current["params"][name] ratio = abs(value_number / old) if ratio < MIN_CHANGE_RATIO or ratio > MAX_CHANGE_RATIO: raise RewardConfigError( f"params.{name} 单轮变化必须在旧值的 {MIN_CHANGE_RATIO}×–{MAX_CHANGE_RATIO}×" ) if value_number == old: raise RewardConfigError(f"params.{name} 没有发生变化") patch["params"][name] = value_number candidate = deepcopy(current) candidate["weights"].update(patch["weights"]) candidate["params"].update(patch["params"]) _cross_validate(candidate, task_id) if constraints is not None: validate_configuration_constraints(candidate, constraints, task_id) return patch def merge_proposal( previous: Any, proposal: Any, constraints: Any | None = None, task_id=FLAT_TASK ) -> dict[str, dict[str, float]]: current = validate_configuration(previous, task_id) patch = validate_proposal(proposal, current, constraints, task_id) merged = deepcopy(current) merged["weights"].update(patch["weights"]) merged["params"].update(patch["params"]) return validate_configuration(merged, task_id) def apply_reward_configuration(env_cfg: Any, value: Any, task_id=FLAT_TASK) -> None: """Apply a validated full config to a fresh mjlab environment config.""" config = validate_configuration(value, task_id) if task_id == OBSTACLE_TASK: from mjlab.managers import RewardTermCfg contact = env_cfg.terminations["illegal_contact"] env_cfg.rewards["obstacle_collision"] = RewardTermCfg( func=contact.func, params=deepcopy(contact.params), weight=config["weights"]["collision_penalty"], ) env_cfg.rewards["obstacle_proximity"].weight = -config["weights"]["avoidance_weight"] env_cfg.rewards["action_rate_l2"].weight = config["weights"]["action_smoothness"] env_cfg.commands["twist"].speed = config["params"]["target_velocity"] return for name, weight in config["weights"].items(): if name not in env_cfg.rewards: raise RewardConfigError(f"环境缺少奖励项:{name}") env_cfg.rewards[name].weight = weight params = config["params"] direct = { "track_linear_velocity.std": ("track_linear_velocity", "std"), "track_angular_velocity.std": ("track_angular_velocity", "std"), "pose.walking_threshold": ("pose", "walking_threshold"), "pose.running_threshold": ("pose", "running_threshold"), "foot_gait.period": ("foot_gait", "period"), "foot_gait.threshold": ("foot_gait", "threshold"), "foot_gait.command_threshold": ("foot_gait", "command_threshold"), "foot_clearance.target_height": ("foot_clearance", "target_height"), "foot_clearance.command_threshold": ("foot_clearance", "command_threshold"), "foot_slip.command_threshold": ("foot_slip", "command_threshold"), "soft_landing.command_threshold": ("soft_landing", "command_threshold"), "stand_still.command_threshold": ("stand_still", "command_threshold"), } for path, (term, parameter) in direct.items(): env_cfg.rewards[term].params[parameter] = params[path] for regime in ("standing", "walking", "running"): key = f"std_{regime}" scale = params[f"pose.{key}_scale"] baseline = env_cfg.rewards["pose"].params[key] env_cfg.rewards["pose"].params[key] = { pattern: float(std) * scale for pattern, std in baseline.items() }