308 lines
14 KiB
Python
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()
|