feat(tuning): release V0.8.2 Agent 界面重构
This commit is contained in:
@@ -93,6 +93,68 @@ def _cross_validate(config: Mapping[str, Mapping[str, float]]) -> None:
|
||||
raise RewardConfigError("pose.walking_threshold 必须小于 pose.running_threshold")
|
||||
|
||||
|
||||
def _path_spec(path: str) -> tuple[str, str, NumericSpec]:
|
||||
if path.startswith("weights."):
|
||||
section, name = "weights", path.removeprefix("weights.")
|
||||
spec = WEIGHT_SPECS.get(name)
|
||||
elif path.startswith("params."):
|
||||
section, name = "params", path.removeprefix("params.")
|
||||
spec = PARAMETER_SPECS.get(name)
|
||||
else:
|
||||
section, name, spec = "", "", None
|
||||
if spec is None:
|
||||
raise RewardConfigError(f"未知参数约束:{path}")
|
||||
return section, name, spec
|
||||
|
||||
|
||||
def validate_constraints(value: Any) -> dict[str, dict[str, float | str]]:
|
||||
"""Validate sparse per-session range/fixed safety constraints."""
|
||||
root = _mapping(value, "constraints")
|
||||
if len(root) > len(WEIGHT_SPECS) + len(PARAMETER_SPECS):
|
||||
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)
|
||||
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) -> None:
|
||||
"""Ensure a complete reward configuration satisfies every session constraint."""
|
||||
config = validate_configuration(value)
|
||||
checked = validate_constraints(constraints)
|
||||
for path, constraint in checked.items():
|
||||
section, name, _ = _path_spec(path)
|
||||
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) -> dict[str, dict[str, float]]:
|
||||
"""Validate a complete configuration and reject missing/unknown fields."""
|
||||
root = _mapping(value, "rewardConfig")
|
||||
@@ -118,8 +180,10 @@ def validate_configuration(value: Any) -> dict[str, dict[str, float]]:
|
||||
return config
|
||||
|
||||
|
||||
def validate_proposal(value: Any, previous: Any) -> dict[str, dict[str, float]]:
|
||||
"""Validate a sparse Agent patch relative to a complete previous config."""
|
||||
def validate_proposal(
|
||||
value: Any, previous: Any, constraints: Any | None = None
|
||||
) -> dict[str, dict[str, float]]:
|
||||
"""Validate a sparse Agent patch relative to a complete previous config and guardrails."""
|
||||
current = validate_configuration(previous)
|
||||
root = _mapping(value, "proposal")
|
||||
if not set(root).issubset({"weights", "params"}):
|
||||
@@ -167,12 +231,16 @@ def validate_proposal(value: Any, previous: Any) -> dict[str, dict[str, float]]:
|
||||
candidate["weights"].update(patch["weights"])
|
||||
candidate["params"].update(patch["params"])
|
||||
_cross_validate(candidate)
|
||||
if constraints is not None:
|
||||
validate_configuration_constraints(candidate, constraints)
|
||||
return patch
|
||||
|
||||
|
||||
def merge_proposal(previous: Any, proposal: Any) -> dict[str, dict[str, float]]:
|
||||
def merge_proposal(
|
||||
previous: Any, proposal: Any, constraints: Any | None = None
|
||||
) -> dict[str, dict[str, float]]:
|
||||
current = validate_configuration(previous)
|
||||
patch = validate_proposal(proposal, current)
|
||||
patch = validate_proposal(proposal, current, constraints)
|
||||
merged = deepcopy(current)
|
||||
merged["weights"].update(patch["weights"])
|
||||
merged["params"].update(patch["params"])
|
||||
|
||||
Reference in New Issue
Block a user