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

308 lines
14 KiB
Python

"""Registered source authority, immutable snapshots, and same-trial continuation."""
import hashlib
import json
import os
import sys
import tempfile
import threading
import unittest
from pathlib import Path
from unittest.mock import patch
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from pretrained_sources import PretrainedSources, SourceError, regular_bytes
from server import ApiError, TrainingManager
from test_tuning_manager import BASE_METRICS, FakeAdvisor
from tuning.manager import TuningError, TuningManager
from tuning.process import GpuLease
from tuning.schema import BASE_REWARD_CONFIGURATION, RewardConfigError, validate_proposal
def fake_validate(_self, directory, checkpoint, *_args):
artifacts = {}
for key, relative in {
"checkpoint": checkpoint,
"onnx": "policy.onnx",
"env": "params/env.yaml",
"agent": "params/agent.yaml",
}.items():
data = (directory / relative).read_bytes()
artifacts[key] = {
"name": Path(relative).name,
"sha256": hashlib.sha256(data).hexdigest(),
"bytes": len(data),
}
return {
"source_id": hashlib.sha256(json.dumps(artifacts, sort_keys=True).encode()).hexdigest(),
"artifacts": artifacts,
"source_iteration": 10000,
}
class RegisteredSourcesTest(unittest.TestCase):
def setUp(self):
self.temp = tempfile.TemporaryDirectory()
self.root = Path(self.temp.name)
self.original = self.root / "original"
(self.original / "params").mkdir(parents=True)
for relative in ("model_10000.pt", "policy.onnx", "params/env.yaml", "params/agent.yaml"):
(self.original / relative).write_bytes(relative.encode())
self.config = self.root / "sources.json"
self.config.write_text(
json.dumps(
{
"allowedRoots": [str(self.original)],
"sources": [
{
"id": "base",
"label": "基础行走",
"checkpoint": str(self.original / "model_10000.pt"),
"onnx": str(self.original / "policy.onnx"),
}
],
}
)
)
self.validator = patch.object(PretrainedSources, "_validate", fake_validate)
self.validator.start()
self.registry = PretrainedSources(
self.config,
self.root / "snapshots",
sys.executable,
Path(__file__).resolve().parents[1] / "rl",
)
def tearDown(self):
self.validator.stop()
self.temp.cleanup()
def test_registered_snapshot_never_rereads_mutated_original_and_checks_sha(self):
bound = self.registry.bind(self.registry.catalog()[0]["id"], "Unitree-Go2-Flat")
original_sha = bound["manifest"]["artifacts"]["checkpoint"]["sha256"]
(self.original / "model_10000.pt").write_bytes(b"changed after registration")
self.assertEqual(
self.registry.bind(self.registry.catalog()[0]["id"], "Unitree-Go2-Flat"), bound
)
args = self.registry.arguments(bound)
checkpoint = Path(args[args.index("--pretrained-checkpoint") + 1])
self.assertEqual(hashlib.sha256(checkpoint.read_bytes()).hexdigest(), original_sha)
restored = PretrainedSources(None, self.root / "snapshots", sys.executable, self.root)
self.assertEqual(restored.arguments(bound), args)
re_registered = PretrainedSources(
self.config, self.root / "snapshots", sys.executable, self.root
)
with self.assertRaisesRegex(SourceError, "已变化"):
re_registered.bind(bound["sourceId"], "Unitree-Go2-Flat")
self.assertEqual(restored.arguments(bound), args)
checkpoint.chmod(0o644)
checkpoint.write_bytes(b"tamper")
with self.assertRaisesRegex(SourceError, "SHA"):
restored.arguments(bound)
checkpoint.parent.chmod(0o755)
checkpoint.unlink()
with self.assertRaises(SourceError):
restored.verify(bound)
def test_rejects_path_symlink_fifo_size_and_invalid_registry(self):
with self.assertRaises(SourceError):
self.registry.bind(str(self.original / "model_10000.pt"), "Unitree-Go2-Flat")
with self.assertRaisesRegex(SourceError, "Rough"):
self.registry.bind(self.registry.catalog()[0]["id"], "Unitree-Go2-Rough")
(self.original / "link.pt").symlink_to(self.original / "model_10000.pt")
with self.assertRaises(SourceError):
regular_bytes(self.original / "link.pt", self.original, 256)
with self.assertRaises(SourceError):
regular_bytes(self.original / ".." / "sources.json", self.original, 256)
with self.assertRaises(SourceError):
regular_bytes(self.original / "policy.onnx", self.original, 1)
os.mkfifo(self.original / "pipe")
with self.assertRaises(SourceError):
regular_bytes(self.original / "pipe", self.original, 256)
self.config.write_text('{"allowedRoots": [], "sources": [], "python": "evil"}')
with self.assertRaises(SourceError):
PretrainedSources(self.config, self.root / "other", sys.executable, self.root)
def test_missing_checkpoint_is_visible_failure_not_random_fallback(self):
(self.original / "model_10000.pt").unlink()
registry = PretrainedSources(self.config, self.root / "other", sys.executable, self.root)
self.assertFalse(registry.catalog()[0]["ready"])
with self.assertRaisesRegex(SourceError, "pt"):
registry.bind("base", "Unitree-Go2-Flat")
def test_training_request_only_id_and_public_identity(self):
manager = TrainingManager(
Path(__file__).resolve().parents[1] / "rl",
sys.executable,
("Unitree-Go2-Flat", "Unitree-Go2-Rough"),
check_environment=False,
sources=self.registry,
)
payload = {
"taskId": "Unitree-Go2-Flat",
"numEnvs": 4,
"maxIterations": 1,
"seed": 42,
"device": "cpu",
"pretrainedSourceId": self.registry.catalog()[0]["id"],
}
config = manager.parse_config(payload)
self.assertEqual(
config.pretrained,
self.registry.bind(self.registry.catalog()[0]["id"], "Unitree-Go2-Flat"),
)
command = manager.command_for(config)
self.assertIn("--pretrained-source-id", command)
self.assertNotIn("--resume-checkpoint", command)
self.assertNotIn(str(self.original), json.dumps(manager.health()["pretrainedSources"]))
for key in ("pretrainedCheckpoint", "allowedRoots", "pretrained", "resumeCheckpoint"):
with self.assertRaises(ApiError):
manager.parse_config({**payload, key: "/etc/passwd"})
with self.assertRaises(ApiError):
manager.parse_config({**payload, "pretrainedSourceId": None})
with self.assertRaises(ApiError):
manager.parse_config({**payload, "taskId": "Unitree-Go2-Rough"})
old = manager.parse_config({k: v for k, v in payload.items() if k != "pretrainedSourceId"})
self.assertIsNone(old.pretrained)
self.assertNotIn("--pretrained-checkpoint", manager.command_for(old))
def tuning(self):
return TuningManager(
Path(__file__).resolve().parents[1] / "rl",
sys.executable,
self.root / "tuning",
GpuLease(),
advisor=FakeAdvisor(),
sources=self.registry,
)
def session(self, manager, mode="approval"):
_, config, objective, fallback = manager.parse_create(
{
"mode": mode,
"pretrainedSourceId": self.registry.catalog()[0]["id"],
"trialCount": 2,
"numEnvs": 4,
"initialIterations": 1,
"middleIterations": 2,
"finalIterations": 3,
}
)
return manager.storage.create_session(mode, config, objective, fallback)
def test_baseline_new_trial_warmstart_and_rung_resumes_only_own_checkpoint(self):
manager = self.tuning()
session = self.session(manager)
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"trial checkpoint")
(directory / "policy.onnx").write_bytes(b"trial policy")
(directory / "initialization.json").write_text(
json.dumps(session["config"]["pretrained"]["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]
self.assertEqual(
train[0][train[0].index("--pretrained-source-id") + 1],
train[1][train[1].index("--pretrained-source-id") + 1],
)
self.assertIn("--resume-checkpoint", train[2])
self.assertNotIn("--pretrained-checkpoint", train[2])
self.assertIn("trial-1-0/model_0.pt", train[2][-1])
self.assertEqual(len([c for c in commands if "--steps-per-seed=1000" in c]), 3)
with self.assertRaises(RewardConfigError):
validate_proposal({"pretrainedSourceId": "other"}, BASE_REWARD_CONFIGURATION)
def test_restart_states_source_constraints_and_explicit_recovery(self):
manager = self.tuning()
sessions = []
for state in ("queued", "paused", "awaiting_approval", "succeeded"):
session = self.session(manager)
trial = manager.storage.create_trial(
session["id"], 0, 0, 1, BASE_REWARD_CONFIGURATION, None, "baseline"
)
manager.storage.update_trial(
trial["id"],
state="completed",
evaluation={"metrics": BASE_METRICS},
eligible=True,
score=0,
)
proposal = manager.storage.create_proposal(
session["id"], trial["id"], {"weights": {"pose": 1.1}}, "proposal", {}, 0.8
)
manager.storage.replace_constraints(
session["id"], 0, {"weights.pose": {"kind": "range", "min": 0, "max": 2}}
)
manager.storage.update_session(session["id"], state=state)
sessions.append((session, proposal, state))
restored = self.tuning() # Actual SQLite reopening, no worker/API calls on construction.
self.assertFalse(restored.workers)
for session, _proposal, previous in sessions:
value = restored.detail(session["id"])
self.assertEqual(
value["state"], "succeeded" if previous == "succeeded" else "interrupted"
)
self.assertEqual(value["config"]["pretrained"], session["config"]["pretrained"])
self.assertEqual(value["control"]["constraintsRevision"], 1)
self.assertEqual(value["mode"], "approval")
self.assertEqual(value["proposals"][0]["state"], "pending")
session, proposal, _ = sessions[0]
# Stop before dispatch: recovery must invalidate pending first, never execute it.
with (
patch.object(restored, "_wait_for_dispatch", side_effect=RuntimeError("test boundary")),
patch.object(restored.advisor, "propose") as advisor,
):
restored._run_session(session["id"], True, threading.Event())
advisor.assert_not_called()
self.assertEqual(restored.storage.get_proposal(proposal["id"])["state"], "rejected")
self.assertIn(
"recovery_invalidated", restored.storage.get_proposal(proposal["id"])["feedback"]
)
self.assertEqual(len(restored.storage.list_trials(session["id"])), 1)
session = sessions[1][0]
with patch.object(restored, "_start_worker") as start:
restored.resume(session["id"])
with self.assertRaises(TuningError):
restored.resume(session["id"])
start.assert_called_once()
session = sessions[2][0]
snapshot = self.registry.verify(session["config"]["pretrained"])
file = snapshot / "policy.onnx"
file.chmod(0o644)
file.write_bytes(b"corrupt")
with patch.object(restored, "_start_worker") as start:
with self.assertRaises(SourceError):
restored.resume(session["id"])
start.assert_not_called()
self.assertEqual(restored.storage.get_session(session["id"])["state"], "interrupted")
if __name__ == "__main__":
unittest.main()