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

476 lines
22 KiB
Python

"""CPU transfer regression; real source and <=4env/1iteration GPU checks are opt-in."""
import copy
import importlib.util
import io
import json
import os
import sys
import tempfile
import unittest
from dataclasses import asdict, replace
from pathlib import Path
from types import SimpleNamespace
ROOT = Path(__file__).resolve().parents[1]
for root in (ROOT, ROOT / "rl"):
sys.path.insert(0, str(root))
HAS_STACK = all(importlib.util.find_spec(m) is not None for m in ("mjlab", "torch", "onnxruntime"))
@unittest.skipUnless(HAS_STACK, "installed RL + CPU ORT stack required")
class PretrainedTest(unittest.TestCase):
@classmethod
def setUpClass(cls):
import torch
torch.set_num_threads(1)
from pretrained import ValidatedSource, make_reference_actor
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
torch.manual_seed(42)
actor = make_reference_actor()
with torch.no_grad():
actor.obs_normalizer._mean.uniform_(-0.2, 0.2)
actor.obs_normalizer._var.uniform_(0.1, 1.2)
actor.obs_normalizer._std.copy_(actor.obs_normalizer._var.sqrt())
actor.obs_normalizer.count.fill_(983138304)
cls.source = ValidatedSource(actor.state_dict(), {"source_id": "synthetic"})
cls.env_cfg = asdict(unitree_go2_flat_env_cfg())
cls.agent_cfg = asdict(unitree_go2_ppo_runner_cfg())
def test_extension_preserves_actions_and_learns_new_columns(self):
import torch
from pretrained import comparison_observations, make_reference_actor, warm_start_actor
from tensordict import TensorDict
source_actor = make_reference_actor().eval()
source_actor.load_state_dict(self.source.actor_state)
base = comparison_observations()
expected = source_actor.mlp(source_actor.obs_normalizer(base)).detach()
for dim in (47, 81, 97):
with self.subTest(dim=dim):
actor = make_reference_actor(dim)
warm_start_actor(actor, self.source)
for key in ("_mean", "_var", "_std"):
value = getattr(actor.obs_normalizer, key)
self.assertTrue(
torch.equal(value[:, :47], self.source.actor_state[f"obs_normalizer.{key}"])
)
self.assertTrue(
torch.equal(
value[:, 47:],
torch.full_like(value[:, 47:], 0 if key == "_mean" else 1),
)
)
self.assertTrue(
torch.equal(
actor.obs_normalizer.count, self.source.actor_state["obs_normalizer.count"]
)
)
self.assertEqual(actor.mlp[0].weight[:, 47:].count_nonzero().item(), 0)
self.assertTrue(actor.mlp[0].weight.requires_grad)
extra = torch.rand(len(base), dim - 47) * 2 - 1
x = torch.cat((base, extra), dim=1)
torch.testing.assert_close(
actor.mlp(actor.obs_normalizer(x)), expected, atol=3e-6, rtol=3e-6
)
initial_error = (actor.mlp(actor.obs_normalizer(x)) - expected).abs().max().item()
x[:, 47:] *= -1
torch.testing.assert_close(
actor.mlp(actor.obs_normalizer(x)), expected, atol=3e-6, rtol=3e-6
)
before = actor.obs_normalizer._mean.clone()
obs = TensorDict({"actor": x}, batch_size=[len(base)])
actor.update_normalization(obs)
self.assertEqual(actor.obs_normalizer.count.item(), 983138304 + len(base))
self.assertLess(
(actor.obs_normalizer._mean[:, :47] - before[:, :47]).abs().max().item(), 1e-6
)
# Updating statistics is not eval identity. Random gravity probes include
# out-of-distribution values in nearly constant source channels.
updated_source = copy.deepcopy(source_actor).train()
updated_source.update_normalization(
TensorDict({"actor": base}, batch_size=[len(base)])
)
updated_expected = updated_source.mlp(updated_source.obs_normalizer(base)).detach()
torch.testing.assert_close(actor(obs), updated_expected, atol=3e-6, rtol=3e-6)
self.assertLess((actor(obs) - expected).abs().max().item(), 2e-4)
if hasattr(self, "evidence"):
self.evidence.setdefault("cpu_transfer", {})[str(dim)] = {
"initial_max_abs_error": initial_error,
"updated_source_max_abs_error": (actor(obs) - updated_expected)
.abs()
.max()
.item(),
"random_batch_action_max_abs_change": (actor(obs) - expected)
.abs()
.max()
.item(),
}
loss = actor(obs).square().mean()
loss.backward()
grad = actor.mlp[0].weight.grad
self.assertTrue(torch.isfinite(grad).all())
if dim > 47:
self.assertGreater(grad[:, 47:].abs().max().item(), 0)
torch.optim.Adam(actor.parameters(), lr=1e-4).step()
if dim > 47:
self.assertGreater(actor.mlp[0].weight[:, 47:].abs().max().item(), 0)
stream = io.BytesIO()
torch.save(actor.state_dict(), stream)
stream.seek(0)
restored = make_reference_actor(dim)
restored.load_state_dict(torch.load(stream, weights_only=True), strict=True)
for k, v in actor.state_dict().items():
self.assertTrue(torch.equal(v, restored.state_dict()[k]), k)
def test_rejects_shapes_nonfinite_and_runtime_activation(self):
import torch
from pretrained import (
PretrainedError,
ValidatedSource,
make_reference_actor,
warm_start_actor,
)
for key, value in (
("mlp.0.weight", torch.zeros(512, 46)),
("obs_normalizer._std", torch.full((1, 47), float("nan"))),
("distribution.std_param", torch.zeros(12)),
):
state = dict(self.source.actor_state)
state[key] = value
with self.assertRaises(PretrainedError):
warm_start_actor(make_reference_actor(81), ValidatedSource(state, {}))
actor = make_reference_actor(81)
actor.mlp[1] = torch.nn.ReLU()
with self.assertRaisesRegex(PretrainedError, "architecture"):
warm_start_actor(actor, self.source)
with self.assertRaises(PretrainedError):
warm_start_actor(make_reference_actor(82), self.source)
def test_semantics_fail_closed(self):
from pretrained import PretrainedError, _plain, validate_semantics
env, agent = _plain(self.env_cfg), _plain(self.agent_cfg)
validate_semantics(env, agent, self.env_cfg, self.agent_cfg)
for edit in (
lambda e: e["observations"]["actor"]["terms"]["phase"]["params"].update(period=0.7),
lambda e: e["observations"]["actor"]["terms"]["joint_vel"].update(scale=0.1),
lambda e: e["actions"]["joint_pos"].update(scale=0.5),
lambda e: e["scene"]["entities"]["robot"]["articulation"]["actuators"][0].update(
armature=0.1
),
):
bad = copy.deepcopy(env)
edit(bad)
with self.assertRaises(PretrainedError):
validate_semantics(bad, agent, self.env_cfg, self.agent_cfg)
with self.assertRaises(PretrainedError):
validate_semantics(env, agent, bad, self.agent_cfg)
bad_agent = copy.deepcopy(agent)
bad_agent["actor"]["activation"] = "relu"
with self.assertRaises(PretrainedError):
validate_semantics(env, bad_agent, self.env_cfg, self.agent_cfg)
def test_compiled_runtime_contract(self):
from unittest.mock import patch
from pretrained import BASE_TERMS, JOINTS, PretrainedError, validate_runtime_contract
metadata = {
"joint_names": JOINTS,
"observation_names": BASE_TERMS + ["forward_depth", "target_error"],
"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,
}
with patch("mjlab.rl.exporter_utils.get_base_metadata", return_value=metadata):
validate_runtime_contract(None)
metadata["joint_names"] = list(reversed(JOINTS))
with self.assertRaisesRegex(PretrainedError, "joint order"):
validate_runtime_contract(None)
def test_safe_yaml_and_allowed_roots(self):
from pretrained import PretrainedError, _read_allowed, _yaml_data
with self.assertRaises(PretrainedError):
_yaml_data(b"x: !!python/object/apply:os.system ['touch /tmp/never-pretrained']")
with self.assertRaises(PretrainedError):
_yaml_data(b"x: 1\nx: 2")
self.assertEqual(_yaml_data(b"axis: {0: 1, 1: 2}"), {"axis": {0: 1, 1: 2}})
self.assertEqual(
_yaml_data(b"x: !!python/name:os.system ''"), {"x": {"symbol": "os.system"}}
)
with tempfile.TemporaryDirectory() as directory, tempfile.TemporaryDirectory() as outside:
root = Path(directory)
secret = Path(outside) / "secret.pt"
secret.write_bytes(b"x")
(root / "link.pt").symlink_to(secret)
for path in (secret, root / "link.pt", root / ".." / Path(outside).name / "secret.pt"):
with self.assertRaises(PretrainedError):
_read_allowed(path, [root], 128)
with self.assertRaises(PretrainedError):
_read_allowed(secret, [Path(outside)], 0)
def test_fresh_runner_only_and_trial_consistency(self):
import torch
from pretrained import PretrainedError, initialize_runner, make_reference_actor
actors = []
for _ in range(2):
actor = make_reference_actor(81)
critic = torch.nn.Linear(108, 1)
original = copy.deepcopy(critic.state_dict())
optimizer = torch.optim.Adam(list(actor.parameters()) + list(critic.parameters()))
runner = SimpleNamespace(
current_learning_iteration=0,
alg=SimpleNamespace(actor=actor, critic=critic, optimizer=optimizer),
)
initialize_runner(runner, self.source)
self.assertFalse(optimizer.state)
self.assertEqual(runner.current_learning_iteration, 0)
for k in original:
self.assertTrue(torch.equal(original[k], critic.state_dict()[k]))
actors.append(actor.state_dict())
runner.current_learning_iteration = 1
with self.assertRaises(PretrainedError):
initialize_runner(runner, self.source)
runner.current_learning_iteration = 0
optimizer.state[actor.mlp[0].weight] = {"step": torch.tensor(1)}
with self.assertRaises(PretrainedError):
initialize_runner(runner, self.source)
for k in actors[0]:
self.assertTrue(torch.equal(actors[0][k], actors[1][k]))
def test_cli_modes_do_not_reinitialize_resume(self):
from scripts.train import TrainConfig, _load_pretrained
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
cfg = TrainConfig(env=unitree_go2_flat_env_cfg(), agent=unitree_go2_ppo_runner_cfg())
self.assertIsNone(_load_pretrained(cfg))
self.assertIsNone(_load_pretrained(replace(cfg, resume_checkpoint="trial.pt")))
with self.assertRaisesRegex(ValueError, "mutually exclusive"):
_load_pretrained(
replace(cfg, resume_checkpoint="trial.pt", pretrained_checkpoint="source.pt")
)
cfg.agent.resume = True
with self.assertRaisesRegex(ValueError, "mutually exclusive"):
_load_pretrained(replace(cfg, pretrained_checkpoint="source.pt"))
with self.assertRaisesRegex(ValueError, "ONNX alone"):
_load_pretrained(replace(cfg, pretrained_onnx="policy.onnx"))
@unittest.skipUnless(
HAS_STACK and os.environ.get("GO2_PRETRAINED_SOURCE"),
"set GO2_PRETRAINED_SOURCE to explicit allowed source directory",
)
class RealPretrainedTest(PretrainedTest):
@classmethod
def setUpClass(cls):
super().setUpClass()
from pretrained import read_pretrained_source
directory = Path(os.environ["GO2_PRETRAINED_SOURCE"])
cls.source = read_pretrained_source(
directory / "model_10000.pt",
allowed_roots=[directory],
target_env=cls.env_cfg,
target_agent=cls.agent_cfg,
)
cls.evidence = {"source": cls.source.manifest}
@classmethod
def tearDownClass(cls):
if os.environ.get("GO2_PRETRAINED_EVIDENCE_DIR"):
directory = Path(os.environ["GO2_PRETRAINED_EVIDENCE_DIR"])
directory.mkdir(parents=True, exist_ok=True)
(directory / "real-source-evidence.json").write_text(
json.dumps(cls.evidence, indent=2) + "\n"
)
def test_wrong_checkpoint_actor_rejected_by_onnx(self):
from pretrained import PretrainedError, make_reference_actor, verify_onnx
# No bulk loading: a fresh random actor is sufficient to prove fail-closed identity.
onnx = (Path(os.environ["GO2_PRETRAINED_SOURCE"]) / "policy.onnx").read_bytes()
with self.assertRaisesRegex(PretrainedError, "does not match"):
verify_onnx(onnx, make_reference_actor())
def test_extended_export_matches_cpu_actor(self):
import numpy as np
import onnxruntime as ort
import torch
from pretrained import comparison_observations, make_reference_actor, warm_start_actor
errors = {}
with tempfile.TemporaryDirectory() as directory:
for dim in (81, 97):
actor = make_reference_actor(dim).eval()
warm_start_actor(actor, self.source)
export = actor.as_onnx(verbose=False)
path = str(Path(directory) / f"policy-{dim}.onnx")
torch.onnx.export(
export,
export.get_dummy_inputs(),
path,
input_names=export.input_names,
output_names=export.output_names,
opset_version=18,
dynamo=False,
)
session = ort.InferenceSession(path, providers=["CPUExecutionProvider"])
x = torch.cat((comparison_observations(), torch.rand(48, dim - 47) * 2 - 1), dim=1)
with torch.no_grad():
expected = export(x).numpy()
actual = np.concatenate(
[session.run(None, {"obs": row[None].numpy()})[0] for row in x]
)
np.testing.assert_allclose(actual, expected, atol=2e-5, rtol=2e-5)
errors[str(dim)] = float(np.abs(actual - expected).max())
self.evidence["extended_onnx_max_abs_error"] = errors
@unittest.skipUnless(
os.environ.get("GO2_PRETRAINED_GPU_SMOKE") == "1", "opt-in <=4env x 1iteration GPU smoke"
)
def test_real_observations_and_one_ppo_iteration(self):
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 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
from tensordict import TensorDict
env_cfg = unitree_go2_obstacle_env_cfg()
env_cfg.scene.num_envs = 4
env_cfg.seed = 42
agent_cfg = unitree_go2_ppo_runner_cfg()
agent_cfg.logger = "tensorboard"
agent_cfg.max_iterations = 1
with tempfile.TemporaryDirectory() as directory:
raw = ManagerBasedRlEnv(env_cfg, device="cuda:0")
self.evidence["compiled_runtime"] = validate_runtime_contract(raw)
raw.platform_deployment = deployment_metadata(
OBSTACLE_TASK, validate_task_config(OBSTACLE_TASK, {}, 42), 42
)
env = RslRlVecEnvWrapper(raw)
try:
runner = VelocityOnPolicyRunner(env, asdict(agent_cfg), directory, "cuda:0")
manifest = initialize_runner(runner, self.source)
actor = runner.alg.actor
cpu_actor = make_reference_actor().eval()
cpu_actor.load_state_dict(self.source.actor_state)
obs = env.get_observations()
base = obs["actor"][:, :47].cpu()
with torch.no_grad():
expected = cpu_actor(TensorDict({"actor": base}, batch_size=[4]))
before = actor(obs).cpu()
torch.testing.assert_close(before, expected, atol=2e-5, rtol=2e-5)
old_mean = actor.obs_normalizer._mean.clone()
actor.update_normalization(obs)
mean_change = (
(actor.obs_normalizer._mean[:, :47] - old_mean[:, :47]).abs().max().item()
)
after = actor(obs).detach().cpu()
self.assertTrue(torch.isfinite(after).all())
torch.testing.assert_close(before, after, atol=2e-5, rtol=2e-5)
loss = actor(obs).square().mean()
loss.backward()
grad = actor.mlp[0].weight.grad[:, 47:]
self.assertTrue(torch.isfinite(grad).all())
self.assertGreater(grad.abs().max().item(), 0)
grad_max = grad.abs().max().item()
runner.alg.optimizer.zero_grad()
critic_before = copy.deepcopy(runner.alg.critic.state_dict())
runner.learn(num_learning_iterations=1, init_at_random_ep_len=True)
self.assertGreater(actor.mlp[0].weight[:, 47:].abs().max().item(), 0)
self.assertTrue(runner.alg.optimizer.state)
self.assertTrue(
any(
not torch.equal(v, runner.alg.critic.state_dict()[k])
for k, v in critic_before.items()
)
)
saved = torch.load(
Path(directory) / "model_0.pt", map_location="cpu", weights_only=True
)
# Same-trial resume uses the installed PPO loader: restore all states.
actor_state = copy.deepcopy(actor.state_dict())
# Production promotion starts a fresh process/runner; do not reload
# into tensors created by the previous rollout's inference_mode.
resumed = VelocityOnPolicyRunner(env, asdict(agent_cfg), directory, "cuda:0")
self.assertTrue(resumed.alg.load(saved, None, strict=True))
resumed.current_learning_iteration = saved["iter"] + 1
self.assertEqual(resumed.current_learning_iteration, 1)
self.assertTrue(resumed.alg.optimizer.state)
actor = resumed.alg.actor
for k, v in actor_state.items():
self.assertTrue(torch.equal(v.cpu(), actor.state_dict()[k].cpu()), k)
import numpy as np
import onnxruntime as ort
session = ort.InferenceSession(
str(Path(directory) / "policy.onnx"), providers=["CPUExecutionProvider"]
)
final_obs = env.get_observations()
with torch.no_grad():
expected_final = actor(final_obs).cpu().numpy()
actual = np.concatenate(
[
session.run(None, {"obs": row[None].cpu().numpy()})[0]
for row in final_obs["actor"]
]
)
np.testing.assert_allclose(actual, expected_final, atol=2e-5, rtol=2e-5)
# Independently run source ORT on observations actually sampled from MuJoCo.
source_session = ort.InferenceSession(
str(Path(os.environ["GO2_PRETRAINED_SOURCE"]) / "policy.onnx"),
providers=["CPUExecutionProvider"],
)
source_actions = np.concatenate(
[source_session.run(None, {"obs": row[None].numpy()})[0] for row in base]
)
np.testing.assert_allclose(source_actions, before.numpy(), atol=2e-5, rtol=2e-5)
self.evidence["gpu_smoke"] = {
"num_envs": 4,
"iterations": 1,
"initialization": manifest,
"real_observation_source_onnx_max_abs_error": float(
np.abs(source_actions - before.numpy()).max()
),
"normalization_batch_action_max_abs_change": (after - before)
.abs()
.max()
.item(),
"normalization_batch_base_mean_max_abs_change": mean_change,
"new_columns_gradient_max": grad_max,
"new_columns_weight_max_after_ppo": actor.mlp[0]
.weight[:, 47:]
.abs()
.max()
.item(),
"export_max_abs_error": float(np.abs(actual - expected_final).max()),
"count_restored": actor.obs_normalizer.count.item(),
"all_actor_tensors_restored_exactly": True,
}
finally:
env.close()
if __name__ == "__main__":
unittest.main()