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

164 lines
8.2 KiB
Python

"""Opt-in real CPU runner integration: four environments, one ONNX-derived PPO iteration."""
import copy
import io
import json
import os
import sys
import tempfile
import unittest
from dataclasses import asdict
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
for root in (ROOT, ROOT / "rl"):
sys.path.insert(0, str(root))
@unittest.skipUnless(os.environ.get("GO2_UPLOAD_RUNNER_OUTPUT"), "opt-in four-env CPU runner smoke")
class UploadedRunnerTest(unittest.TestCase):
def test_real_upload_runner_ppo_export_and_same_trial_restore(self):
import numpy as np
import onnxruntime as ort
import torch
import warp as wp
if not hasattr(wp, "context"):
from warp._src import context
wp.context = context
from mjlab.envs import ManagerBasedRlEnv
from mjlab.rl import RslRlVecEnvWrapper
from pretrained import initialize_runner, make_reference_actor, validate_runtime_contract
from pretrained_sources import PretrainedSources
from pretrained_upload import read_uploaded_source
from src.tasks.obstacle_avoidance.env_cfg import unitree_go2_obstacle_env_cfg
from src.tasks.velocity.config.go2.rl_cfg import unitree_go2_ppo_runner_cfg
from src.tasks.velocity.rl.runner import VelocityOnPolicyRunner
from task_config import OBSTACLE_TASK, deployment_metadata, validate_task_config
torch.set_num_threads(1)
source_root = Path(os.environ["GO2_UPLOAD_REAL_DIR"])
output = Path(os.environ["GO2_UPLOAD_RUNNER_OUTPUT"])
output.mkdir(parents=True, exist_ok=True)
# Only these explicit files; each decoder sees its independent empty root.
original_onnx = (source_root / "policy.onnx").read_bytes()
options = ort.SessionOptions()
options.intra_op_num_threads = options.inter_op_num_threads = 1
original = ort.InferenceSession(original_onnx, options, providers=["CPUExecutionProvider"])
evidence = {}
for fmt, data in (
("pt", (source_root / "model_10000.pt").read_bytes()),
("onnx", original_onnx),
):
with tempfile.TemporaryDirectory() as store:
registry = PretrainedSources(None, Path(store), sys.executable, ROOT / "rl")
record = registry.receive_upload(
io.BytesIO(data), len(data), fmt, "go2-legacy47-v1", f"single.{fmt}"
)
bound = record["initialization"]
directory = registry.verify(bound)
cfg = unitree_go2_obstacle_env_cfg()
cfg.scene.num_envs = 4
cfg.seed = 42
agent = unitree_go2_ppo_runner_cfg()
agent.logger = "tensorboard"
agent.max_iterations = 1
source = read_uploaded_source(
directory / "actor.pt",
allowed_roots=[directory],
manifest_path=directory / "upload.json",
target_env=asdict(cfg),
target_agent=asdict(agent),
)
raw = ManagerBasedRlEnv(cfg, device="cpu")
env = RslRlVecEnvWrapper(raw)
try:
validate_runtime_contract(raw)
raw.platform_deployment = deployment_metadata(
OBSTACLE_TASK, validate_task_config(OBSTACLE_TASK, {}, 42), 42
)
log = output / fmt
log.mkdir(exist_ok=True)
runner = VelocityOnPolicyRunner(env, asdict(agent), str(log), "cpu")
raw.platform_initialization = initialize_runner(runner, source)
(log / "initialization.json").write_text(
json.dumps(raw.platform_initialization)
)
actor = runner.alg.actor
obs = env.get_observations()
reference = make_reference_actor().eval()
reference.load_state_dict(source.actor_state)
with torch.no_grad():
before = actor(obs)
expected = reference.mlp(reference.obs_normalizer(obs["actor"][:, :47]))
torch.testing.assert_close(before, expected, atol=2e-5, rtol=2e-5)
onnx_actions = np.concatenate(
[
original.run(None, {"obs": row[None].numpy()})[0]
for row in obs["actor"][:, :47]
]
)
np.testing.assert_allclose(before.numpy(), onnx_actions, atol=2e-5, rtol=2e-5)
self.assertEqual(actor.mlp[0].weight[:, 47:].count_nonzero().item(), 0)
self.assertFalse(runner.alg.optimizer.state)
actor(obs).square().mean().backward()
gradient = actor.mlp[0].weight.grad[:, 47:]
self.assertTrue(torch.isfinite(gradient).all())
self.assertGreater(gradient.abs().max().item(), 0)
evidence[fmt] = {
"source_id": record["id"],
"real_observation_ort_error": float(
np.abs(before.numpy() - onnx_actions).max()
),
"new_column_gradient_max": gradient.abs().max().item(),
}
runner.alg.optimizer.zero_grad()
if fmt == "pt":
continue # Only ONNX-derived branch runs the one approved PPO iteration.
runner.learn(num_learning_iterations=1, init_at_random_ep_len=True)
self.assertGreater(actor.mlp[0].weight[:, 47:].abs().max().item(), 0)
saved = torch.load(log / "model_0.pt", map_location="cpu", weights_only=True)
updated = copy.deepcopy(actor.state_dict())
resumed = VelocityOnPolicyRunner(env, asdict(agent), str(log), "cpu")
self.assertTrue(resumed.alg.load(saved, None, strict=True))
self.assertTrue(resumed.alg.optimizer.state)
for key, tensor in updated.items():
self.assertTrue(
torch.equal(tensor, resumed.alg.actor.state_dict()[key]), key
)
self.assertGreater(resumed.alg.actor.obs_normalizer.count.item(), 1_000_000)
exported = ort.InferenceSession(
str(log / "policy.onnx"), options, providers=["CPUExecutionProvider"]
)
metadata = json.loads(
exported.get_modelmeta().custom_metadata_map["pretrained_initialization"]
)
self.assertEqual(metadata["sourceFormat"], "onnx")
self.assertEqual(
metadata["uploadSha256"], source.manifest["artifacts"]["upload"]["sha256"]
)
self.assertNotIn(str(source_root), json.dumps(metadata))
final_obs = env.get_observations()
with torch.no_grad():
expected = resumed.alg.actor(final_obs).numpy()
actual = np.concatenate(
[
exported.run(None, {"obs": row[None].numpy()})[0]
for row in final_obs["actor"]
]
)
np.testing.assert_allclose(actual, expected, atol=2e-5, rtol=2e-5)
evidence[fmt].update(
ppo_iterations=1,
num_envs=4,
count_restored=resumed.alg.actor.obs_normalizer.count.item(),
export_max_error=float(np.abs(actual - expected).max()),
new_columns_max_after_ppo=actor.mlp[0].weight[:, 47:].abs().max().item(),
source_metadata=metadata,
)
finally:
env.close()
(output / "evidence.json").write_text(json.dumps(evidence, indent=2))
print("REAL_RUNNER_UPLOAD_EVIDENCE=" + json.dumps(evidence))