feat(training): release V0.9.1 避障训练与基础策略迁移
web-platform-ci / TypeScript, lint, unit, build (push) Has been cancelled
web-platform-ci / Playwright E2E (push) Has been cancelled

This commit is contained in:
2026-09-08 10:50:13 +08:00
parent fa5485049a
commit 438e56bcc8
113 changed files with 15027 additions and 539 deletions
@@ -0,0 +1,735 @@
"""Single-file upload security, strict graph reconstruction, and shared warm-start seam."""
import copy
import errno
import hashlib
import http.client
import importlib.util
import io
import json
import os
import socket
import sys
import tempfile
import threading
import unittest
from contextlib import contextmanager
from dataclasses import asdict
from http.server import ThreadingHTTPServer
from pathlib import Path
from unittest.mock import patch
ROOT = Path(__file__).resolve().parents[1]
for root in (ROOT, ROOT / "rl"):
sys.path.insert(0, str(root))
from pretrained_sources import PretrainedSources, SourceError # noqa: E402
from server import TrainingManager, TrainingRequestHandler # noqa: E402
def fixture_onnx(state):
import onnx
from onnx import helper, numpy_helper
from pretrained_upload import SEMANTICS
tensors = [
numpy_helper.from_array(value.numpy(), key)
for key, value in state.items()
if key.startswith("mlp.") or key == "obs_normalizer._mean"
]
tensors.append(
numpy_helper.from_array((state["obs_normalizer._std"] + 0.01).numpy(), "onnx::Div_24")
)
nodes = [
helper.make_node("Sub", ["obs", "obs_normalizer._mean"], ["centered"]),
helper.make_node("Div", ["centered", "onnx::Div_24"], ["normalized"]),
]
previous = "normalized"
for i in (0, 2, 4, 6):
output = "actions" if i == 6 else f"linear{i}"
nodes.append(
helper.make_node(
"Gemm", [previous, f"mlp.{i}.weight", f"mlp.{i}.bias"], [output], transB=1
)
)
previous = output
if i < 6:
previous = f"elu{i}"
nodes.append(helper.make_node("Elu", [output], [previous], alpha=1.0))
model = helper.make_model(
helper.make_graph(
nodes,
"actor",
[helper.make_tensor_value_info("obs", onnx.TensorProto.FLOAT, [1, 47])],
[helper.make_tensor_value_info("actions", onnx.TensorProto.FLOAT, [1, 12])],
tensors,
),
opset_imports=[helper.make_opsetid("", 18)],
ir_version=8,
)
helper.set_model_props(
model, {key: ",".join(map(str, value)) for key, value in SEMANTICS.items()}
)
return model
@unittest.skipUnless(
all(importlib.util.find_spec(m) is not None for m in ("torch", "mjlab", "onnx", "onnxruntime")),
"installed CPU RL/ONNX stack required",
)
class UploadTest(unittest.TestCase):
@classmethod
def setUpClass(cls):
import torch
from pretrained import make_reference_actor
torch.set_num_threads(1)
torch.manual_seed(123)
actor = make_reference_actor()
actor.obs_normalizer.count.fill_(983138304)
cls.state = actor.state_dict()
stream = io.BytesIO()
torch.save({"actor_state_dict": cls.state, "iter": 10000}, stream)
cls.pt = stream.getvalue()
cls.onnx = fixture_onnx(cls.state).SerializeToString()
def setUp(self):
self.temp = tempfile.TemporaryDirectory()
self.root = Path(self.temp.name)
self.registry = PretrainedSources(None, self.root / "store", sys.executable, ROOT / "rl")
def tearDown(self):
self.temp.cleanup()
def upload(self, data=None, fmt="pt", name="policy.pt"):
data = self.pt if data is None else data
return self.registry.receive_upload(
io.BytesIO(data), len(data), fmt, "go2-legacy47-v1", name
)
def test_single_file_persistence_dedup_and_binding_cli(self):
record = self.upload(name="../../etc/evil\n.pt")
self.assertEqual(record["label"], "evil_.pt")
bound = record["initialization"]
directory = self.registry.verify(bound)
self.assertEqual(
{p.name for p in directory.iterdir()},
{"upload.pt", "actor.pt", "upload.json", "label.json"},
)
self.assertEqual(self.upload(name="renamed.pt"), record)
restored = PretrainedSources(None, self.root / "store", sys.executable, ROOT / "rl")
self.assertEqual(restored.catalog(), [record])
self.assertEqual(restored.bind(record["id"], "Unitree-Go2-Flat"), bound)
args = restored.arguments(bound)
self.assertIn("--pretrained-upload-manifest", args)
self.assertNotIn("--resume-checkpoint", args)
manager = TrainingManager(
ROOT / "rl",
sys.executable,
("Unitree-Go2-Flat",),
check_environment=False,
sources=restored,
)
cfg = manager.parse_config(
{
"taskId": "Unitree-Go2-Flat",
"numEnvs": 4096,
"maxIterations": 1,
"seed": 42,
"device": "cpu",
"pretrainedSourceId": record["id"],
}
)
self.assertEqual(cfg.pretrained, bound)
self.assertIn("--pretrained-upload-manifest", manager.command_for(cfg))
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
train = TrainConfig(
env=unitree_go2_flat_env_cfg(),
agent=unitree_go2_ppo_runner_cfg(),
pretrained_checkpoint=str(directory / "actor.pt"),
pretrained_upload_manifest=str(directory / "upload.json"),
pretrained_allowed_roots=[str(directory)],
pretrained_source_id=record["id"],
)
source = _load_pretrained(train)
self.assertEqual(source.manifest, bound["manifest"])
from pretrained import public_initialization_metadata
public = public_initialization_metadata(source.manifest)
self.assertEqual(public["uploadSha256"], hashlib.sha256(self.pt).hexdigest())
self.assertEqual(public["templateId"], "go2-legacy47-v1")
self.assertNotIn(str(self.root), json.dumps(public))
# Changing upload content never changes the previously bound job.
newer = self.upload(self.onnx, "onnx")
self.assertNotEqual(newer["id"], record["id"])
self.assertEqual(restored.arguments(cfg.pretrained), args)
(directory / "actor.pt").chmod(0o644)
(directory / "actor.pt").write_bytes(b"tamper")
with self.assertRaisesRegex(SourceError, "SHA"):
restored.arguments(bound)
broken = PretrainedSources(None, self.root / "store", sys.executable, ROOT / "rl")
self.assertFalse(next(r for r in broken.catalog() if r["id"] == record["id"])["ready"])
with self.assertRaises(SourceError):
broken.bind(record["id"], "Unitree-Go2-Flat")
def test_tuning_same_binding_new_trial_and_own_rung_resume(self):
from test_tuning_manager import BASE_METRICS, FakeAdvisor
from tuning.manager import TuningManager
from tuning.process import GpuLease
from tuning.schema import BASE_REWARD_CONFIGURATION
record = self.upload(self.onnx, "onnx")
manager = TuningManager(
ROOT / "rl",
sys.executable,
self.root / "tuning",
GpuLease(),
advisor=FakeAdvisor(),
sources=self.registry,
)
_, config, objective, fallback = manager.parse_create(
{
"mode": "approval",
"pretrainedSourceId": record["id"],
"trialCount": 2,
"numEnvs": 4096,
"initialIterations": 1,
"middleIterations": 2,
"finalIterations": 3,
}
)
session = manager.storage.create_session("approval", config, objective, fallback)
self.assertEqual(config["pretrained"], record["initialization"])
from pretrained import public_initialization_metadata
public = public_initialization_metadata(config["pretrained"]["manifest"])
self.assertEqual(public["sourceFormat"], "onnx")
self.assertIsNone(public["sourceIteration"])
self.assertEqual(public["derivedFields"]["normalizer_count"]["value"], 1_000_000)
self.assertNotIn(
"onnxSha256", public
) # Do not confuse generated checkpoint with user upload.
manager.cancel_events[session["id"]] = threading.Event()
commands = []
def run(_session, command, _cwd, _env, _log):
commands.append(command)
if "--output-dir" in command:
directory = Path(command[command.index("--output-dir") + 1])
(directory / "model_0.pt").write_bytes(b"own checkpoint with updated statistics")
(directory / "policy.onnx").write_bytes(b"trial policy")
(directory / "initialization.json").write_text(
json.dumps(record["initialization"]["manifest"])
)
else:
Path(command[command.index("--output") + 1]).write_text(
json.dumps({"metrics": BASE_METRICS})
)
return 0
with patch.object(manager, "_run_command", run):
parent = None
for number, rung in ((0, 0), (1, 0), (1, 1)):
trial = manager.storage.create_trial(
session["id"],
number,
rung,
rung + 1,
BASE_REWARD_CONFIGURATION,
None,
f"trial-{number}-{rung}",
)
result = manager._execute_trial(session, trial, parent if rung else None)
parent = manager._session_root(session["id"]) / result["checkpointPath"]
train = [c for c in commands if "--output-dir" in c]
for command in train[:2]:
self.assertIn("--pretrained-upload-manifest", command)
self.assertEqual(command[command.index("--pretrained-source-id") + 1], record["id"])
self.assertNotIn("--pretrained-upload-manifest", train[2])
self.assertNotIn("--pretrained-checkpoint", train[2])
self.assertIn("trial-1-0/model_0.pt", train[2][-1])
restored_sources = PretrainedSources(None, self.registry.store, sys.executable, ROOT / "rl")
restored = TuningManager(
ROOT / "rl",
sys.executable,
self.root / "tuning",
GpuLease(),
advisor=FakeAdvisor(),
sources=restored_sources,
)
self.assertEqual(
restored.detail(session["id"])["config"]["pretrained"], record["initialization"]
)
self.assertFalse(restored.workers)
def test_input_limits_disconnect_timeout_capacity_and_cleanup(self):
for length, fmt, template in (
(0, "pt", "go2-legacy47-v1"),
(256 * 1024**2 + 1, "pt", "go2-legacy47-v1"),
(64 * 1024**2 + 1, "onnx", "go2-legacy47-v1"),
(1, "zip", "go2-legacy47-v1"),
(1, "pt", ""),
):
with (
self.subTest(length=length, fmt=fmt, template=template),
self.assertRaises(SourceError),
):
self.registry.receive_upload(io.BytesIO(b"x"), length, fmt, template, "x")
with self.assertRaisesRegex(SourceError, "中断"):
self.registry.receive_upload(io.BytesIO(b"x"), 100, "pt", "go2-legacy47-v1", "x")
with (
patch("pretrained_sources.time.monotonic", side_effect=[0, 61]),
self.assertRaisesRegex(SourceError, "超时"),
):
self.upload()
trickle = unittest.mock.Mock()
trickle.read.side_effect = AssertionError("buffered read could bypass total deadline")
trickle.read1.return_value = b"x"
with (
patch("pretrained_sources.time.monotonic", side_effect=[0, 1, 61]),
self.assertRaisesRegex(SourceError, "超时"),
):
self.registry.receive_upload(trickle, 10, "pt", "go2-legacy47-v1", "x")
self.assertEqual(trickle.read1.call_count, 1)
stream = unittest.mock.Mock()
stream.read1.side_effect = TimeoutError()
with self.assertRaises(SourceError):
self.registry.receive_upload(stream, 10, "pt", "go2-legacy47-v1", "x")
with self.registry.upload_slot(1), self.assertRaisesRegex(SourceError, "正在上传"):
self.upload()
self.assertEqual(list(self.registry.store.iterdir()), [])
import subprocess
with (
patch(
"pretrained_sources.subprocess.run",
side_effect=subprocess.TimeoutExpired("validator", 60),
),
self.assertRaisesRegex(SourceError, "超时"),
):
self.upload()
self.assertEqual(list(self.registry.store.iterdir()), [])
huge = self.registry.store / "quota"
with huge.open("wb") as stream:
stream.truncate(2 * 1024**3)
with self.assertRaisesRegex(SourceError, "上限"):
self.upload()
huge.unlink()
for n in range(32):
(self.registry.store / f"{n:064x}").mkdir()
with self.assertRaisesRegex(SourceError, "上限"):
self.upload()
def test_unsafe_pickle_invalid_checkpoint_and_zip_never_publish(self):
import torch
marker = self.root / "unsafe-executed"
class Unsafe:
def __reduce__(self):
return (os.system, (f"touch {marker}",))
bad = dict(self.state)
bad["mlp.0.weight"] = torch.full((512, 47), float("nan"))
wrong = dict(self.state)
wrong["mlp.0.weight"] = torch.zeros(512, 97)
cases = [b"PK\x03\x04bad zip", b"not a model"]
for payload in (
{"actor_state_dict": bad},
{"actor_state_dict": wrong},
{"actor_state_dict": self.state, "metadata": {"joint_names": ["wrong"]}},
{"actor_state_dict": self.state, "evil": Unsafe()},
):
stream = io.BytesIO()
torch.save(payload, stream)
cases.append(stream.getvalue())
for data in cases:
with self.subTest(size=len(data)), self.assertRaises(SourceError):
self.upload(data)
self.assertEqual(self.registry.catalog(), [])
self.assertEqual(list(self.registry.store.iterdir()), [])
self.assertFalse((self.root / "unsafe-executed").exists())
def test_strict_graph_rejects_bad_edges_ops_attributes_external_and_metadata(self):
import onnx
from pretrained import PretrainedError
from pretrained_upload import _onnx
base = onnx.load_model_from_string(self.onnx)
def external(m):
m.graph.initializer[0].data_location = onnx.TensorProto.EXTERNAL
m.graph.initializer[0].external_data.add(key="location", value="/etc/passwd")
def attrs(m):
m.graph.node[2].attribute.add(name="transA", type=onnx.AttributeProto.INT, i=1)
def metadata(m):
next(p for p in m.metadata_props if p.key == "joint_names").value = "wrong"
def extra(m):
m.graph.node.append(onnx.helper.make_node("Identity", ["actions"], ["branch"]))
def shape(m):
m.graph.input[0].type.tensor_type.shape.dim[1].dim_value = 97
edits = [
external,
attrs,
metadata,
extra,
shape,
lambda m: setattr(m.graph.node[2], "domain", "evil"),
lambda m: setattr(m.graph.node[3], "op_type", "Relu"),
lambda m: m.graph.node[4].input.__setitem__(0, "normalized"),
lambda m: m.graph.node[0].input.reverse(),
lambda m: m.graph.node[4].input.__setitem__(1, "mlp.0.weight"),
lambda m: setattr(m.graph.output[0], "name", "linear0"),
lambda m: setattr(m.graph.initializer[0], "data_type", onnx.TensorProto.DOUBLE),
]
for edit in edits:
model = copy.deepcopy(base)
edit(model)
with self.subTest(edit=edit), self.assertRaises((PretrainedError, ValueError)):
_onnx(model.SerializeToString())
bad = copy.deepcopy(base)
external(bad)
with self.assertRaisesRegex(SourceError, "external"):
self.upload(bad.SerializeToString(), "onnx")
self.assertEqual(self.registry.catalog(), [])
def test_http_upload_storage_errors_are_503_but_read_failures_are_400(self):
manager = TrainingManager(
ROOT / "rl",
sys.executable,
("Unitree-Go2-Flat",),
check_environment=False,
sources=self.registry,
)
reset_timeouts = []
class Handler(TrainingRequestHandler):
access_token = "upload-test-token"
read_error = None
def log_message(self, *_args):
pass
def _upload(self):
original_stream = self.rfile
if self.read_error is not None:
self.rfile = unittest.mock.Mock(wraps=original_stream)
self.rfile.read1.side_effect = self.read_error
try:
super()._upload()
finally:
self.rfile = original_stream
reset_timeouts.append(self.connection.gettimeout())
Handler.manager = manager
server = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
original_open = Path.open
url = (
"/api/training/pretrained-sources/upload"
"?format=pt&template=go2-legacy47-v1&name=policy.pt"
)
try:
for phase, error in (
("open", OSError(errno.ENOSPC, "disk full", "/private/upload.pt")),
("open", OSError(errno.EACCES, "permission denied", "/private/upload.pt")),
("write", OSError(errno.ENOSPC, "disk full", "/private/upload.pt")),
("close", OSError(errno.ENOSPC, "flush failed", "/private/upload.pt")),
("read", ConnectionResetError(errno.ECONNRESET, "connection reset")),
("read", TimeoutError("read timeout")),
("truncated", None),
):
with self.subTest(phase=phase, error=error):
Handler.read_error = error if phase == "read" else None
@contextmanager
def faulty_output(path, *args, phase=phase, error=error, **kwargs):
with original_open(path, *args, **kwargs) as output:
if phase == "write":
output = unittest.mock.Mock(wraps=output)
output.write.side_effect = error
yield output
if phase == "close":
raise error
def open_file(
path,
*args,
phase=phase,
error=error,
output_factory=faulty_output,
**kwargs,
):
if path.name == "upload.pt" and args == ("xb",):
if phase == "open":
raise error
return output_factory(path, *args, **kwargs)
return original_open(path, *args, **kwargs)
conn = http.client.HTTPConnection(*server.server_address, timeout=10)
try:
with patch.object(Path, "open", open_file):
conn.request(
"POST",
url,
body=b"x",
headers={
"Authorization": "Bearer upload-test-token",
"Content-Type": "application/octet-stream",
"Content-Length": "2" if phase == "truncated" else "1",
},
)
if phase == "truncated":
conn.sock.shutdown(socket.SHUT_WR)
response = conn.getresponse()
message = json.loads(response.read())["error"]
finally:
conn.close()
storage_error = phase in ("open", "write", "close")
self.assertEqual(response.status, 503 if storage_error else 400)
self.assertIn("磁盘空间/权限" if storage_error else "中断", message)
self.assertNotIn("/private", message)
self.assertNotIn(str(self.root), message)
self.assertEqual(reset_timeouts[-1], 10)
self.assertEqual(list(self.registry.store.iterdir()), [])
self.assertEqual(self.registry.catalog(), [])
self.assertEqual(manager.jobs, {})
finally:
server.shutdown()
server.server_close()
thread.join()
def test_http_auth_json_limit_binary_headers_and_incomplete_body(self):
manager = TrainingManager(
ROOT / "rl",
sys.executable,
("Unitree-Go2-Flat",),
check_environment=False,
sources=self.registry,
)
class Handler(TrainingRequestHandler):
access_token = "upload-test-token"
def log_message(self, *_args):
pass
Handler.manager = manager
server = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
url = (
"/api/training/pretrained-sources/upload"
"?format=pt&template=go2-legacy47-v1&name=..%2Ftest.pt"
)
headers = {
"Authorization": "Bearer upload-test-token",
"Content-Type": "application/octet-stream",
}
def request(path=url, data=b"x", extra=None):
conn = http.client.HTTPConnection(*server.server_address, timeout=30)
conn.request("POST", path, body=data, headers={**headers, **(extra or {})})
response = conn.getresponse()
status, body = response.status, response.read()
conn.close()
return status, json.loads(body)
try:
self.assertEqual(request(extra={"Authorization": "Bearer wrong"})[0], 401)
self.assertEqual(request(extra={"Origin": "https://evil.example"})[0], 403)
self.assertEqual(request(extra={"Host": "evil.example"})[0], 403)
self.assertEqual(request(extra={"Content-Type": "application/json"})[0], 400)
self.assertEqual(request(extra={"Content-Length": str(256 * 1024**2 + 1)})[0], 413)
self.assertEqual(request(extra={"Content-Encoding": "gzip"})[0], 400)
self.assertEqual(request(extra={"Transfer-Encoding": "chunked"})[0], 400)
self.assertEqual(
request(path=url.replace("template=go2-legacy47-v1", "template="))[0], 400
)
self.assertEqual(
request(path="/api/training/jobs", data=b" " * (128 * 1024 + 1))[0], 413
)
# A blocked decoder is isolated from health/catalog requests.
import subprocess
from types import SimpleNamespace
Handler.tuning_manager = SimpleNamespace(capability=lambda: {})
started, release = threading.Event(), threading.Event()
original_run = subprocess.run
def blocked(*args, **kwargs):
started.set()
release.wait(10)
return original_run(*args, **kwargs)
outcomes = []
with patch("pretrained_sources.subprocess.run", side_effect=blocked):
worker = threading.Thread(target=lambda: outcomes.append(request(data=self.pt)))
worker.start()
self.assertTrue(started.wait(5))
health = http.client.HTTPConnection(*server.server_address, timeout=3)
health.request("GET", "/api/training/health", headers=headers)
response = health.getresponse()
self.assertEqual(response.status, 200)
self.assertTrue(json.loads(response.read())["pretrainedUpload"]["enabled"])
health.close()
release.set()
worker.join(30)
status, record = outcomes[0]
self.assertEqual(status, 201)
self.assertEqual(record["label"], "test.pt")
status, onnx_record = request(
path=url.replace("format=pt", "format=onnx"), data=self.onnx
)
self.assertEqual(status, 201)
self.assertEqual(onnx_record["initialization"]["manifest"]["sourceFormat"], "onnx")
raw = socket.create_connection(server.server_address, timeout=10)
raw.sendall(
(
f"POST {url} HTTP/1.0\r\nHost: localhost\r\n"
"Authorization: Bearer upload-test-token\r\n"
"Content-Type: application/octet-stream\r\nContent-Length: 1000\r\n\r\nx"
).encode()
)
raw.shutdown(socket.SHUT_WR)
self.assertIn(b"400", raw.recv(4096))
raw.close()
self.assertFalse(list(self.registry.store.glob("upload-*")))
finally:
server.shutdown()
server.server_close()
thread.join()
def test_onnx_expansion_and_checkpoint_preserve_updated_count(self):
import torch
from pretrained import (
ValidatedSource,
comparison_observations,
make_reference_actor,
warm_start_actor,
)
from pretrained_upload import _onnx
from tensordict import TensorDict
state, _, identity = _onnx(self.onnx)
self.assertLess(identity["max_abs_error"], 2e-5)
base = comparison_observations()
reference = make_reference_actor().eval()
reference.load_state_dict(state)
expected = reference.mlp(reference.obs_normalizer(base)).detach()
for dim in (81, 97):
actor = make_reference_actor(dim)
warm_start_actor(actor, ValidatedSource(state, {}))
self.assertEqual(actor.mlp[0].weight[:, 47:].count_nonzero().item(), 0)
values = torch.cat((base, torch.ones(len(base), dim - 47)), dim=1)
torch.testing.assert_close(
actor.mlp(actor.obs_normalizer(values)), expected, atol=2e-5, rtol=2e-5
)
batch = TensorDict({"actor": values}, batch_size=[len(base)])
actor.update_normalization(batch)
self.assertEqual(actor.obs_normalizer.count.item(), 1_000_048)
self.assertGreater(actor.obs_normalizer._mean[:, 47:].min().item(), 0)
actor(batch).square().mean().backward()
grad = actor.mlp[0].weight.grad[:, 47:]
self.assertTrue(torch.isfinite(grad).all())
self.assertGreater(grad.abs().max().item(), 0)
torch.optim.Adam(actor.parameters(), lr=1e-4).step()
buffer = io.BytesIO()
torch.save(actor.state_dict(), buffer)
buffer.seek(0)
restored = make_reference_actor(dim)
restored.load_state_dict(torch.load(buffer, weights_only=True))
for key, tensor in actor.state_dict().items():
self.assertTrue(torch.equal(tensor, restored.state_dict()[key]), key)
@unittest.skipUnless(
os.environ.get("GO2_UPLOAD_REAL_DIR"), "opt-in read-only real single-file probes"
)
class RealSingleFileUploadTest(UploadTest):
"""Inherited tests run against real uploads without ever requesting a sidecar."""
@classmethod
def setUpClass(cls):
import torch
torch.set_num_threads(1)
root = Path(os.environ["GO2_UPLOAD_REAL_DIR"])
# Explicit two files read separately as bytes; each upload receives only one.
cls.pt = (root / "model_10000.pt").read_bytes()
cls.onnx = (root / "policy.onnx").read_bytes()
cls.state = torch.load(io.BytesIO(cls.pt), map_location="cpu", weights_only=True)[
"actor_state_dict"
]
def test_real_independent_uploads_ort_and_batch_drift(self):
import torch
from pretrained import comparison_observations, make_reference_actor, verify_onnx
from pretrained_upload import read_uploaded_source
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
evidence = {}
for fmt, data in (("pt", self.pt), ("onnx", self.onnx)):
with tempfile.TemporaryDirectory() as empty:
registry = PretrainedSources(None, Path(empty), sys.executable, ROOT / "rl")
record = registry.receive_upload(
io.BytesIO(data), len(data), fmt, "go2-legacy47-v1", f"single.{fmt}"
)
directory = registry.verify(record["initialization"])
source = read_uploaded_source(
directory / "actor.pt",
allowed_roots=[directory],
manifest_path=directory / "upload.json",
target_env=asdict(unitree_go2_flat_env_cfg()),
target_agent=asdict(unitree_go2_ppo_runner_cfg()),
)
actor = make_reference_actor().eval()
actor.load_state_dict(source.actor_state)
identity = verify_onnx(self.onnx, actor)
self.assertLess(identity["max_abs_error"], 2e-5)
evidence[fmt] = {
"sourceId": record["id"],
"sha256": hashlib.sha256(data).hexdigest(),
"identity": identity,
"normalizer_count": actor.obs_normalizer.count.item(),
"updates": {},
}
probes = comparison_observations()[32:]
before = actor.mlp(actor.obs_normalizer(probes)).detach()
# Ordinary UI default=4096 and service maximum=16384 environments.
for batch_size in (4096, 16384):
updated = copy.deepcopy(actor).train()
updated.obs_normalizer.update(probes.repeat(batch_size // len(probes), 1))
after = updated.mlp(updated.obs_normalizer(probes)).detach()
evidence[fmt]["updates"][str(batch_size)] = {
"rate": batch_size / updated.obs_normalizer.count.item(),
"max_action_drift": (after - before).abs().max().item(),
"max_mean_drift": (
updated.obs_normalizer._mean - actor.obs_normalizer._mean
)
.abs()
.max()
.item(),
"count_after": updated.obs_normalizer.count.item(),
}
self.assertTrue(torch.isfinite(after).all())
self.assertEqual(
PretrainedSources(None, Path(empty), sys.executable, ROOT / "rl").catalog(),
[record],
)
print("REAL_UPLOAD_EVIDENCE=" + json.dumps(evidence, sort_keys=True))
if __name__ == "__main__":
unittest.main()