49 lines
1.8 KiB
Python
49 lines
1.8 KiB
Python
"""Service-owned CPU validator; JSON stdin is never accepted directly from HTTP."""
|
|
|
|
import json
|
|
import sys
|
|
from dataclasses import asdict
|
|
from pathlib import Path
|
|
|
|
ROOT = Path(__file__).resolve().parents[1]
|
|
for root in (ROOT, ROOT.parent):
|
|
sys.path.insert(0, str(root))
|
|
|
|
|
|
def main():
|
|
from pretrained import read_pretrained_source
|
|
from task_config import OBSTACLE_TASK
|
|
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
|
|
|
|
payload = json.loads(sys.stdin.read(128 * 1024))
|
|
task_id = payload["taskId"]
|
|
if task_id == "Unitree-Go2-Flat":
|
|
cfg = unitree_go2_flat_env_cfg()
|
|
elif task_id == OBSTACLE_TASK:
|
|
from src.tasks.obstacle_avoidance.env_cfg import unitree_go2_obstacle_env_cfg, apply_obstacle_configuration
|
|
cfg = unitree_go2_obstacle_env_cfg()
|
|
if payload.get("taskConfig") is not None:
|
|
apply_obstacle_configuration(cfg, payload["taskConfig"])
|
|
else:
|
|
raise ValueError("基础策略不兼容该任务;支持Flat47与Obstacle81/97,不支持Rough")
|
|
directory = Path(payload["directory"])
|
|
options = {}
|
|
if payload.get("uploaded"):
|
|
from pretrained_upload import read_uploaded_source
|
|
read_pretrained_source = read_uploaded_source
|
|
options["manifest_path"] = directory / "upload.json"
|
|
source = read_pretrained_source(
|
|
directory / payload["checkpoint"], allowed_roots=[directory], target_env=asdict(cfg),
|
|
target_agent=asdict(unitree_go2_ppo_runner_cfg()), **options,
|
|
)
|
|
print(json.dumps(source.manifest))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
try:
|
|
main()
|
|
except Exception as error:
|
|
print(f"基础策略不兼容或缺少配套文件/依赖:{error}", file=sys.stderr)
|
|
sys.exit(1)
|