164 lines
8.2 KiB
Python
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))
|