736 lines
31 KiB
Python
736 lines
31 KiB
Python
"""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()
|