Files
Mujoco_WASM/training_server/mobile_manipulator/export_onnx.py
T
chenlin f3a8a38acd
web-platform-ci / Standalone decision service (no cloud credentials) (push) Has been cancelled
web-platform-ci / TypeScript, lint, unit, build (push) Has been cancelled
web-platform-ci / Playwright E2E (push) Has been cancelled
lekiwi-compatibility / cpu-compatibility (push) Has been cancelled
web-platform-ci / Standalone decision service (no cloud credentials) (pull_request) Has been cancelled
web-platform-ci / TypeScript, lint, unit, build (pull_request) Has been cancelled
web-platform-ci / Playwright E2E (pull_request) Has been cancelled
lekiwi-compatibility / cpu-compatibility (pull_request) Has been cancelled
feat: release v1.0.1 CADWorld 网站与 LeKiwi 智能抓放
集成同源 BYOK 会话隔离、精简模型设置、官方订阅入口和 HTTPS 发布运维;保留本地训练/调参与控制能力。同步 npm 版本及 CHANGELOG,记录公网真实 API 验收仍待用户凭据。
2026-09-24 09:57:41 +08:00

121 lines
4.3 KiB
Python

"""Export a trusted PyTorch actor accepting normalized [1,92], returning [1,12].
CLI input is a TorchScript actor, not an entire PPO checkpoint. Any training-time
VecNormalize must be folded into the actor before export. Never load untrusted .pt.
"""
import argparse
import hashlib
import json
from pathlib import Path
import numpy as np
import torch
from .kernel import STAGES, TASK, validate_config
class BoundedActor(torch.nn.Module):
def __init__(self, actor):
super().__init__()
self.actor = actor
def forward(self, observation):
return self.actor(observation).clamp(-1, 1)
def export_policy(actor, package, output, stage="navigate"):
import onnxruntime as ort
if stage not in STAGES:
raise ValueError("invalid training stage")
package, output = Path(package), Path(output)
config_bytes = (package / "robot.json").read_bytes()
config = json.loads(config_bytes)
validate_config(config)
if json.loads((package / "task.json").read_text()) != TASK:
raise ValueError("task contract mismatch")
# Require the compact browser-exported config for the deployment fingerprint.
# Reformatting robot.json changes its hash; export again rather than guessing.
actor = BoundedActor(actor).cpu().eval()
example = torch.zeros((1, TASK["observationSize"]), dtype=torch.float32)
with torch.no_grad():
result = actor(example)
if (
result.shape != (1, TASK["actionSize"])
or result.dtype != torch.float32
or not torch.isfinite(result).all()
):
raise ValueError("actor must return finite float32 [1,12]")
output.parent.mkdir(parents=True, exist_ok=True)
torch.onnx.export(
actor,
example,
str(output),
input_names=["observation"],
output_names=["action"],
opset_version=17,
dynamo=False,
)
session = ort.InferenceSession(str(output), providers=["CPUExecutionProvider"])
rng = np.random.default_rng(42)
for _ in range(5):
obs = rng.uniform(-1, 1, (1, TASK["observationSize"])).astype(np.float32)
actual = session.run(["action"], {"observation": obs})[0]
with torch.no_grad():
expected = actor(torch.from_numpy(obs)).numpy()
np.testing.assert_allclose(actual, expected, atol=1e-5, rtol=1e-5)
metadata = {
"taskId": TASK["id"],
"actionSemantics": TASK["actionSemantics"],
"trainingStage": stage,
"robotId": config["id"],
"observationSize": TASK["observationSize"],
"actionSize": TASK["actionSize"],
"controlDt": TASK["controlDt"],
"normalized": True,
"modelSha256": hashlib.sha256(output.read_bytes()).hexdigest(),
"robotConfigSha256": hashlib.sha256(config_bytes).hexdigest(),
"sceneSha256": hashlib.sha256(
(package / json.loads((package / "environment.json").read_text())["scene"]).read_bytes()
).hexdigest(),
"input": {"name": "observation", "dtype": "float32", "shape": [1, TASK["observationSize"]]},
"output": {"name": "action", "dtype": "float32", "shape": [1, TASK["actionSize"]]},
}
output.with_suffix(".json").write_text(json.dumps(metadata, indent=2) + "\n")
return metadata
class SmokeActor(torch.nn.Module):
"""Untrained hold policy; only validates transport/shape, NOT task competence."""
def __init__(self, config):
super().__init__()
action = torch.zeros((1, TASK["actionSize"]))
self.register_buffer("action", action)
def forward(self, observation):
return observation[:, : TASK["actionSize"]] * 0 + self.action
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--package", required=True, help="extracted browser training bundle")
choice = parser.add_mutually_exclusive_group(required=True)
choice.add_argument("--actor", help="trusted TorchScript actor.pt")
choice.add_argument(
"--smoke", action="store_true", help="UNTRAINED hold policy for wiring tests"
)
parser.add_argument("--output", required=True)
args = parser.parse_args()
actor = (
SmokeActor(json.loads((Path(args.package) / "robot.json").read_text()))
if args.smoke
else torch.jit.load(args.actor, map_location="cpu")
)
export_policy(actor, args.package, args.output)
if __name__ == "__main__":
main()