"""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()