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

337 lines
13 KiB
Python

"""Single-file Go2 legacy47 import. No adjacent files or arbitrary ONNX conversion."""
import hashlib
import io
import json
from dataclasses import asdict
from pathlib import Path
import torch
from pretrained import (
BASE_TERMS,
JOINTS,
NORMALIZATION_POLICY,
PretrainedError,
ValidatedSource,
_require,
make_reference_actor,
validate_actor_state,
validate_semantics,
verify_onnx,
)
TEMPLATE = "go2-legacy47-v1"
SYNTHETIC_COUNT = 1_000_000
UPLOAD_LIMITS = {"pt": 256 * 1024**2, "onnx": 64 * 1024**2}
SEMANTICS = {
"joint_names": JOINTS,
"observation_names": BASE_TERMS,
"command_names": ["twist"],
"action_scale": [0.25],
"joint_stiffness": [20, 20, 40] * 4,
"joint_damping": [1, 1, 2] * 4,
"default_joint_pos": [-0.1, 0.9, -1.8, 0.1, 0.9, -1.8] * 2,
}
def validate_metadata(metadata, *, required=False):
"""Check recognized embedded semantics; absence remains a user template assumption."""
import numpy as np
_require(isinstance(metadata, dict), "模型metadata必须是对象")
checked = []
for key, expected in SEMANTICS.items():
if key not in metadata:
_require(not required, f"ONNX缺少语义metadata: {key}")
continue
actual = metadata[key]
if isinstance(actual, str):
actual = actual.split(",")
if isinstance(expected[0], str):
_require(actual == expected, f"模型metadata冲突: {key}")
else:
try:
actual = np.asarray(actual, dtype=np.float64)
_require(
actual.shape == (len(expected),)
and np.isfinite(actual).all()
and np.allclose(actual, expected, atol=1e-6, rtol=0),
f"模型metadata冲突: {key}",
)
except (TypeError, ValueError) as error:
raise PretrainedError(f"模型metadata无效: {key}") from error
checked.append(key)
for key in ("contract", "templateId"):
if key in metadata:
_require(metadata[key] == TEMPLATE, "模型模板metadata冲突")
checked.append(key)
return checked
def _checkpoint(data):
try:
checkpoint = torch.load(io.BytesIO(data), map_location="cpu", weights_only=True)
except Exception as error:
raise PretrainedError(
"不支持或不安全的.pt;仅支持weights_only Go2 legacy47 actor checkpoint"
) from error
_require(
isinstance(checkpoint, dict) and "actor_state_dict" in checkpoint,
"请选择含actor_state_dict的Go2 legacy47 .pt,不接受ZIP工程或任意模型",
)
state = checkpoint["actor_state_dict"]
if isinstance(state, dict) and isinstance(state.get("mlp.0.weight"), torch.Tensor):
_require(
tuple(state["mlp.0.weight"].shape) == (512, 47),
"当前仅支持Go2 legacy47输入;81/97 checkpoint请使用原trial续训,而不是基础策略上传",
)
validate_actor_state(state)
iteration = checkpoint.get("iter")
_require(
iteration is None or (type(iteration) is int and iteration >= 0), "无效checkpoint iteration"
)
checked = validate_metadata(checkpoint)
for key in ("metadata", "infos"):
if key in checkpoint:
checked.extend(validate_metadata(checkpoint[key]))
nested = checkpoint[key].get("metadata")
if nested is not None:
checked.extend(validate_metadata(nested))
return state, iteration, sorted(set(checked))
def _onnx(data):
import onnx
from onnx import helper, numpy_helper
model = onnx.load_model_from_string(data)
graph = model.graph
_require(
not model.functions
and not model.training_info
and len(model.opset_import) == 1
and model.opset_import[0].domain == ""
and model.opset_import[0].version in (17, 18),
"仅支持标准opset17/18受限MLP ONNX",
)
_require(
not graph.sparse_initializer and not graph.quantization_annotation, "不支持稀疏或量化ONNX"
)
_require(len(graph.input) == len(graph.output) == 1, "ONNX必须单输入单输出")
for value, shape in ((graph.input[0], [1, 47]), (graph.output[0], [1, 12])):
tensor = value.type.tensor_type
_require(
tensor.elem_type == onnx.TensorProto.FLOAT
and [d.dim_value for d in tensor.shape.dim] == shape
and all(not d.dim_param for d in tensor.shape.dim),
"ONNX仅支持float32 [1,47] -> [1,12];其他输入请使用对应训练器",
)
expected_shapes = {"obs_normalizer._mean": (1, 47), "onnx::Div_24": (1, 47)}
for i, shape in zip((0, 2, 4, 6), ((512, 47), (256, 512), (128, 256), (12, 128)), strict=True):
expected_shapes[f"mlp.{i}.weight"] = shape
expected_shapes[f"mlp.{i}.bias"] = (shape[0],)
_require(
len(graph.initializer) == len(expected_shapes)
and {t.name for t in graph.initializer} == set(expected_shapes),
"ONNX initializer不符合受支持MLP",
)
tensors = {}
for tensor in graph.initializer:
_require(
tensor.data_location == onnx.TensorProto.DEFAULT and not tensor.external_data,
"拒绝ONNX external data;必须是单个自包含文件",
)
_require(
tensor.data_type == onnx.TensorProto.FLOAT
and tuple(tensor.dims) == expected_shapes[tensor.name],
"ONNX tensor类型/shape不支持",
)
tensors[tensor.name] = torch.from_numpy(numpy_helper.to_array(tensor).copy())
_require(bool(torch.isfinite(tensors[tensor.name]).all()), "ONNX tensor含非有限值")
# Exact dataflow, not merely op/tensor names: no branch, reorder, alias or extra op.
expected_ops = ["Sub", "Div", "Gemm", "Elu", "Gemm", "Elu", "Gemm", "Elu", "Gemm"]
_require(
[n.op_type for n in graph.node] == expected_ops, "ONNX必须是Sub/Div及4层Gemm+3层ELU精确链路"
)
previous = graph.input[0].name
seen = set(tensors) | {previous}
_require(len(seen) == len(tensors) + 1, "ONNX输入与initializer重名")
layer = 0
for node in graph.node:
_require(node.domain == "" and not node.overload, "拒绝ONNX custom op")
attributes = {a.name: helper.get_attribute_value(a) for a in node.attribute}
_require(len(attributes) == len(node.attribute), "重复ONNX属性")
if node.op_type in ("Sub", "Div"):
inputs = [previous, "obs_normalizer._mean" if node.op_type == "Sub" else "onnx::Div_24"]
_require(not attributes, "不支持normalizer算子属性")
elif node.op_type == "Gemm":
inputs = [previous, f"mlp.{layer}.weight", f"mlp.{layer}.bias"]
layer += 2
_require(
set(attributes) <= {"alpha", "beta", "transA", "transB"}
and attributes.get("alpha", 1.0) == 1.0
and attributes.get("beta", 1.0) == 1.0
and attributes.get("transA", 0) == 0
and attributes.get("transB", 0) == 1,
"不支持Gemm缩放/转置属性",
)
else:
inputs = [previous]
_require(
set(attributes) <= {"alpha"} and attributes.get("alpha", 1.0) == 1.0,
"仅支持ELU alpha=1",
)
_require(
list(node.input) == inputs
and len(node.output) == 1
and node.output[0]
and node.output[0] not in seen,
"ONNX实际连边/输出不符合受支持MLP",
)
previous = node.output[0]
seen.add(previous)
_require(previous == graph.output[0].name, "ONNX输出必须是最后Gemm结果")
metadata = {p.key: p.value for p in model.metadata_props}
_require(len(metadata) == len(model.metadata_props), "重复ONNX metadata")
checked = validate_metadata(metadata, required=True)
onnx.checker.check_model(model, full_check=True)
actor = make_reference_actor()
_require(
actor.obs_normalizer.eps == 0.01
and bool((actor.state_dict()["distribution.std_param"] == 1).all()),
"目标训练默认normalizer/exploration已变化,需要新模板",
)
state = actor.state_dict()
for key in state:
if key in tensors:
state[key] = tensors[key]
std = tensors["onnx::Div_24"] - 0.01
_require(bool((std > 0).all()), "ONNX denominator必须大于模板epsilon=.01")
state["obs_normalizer._std"] = std
state["obs_normalizer._var"] = std.square()
state["obs_normalizer.count"].fill_(SYNTHETIC_COUNT)
validate_actor_state(state)
actor.load_state_dict(state)
identity = verify_onnx(data, actor)
return state, checked, identity
def source_identity(fmt, digest):
return hashlib.sha256(f"{TEMPLATE}:{fmt}:{digest}".encode()).hexdigest()
def import_upload(path, fmt, template, directory):
"""Run only inside the resource-limited validator process."""
_require(template == TEMPLATE, "必须明确确认go2-legacy47-v1模板")
_require(fmt in UPLOAD_LIMITS, "仅支持单个.pt或.onnx,不支持ZIP")
data = Path(path).read_bytes()
_require(0 < len(data) <= UPLOAD_LIMITS[fmt], "上传文件为空或过大")
if fmt == "pt":
state, iteration, checked = _checkpoint(data)
identity = None
else:
state, checked, identity = _onnx(data)
iteration = None
from pretrained import comparison_observations
actor = make_reference_actor().eval()
actor.load_state_dict(state)
with torch.inference_mode():
output = actor.mlp(actor.obs_normalizer(comparison_observations()))
_require(bool(torch.isfinite(output).all()), "actor在随机/物理probe上产生非有限动作")
directory = Path(directory)
actor_path = directory / "actor.pt"
torch.save({"actor_state_dict": state}, actor_path)
digest = hashlib.sha256(data).hexdigest()
manifest = {
"schema_version": 1,
"mode": "pretrained-warm-start",
"contract": TEMPLATE,
"sourceFormat": fmt,
"source_id": source_identity(fmt, digest),
"source_iteration": iteration,
"source_actor_dim": 47,
"source_normalizer_count": state["obs_normalizer.count"].item(),
"normalization": NORMALIZATION_POLICY
if fmt == "pt"
else "synthetic-count/unit-new-features",
"template_confirmation": {
"id": TEMPLATE,
"confirmed_by": "user",
"assumptions": (
"Go2 legacy47 observation physics, ordering, 50Hz and action semantics; "
"not verified source env.yaml"
),
},
"verified_facts": {
"actor_tensor_shapes": "47-512-256-128-12/float32",
"metadata_fields": checked,
"finite_output_probes": 48,
"activation": "graph-verified-ELU" if fmt == "onnx" else "user-template-assumed-ELU",
},
"derived_fields": {}
if fmt == "pt"
else {
"normalizer_count": {"policy": "synthetic", "value": SYNTHETIC_COUNT},
"normalizer_std": "denominator - 0.01",
"normalizer_var": "std squared",
"epsilon": 0.01,
"exploration_std": {"policy": "fresh-target-default", "value": 1.0},
},
"onnx_identity": identity,
"critic": "fresh-target-initialization",
"optimizer": "fresh",
"iteration": 0,
"base_observation_terms": BASE_TERMS,
"joint_names": JOINTS,
"artifacts": {
"upload": {"name": f"upload.{fmt}", "sha256": digest, "bytes": len(data)},
"checkpoint": {
"name": "actor.pt",
"sha256": hashlib.sha256(actor_path.read_bytes()).hexdigest(),
"bytes": actor_path.stat().st_size,
"origin": "service-derived-actor-only",
},
},
}
return manifest
def read_uploaded_source(checkpoint, *, allowed_roots, manifest_path, target_env, target_agent):
"""Load only the service-derived actor and bound provenance, never adjacent sidecars."""
from pretrained_sources import regular_bytes
roots = [Path(p).resolve(strict=True) for p in allowed_roots]
checkpoint, manifest_path = Path(checkpoint), Path(manifest_path)
root = next(
(r for r in roots if checkpoint.is_relative_to(r) and manifest_path.is_relative_to(r)), None
)
_require(root is not None, "上传artifact不在受控根内")
manifest = json.loads(regular_bytes(manifest_path, root, 64 * 1024))
_require(
manifest.get("contract") == TEMPLATE and manifest.get("sourceFormat") in UPLOAD_LIMITS,
"无效上传manifest模板/格式",
)
artifacts = manifest["artifacts"]
_require(
manifest["source_id"]
== source_identity(manifest["sourceFormat"], artifacts["upload"]["sha256"]),
"上传原始SHA身份不匹配",
)
data = regular_bytes(checkpoint, root, UPLOAD_LIMITS["pt"])
_require(
hashlib.sha256(data).hexdigest() == artifacts["checkpoint"]["sha256"], "上传actor SHA不匹配"
)
from pretrained import _plain
from src.tasks.velocity.config.go2.env_cfgs import unitree_go2_flat_env_cfg
from src.tasks.velocity.config.go2.rl_cfg import unitree_go2_ppo_runner_cfg
validate_semantics(
_plain(asdict(unitree_go2_flat_env_cfg())),
_plain(asdict(unitree_go2_ppo_runner_cfg())),
target_env,
target_agent,
)
state, _, _ = _checkpoint(data)
return ValidatedSource(state, manifest)