From 438e56bcc887e5b0453990a54a09258711b6a17e Mon Sep 17 00:00:00 2001 From: cen617-code <1057290604@qq.com> Date: Tue, 8 Sep 2026 10:50:13 +0800 Subject: [PATCH 1/4] =?UTF-8?q?feat(training):=20release=20V0.9.1=20?= =?UTF-8?q?=E9=81=BF=E9=9A=9C=E8=AE=AD=E7=BB=83=E4=B8=8E=E5=9F=BA=E7=A1=80?= =?UTF-8?q?=E7=AD=96=E7=95=A5=E8=BF=81=E7=A7=BB?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- AGENTS.md | 3 +- CHANGELOG.md | 96 +++ package-lock.json | 4 +- package.json | 2 +- training_server/OBSTACLE_AVOIDANCE.md | 188 +++++ training_server/PRETRAINED.md | 178 +++++ training_server/README.md | 24 +- training_server/pretrained.py | 502 ++++++++++++ training_server/pretrained_sources.py | 447 +++++++++++ training_server/pretrained_upload.py | 336 ++++++++ training_server/rl/requirements.txt | 4 + training_server/rl/scripts/evaluate.py | 62 +- .../rl/scripts/evaluate_obstacle.py | 241 ++++++ training_server/rl/scripts/train.py | 127 ++- .../rl/scripts/validate_pretrained.py | 48 ++ training_server/rl/scripts/validate_upload.py | 43 + .../src/tasks/obstacle_avoidance/__init__.py | 17 + .../src/tasks/obstacle_avoidance/env_cfg.py | 89 +++ .../rl/src/tasks/obstacle_avoidance/mdp.py | 231 ++++++ .../src/tasks/obstacle_avoidance/terrain.py | 70 ++ .../rl/src/tasks/velocity/rl/runner.py | 15 + training_server/server.py | 162 +++- training_server/task_config.py | 479 ++++++++++++ .../tests/fixtures/custom-boxes.json | 14 + .../tests/generate_multi_ring_golden.py | 90 +++ training_server/tests/test_custom_boxes.py | 254 ++++++ training_server/tests/test_multi_ring.py | 214 +++++ training_server/tests/test_obstacle_env.py | 223 ++++++ training_server/tests/test_obstacle_tuning.py | 512 ++++++++++++ training_server/tests/test_pretrained.py | 475 +++++++++++ .../tests/test_pretrained_sources.py | 307 ++++++++ .../tests/test_pretrained_upload.py | 735 ++++++++++++++++++ .../tests/test_pretrained_upload_runner.py | 163 ++++ .../tests/test_reward_preset_tasks.py | 137 ++++ training_server/tests/test_server.py | 74 +- training_server/tests/test_task_config.py | 83 ++ training_server/tuning/advisor.py | 7 +- training_server/tuning/manager.py | 280 ++++++- training_server/tuning/obstacle_scoring.py | 130 ++++ training_server/tuning/schema.py | 107 ++- training_server/tuning/storage.py | 85 +- web_platform/TRAINING.md | 80 ++ web_platform/e2e/multiRing.spec.ts | 153 ++++ web_platform/e2e/obstacle.spec.ts | 291 +++++++ web_platform/e2e/pretrainedNavigation.spec.ts | 209 +++++ web_platform/e2e/pretrainedUpload.spec.ts | 177 +++++ web_platform/fixtures/obstacle/README.md | 9 + .../fixtures/obstacle/legacy-flat.onnx | Bin 0 -> 177 bytes .../fixtures/obstacle/legacy-wrong-shape.onnx | Bin 0 -> 177 bytes .../fixtures/obstacle/multi-wrong-graph.onnx | Bin 0 -> 8573 bytes .../fixtures/obstacle/multi-zero-action.onnx | Bin 0 -> 8573 bytes .../fixtures/obstacle/ort-init-failure.onnx | Bin 0 -> 4510 bytes .../fixtures/obstacle/wrong-graph.onnx | Bin 0 -> 4563 bytes .../fixtures/obstacle/zero-action.onnx | Bin 0 -> 4568 bytes web_platform/src/app/App.tsx | 207 +++-- .../app/components/WorkspaceToolsPanel.tsx | 31 +- .../components/charts/ScalarChart.test.tsx | 47 ++ .../src/components/charts/ScalarChart.tsx | 216 +++++ .../src/map/customTrainingMap.test.ts | 137 ++++ web_platform/src/map/trainingMap.test.ts | 191 +++++ web_platform/src/map/trainingMap.ts | 199 +++++ web_platform/src/rl/RLPolicyPanel.test.tsx | 64 ++ web_platform/src/rl/RLPolicyPanel.tsx | 109 ++- web_platform/src/rl/deployment.test.ts | 127 +++ web_platform/src/rl/deployment.ts | 406 ++++++++++ .../src/rl/fixtures/go2CompiledDynamics.json | 279 +++++++ .../src/rl/fixtures/multiRingDeployment.json | 221 ++++++ .../src/rl/fixtures/multiRingGolden.json | 336 ++++++++ .../src/rl/fixtures/obstacleDeployment.json | 212 +++++ .../src/rl/fixtures/obstacleRayGolden.json | 90 +++ .../Go2ObstacleAvoidanceBindings.test.ts | 55 ++ .../runtime/Go2ObstacleAvoidanceBindings.ts | 165 ++++ .../src/rl/runtime/Go2wPolicyBindings.ts | 13 +- .../src/rl/runtime/OnnxPolicyRuntime.test.ts | 217 ++++++ .../src/rl/runtime/OnnxPolicyRuntime.ts | 72 +- .../src/rl/tasks/go2ObstacleAvoidance.test.ts | 114 +++ .../src/rl/tasks/go2ObstacleAvoidance.ts | 160 ++++ web_platform/src/rl/tasks/multiRing.test.ts | 90 +++ web_platform/src/rl/tasks/raycastBenchmark.ts | 60 ++ web_platform/src/rl/types.ts | 10 +- ...PhysicsAdapter.trainingTransaction.test.ts | 259 ++++++ web_platform/src/simulation/PhysicsAdapter.ts | 150 +++- .../SimulationSession.policyLifecycle.test.ts | 71 ++ .../src/simulation/SimulationSession.ts | 134 +++- .../src/training/LocalTrainingClient.test.ts | 63 ++ .../src/training/LocalTrainingClient.ts | 21 + .../src/training/LocalTrainingPanel.test.tsx | 310 ++++++++ .../src/training/LocalTrainingPanel.tsx | 474 ++++++++++- .../training/PretrainedSourceSelect.test.tsx | 358 +++++++++ .../src/training/PretrainedSourceSelect.tsx | 123 +++ .../src/training/PretrainedUpload.tsx | 111 +++ .../training/TrainingMetricHistory.test.ts | 63 ++ .../src/training/TrainingMetricHistory.ts | 93 +++ .../training/TrainingMetricsPanel.test.tsx | 36 + .../src/training/TrainingMetricsPanel.tsx | 81 ++ .../src/training/pretrainedSelection.ts | 21 + .../src/training/trainingLosses.test.ts | 18 + web_platform/src/training/trainingLosses.ts | 14 + web_platform/src/training/types.ts | 95 ++- .../src/tuning/AgentDecisionTimeline.tsx | 5 +- .../src/tuning/MetricsComparisonBoard.tsx | 20 +- web_platform/src/tuning/ScalarChart.tsx | 206 +---- web_platform/src/tuning/TuningApp.tsx | 167 +++- web_platform/src/tuning/TuningConsole.tsx | 6 + .../src/tuning/TuningControlToolbar.tsx | 19 +- web_platform/src/tuning/TuningSessionRail.tsx | 7 +- web_platform/src/tuning/domain.ts | 32 +- .../src/tuning/obstacleTuning.test.tsx | 93 +++ web_platform/src/tuning/tuningStore.ts | 23 +- .../viewer/MuJoCoViewer.perception.test.ts | 36 + web_platform/src/viewer/MuJoCoViewer.ts | 104 +++ .../src/viewer/NavigationGoal.test.ts | 178 +++++ web_platform/src/viewer/NavigationGoal.ts | 170 ++++ 113 files changed, 15027 insertions(+), 539 deletions(-) create mode 100644 training_server/OBSTACLE_AVOIDANCE.md create mode 100644 training_server/PRETRAINED.md create mode 100644 training_server/pretrained.py create mode 100644 training_server/pretrained_sources.py create mode 100644 training_server/pretrained_upload.py create mode 100644 training_server/rl/scripts/evaluate_obstacle.py create mode 100644 training_server/rl/scripts/validate_pretrained.py create mode 100644 training_server/rl/scripts/validate_upload.py create mode 100644 training_server/rl/src/tasks/obstacle_avoidance/__init__.py create mode 100644 training_server/rl/src/tasks/obstacle_avoidance/env_cfg.py create mode 100644 training_server/rl/src/tasks/obstacle_avoidance/mdp.py create mode 100644 training_server/rl/src/tasks/obstacle_avoidance/terrain.py create mode 100644 training_server/task_config.py create mode 100644 training_server/tests/fixtures/custom-boxes.json create mode 100644 training_server/tests/generate_multi_ring_golden.py create mode 100644 training_server/tests/test_custom_boxes.py create mode 100644 training_server/tests/test_multi_ring.py create mode 100644 training_server/tests/test_obstacle_env.py create mode 100644 training_server/tests/test_obstacle_tuning.py create mode 100644 training_server/tests/test_pretrained.py create mode 100644 training_server/tests/test_pretrained_sources.py create mode 100644 training_server/tests/test_pretrained_upload.py create mode 100644 training_server/tests/test_pretrained_upload_runner.py create mode 100644 training_server/tests/test_reward_preset_tasks.py create mode 100644 training_server/tests/test_task_config.py create mode 100644 training_server/tuning/obstacle_scoring.py create mode 100644 web_platform/TRAINING.md create mode 100644 web_platform/e2e/multiRing.spec.ts create mode 100644 web_platform/e2e/obstacle.spec.ts create mode 100644 web_platform/e2e/pretrainedNavigation.spec.ts create mode 100644 web_platform/e2e/pretrainedUpload.spec.ts create mode 100644 web_platform/fixtures/obstacle/README.md create mode 100644 web_platform/fixtures/obstacle/legacy-flat.onnx create mode 100644 web_platform/fixtures/obstacle/legacy-wrong-shape.onnx create mode 100644 web_platform/fixtures/obstacle/multi-wrong-graph.onnx create mode 100644 web_platform/fixtures/obstacle/multi-zero-action.onnx create mode 100644 web_platform/fixtures/obstacle/ort-init-failure.onnx create mode 100644 web_platform/fixtures/obstacle/wrong-graph.onnx create mode 100644 web_platform/fixtures/obstacle/zero-action.onnx create mode 100644 web_platform/src/components/charts/ScalarChart.test.tsx create mode 100644 web_platform/src/components/charts/ScalarChart.tsx create mode 100644 web_platform/src/map/customTrainingMap.test.ts create mode 100644 web_platform/src/map/trainingMap.test.ts create mode 100644 web_platform/src/map/trainingMap.ts create mode 100644 web_platform/src/rl/RLPolicyPanel.test.tsx create mode 100644 web_platform/src/rl/deployment.test.ts create mode 100644 web_platform/src/rl/deployment.ts create mode 100644 web_platform/src/rl/fixtures/go2CompiledDynamics.json create mode 100644 web_platform/src/rl/fixtures/multiRingDeployment.json create mode 100644 web_platform/src/rl/fixtures/multiRingGolden.json create mode 100644 web_platform/src/rl/fixtures/obstacleDeployment.json create mode 100644 web_platform/src/rl/fixtures/obstacleRayGolden.json create mode 100644 web_platform/src/rl/runtime/Go2ObstacleAvoidanceBindings.test.ts create mode 100644 web_platform/src/rl/runtime/Go2ObstacleAvoidanceBindings.ts create mode 100644 web_platform/src/rl/runtime/OnnxPolicyRuntime.test.ts create mode 100644 web_platform/src/rl/tasks/go2ObstacleAvoidance.test.ts create mode 100644 web_platform/src/rl/tasks/go2ObstacleAvoidance.ts create mode 100644 web_platform/src/rl/tasks/multiRing.test.ts create mode 100644 web_platform/src/rl/tasks/raycastBenchmark.ts create mode 100644 web_platform/src/simulation/PhysicsAdapter.trainingTransaction.test.ts create mode 100644 web_platform/src/simulation/SimulationSession.policyLifecycle.test.ts create mode 100644 web_platform/src/training/PretrainedSourceSelect.test.tsx create mode 100644 web_platform/src/training/PretrainedSourceSelect.tsx create mode 100644 web_platform/src/training/PretrainedUpload.tsx create mode 100644 web_platform/src/training/TrainingMetricHistory.test.ts create mode 100644 web_platform/src/training/TrainingMetricHistory.ts create mode 100644 web_platform/src/training/TrainingMetricsPanel.test.tsx create mode 100644 web_platform/src/training/TrainingMetricsPanel.tsx create mode 100644 web_platform/src/training/pretrainedSelection.ts create mode 100644 web_platform/src/training/trainingLosses.test.ts create mode 100644 web_platform/src/training/trainingLosses.ts create mode 100644 web_platform/src/tuning/obstacleTuning.test.tsx create mode 100644 web_platform/src/viewer/MuJoCoViewer.perception.test.ts create mode 100644 web_platform/src/viewer/NavigationGoal.test.ts create mode 100644 web_platform/src/viewer/NavigationGoal.ts diff --git a/AGENTS.md b/AGENTS.md index 52914f36..d20b405c 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -1,3 +1,4 @@ Please always speak chinese. python 虚拟环境路径在/home/cen/Embodied_Workspace/Mujoco_Projects/mujoco/.venv/bin/activate -系统是Ubuntu 24.04 LTS \ No newline at end of file +系统是Ubuntu 24.04 LTS +每次提交版本前,更新CHANGELOG.md文件 \ No newline at end of file diff --git a/CHANGELOG.md b/CHANGELOG.md index c7eceeeb..86437bf1 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,6 +2,82 @@ 本项目的重要变更记录在此文件中,版本标签沿用仓库现有的 `V主版本.次版本[.修订版本]` 格式。 +## [0.9.1] - 2026-09-08 + +- Go2避障训练改为每个episode从同一连通自由区域随机采样起终点并随机化初始朝向,保证障碍/边界净空和最小目标距离;固定三seed评估及浏览器部署仍使用声明的参考起终点,兼顾泛化与可比性。 + +- 修正基础策略上传存储错误分类:文件打开/写入/关闭失败返回503并提示检查磁盘空间/权限;流读取超时或连接截断仍返回400,均保留临时清理与socket超时复位。 + +- 完成普通训练与自调参共用单文件.pt/.onnx上传入口:默认服务直接连接、Go2 legacy47模板确认、自动选择内容ID/文件SHA、取消及同文件重试;任务/连接epoch变化保留旧选择并隔离晚到响应,上传中禁止启动。补双面板真实HTTP浏览器验收与4环境×1iteration CPU ONNX-derived PPO/导出/统计恢复检查(仅链路验证)。 + +- 新增Go2 legacy47单文件.pt/受限MLP ONNX上传后端,无需管理员注册/邻接配置;认证有界二进制接收、CPU限额安全校验、持久化内容SHA来源,复用普通训练/自调参warm-start与rung自身resume。ONNX只继承确定性网络,明确模板假设、合成normalizer count及fresh训练状态;双UI文件选择待接入。 + +- 修复跨任务奖励preset混入Flat:以持久化来源session恢复权威taskId,读写与训练入口完整验证任务schema;跨任务或损坏来源在创建作业前拒绝,Flat菜单仅展示已确认Flat的preset。 +- 修复导航设定拖回起点误触:超过5px后持续记为拖动,取消/Esc/清理/卸载清空整个手势,下一次正常点击不受影响。 + +- 新增显式multi_ring_raycast三层48射线/97维部署,旧缺省32射线/81维不变;task/ONNX严格模式、pitch/yaw角度及顺序白名单,浏览器观测、真实ORT shape、导航面板/PiP动态线数贯通。 +- multi奖励仅按已验证标准底板顶面几何分类过滤地面,观测保留实际地板距离;缓存静态box/方向并优化解析slab,补真实CPU/WASM/GPU短smoke及48ray整帧基准。仍存在层间/侧后/坑盲区,未实现局部高程图。 +- 修复训练CLI未导入仓库任务注册模块导致直接启动Flat/Rough/Obstacle失败;保留独立解释器三seed评估与调参导航速度。事务导入失败面板改为保留完整diagnostic detail,不再用“模型编译失败”摘要掩盖ORT shape/初始化错误。 + +- 新增已应用场景静态碰撞多实例→权威custom_boxes导出,保留世界坐标并诚实标记AABB/标准底板近似;前后端严格校验布局、摩擦、数量及起终点圆形安全区,贯通任务JSON/ONNX与最小坐标配置。 + +- 新增 Go2 避障目标贴地呼吸信标、一次性点击设定目标模式、实时目标坐标/距离和目标复位;保持81维观测与异步held-action,输入与地图操纵器互斥。 +- 本地训练卡片新增可折叠Loss/Reward趋势,按迭代去重补全并有界保留500点;复用共享Canvas图表,综合指标独立缩放,曲线平滑但保留原值。 + +- 新增 Go2 前视32射线避障训练任务、自定义地形参数、配套碰撞地图与81维ONNX浏览器部署闭环,支持PiP和射线显示。 +- 自定义地图与策略采用候选会话事务,真实ONNX初始化/graph校验或地图绑定失败时保留原场景、策略和物理状态;默认Flat兼容旧无metadata的47维导出。 +- 对齐原Go2训练模型的关节armature、足端接触参数和自碰撞mask,补充真实WASM/CPU射线、动力学参数及浏览器回归测试。短训练仅验证链路,不代表避障收敛。 + +## [0.8.3] - 2026-09-04 + +### 新增 + +- 增加视口浮动地图工具栏(`MapViewportTools`),在视口上方提供变换模式切换、坐标空间切换和快捷对齐工具。 +- 增加三维场景大纲树(`SceneOutliner`),直观查看场景中机器人、地图实例及层级结构。 +- 增加视口快捷键挂载(`useMapEditorShortcuts`),支持快速切换操纵器模式、聚焦对象和撤销操作。 +- 增加属性检查器面板(`MapObjectInspector` 与 `RobotInspector`),支持精确调节地图对象位姿与机器人关节/执行器。 +- 增加右侧栏自适应 Tab 与多工具工作区容器(`RightSidebarTabs`、`WorkspaceToolsPanel`)。 +- 增加拖拽式连续数值调节输入组件(`ScrubbableNumberInput`)与通用垂直双栏拆分面板(`VerticalSplitPane`)。 + +### 变更 + +- 重构并收敛模型控制侧栏(`ModelControlsSidebar`)与工程侧栏(`ProjectSidebar`),移除冗余的行内变换逻辑。 +- 优化地图编辑面板(`MapEditorPanel`)与物理地图面板(`PhysicalMapPanel`),与统一检查器架构对齐。 +- 清理已落地的历史设计草案与计划文档。 + +## [0.8.2] - 2026-09-03 + +### 新增 + +- 将自调参工作台重构为高内聚的组件群:Session 列表导航轨(`TuningSessionRail`)、控制工具栏(`TuningControlToolbar`)、Agent 决策时间线(`AgentDecisionTimeline`)、指标对比看板(`MetricsComparisonBoard`)、排行榜(`TuningLeaderboard`)、日志控制台(`TuningConsole`)与参数差异对比器(`RewardConfigDiffEditor`)。 +- 建立自调参集中式状态管理(`tuningStore.ts`)与自适应轮询逻辑(`useTuningPolling.ts`)。 +- 引入标量环形缓冲区(`ScalarRingBuffer.ts`),优化长周期 TensorBoard 曲线的高频更新与渲染性能。 +- 服务端增加单步执行令牌(Step Token)调度,支持逐轮审批模式下的受控单步推进。 +- 服务端增加参数约束 CAS 校验护栏与会话级别约束同步。 +- 支持同会话内安全 Trial 的一键回滚与基准重设(Rollback)。 + +## [0.8.1] - 2026-09-03 + +### 新增 + +- 统一地图资产库与多实例堆叠系统(`MapAssetLibrary`、`MapStackComposer`),支持同源地图资产多次放置与独立位姿变换。 +- 新增参数化地图实时预览图层(`ParametricMapPreviewLayer`),在视口中实时预览程序化地形与碰撞体变换。 +- 新增程序地形贴合与落位计算(`sceneSurface.ts`),支持认证资产在程序化地形上自动重力落位。 +- 引入地图场景草稿事务化编译(`mapSceneDraft`、`EditableMapDraftCommit`),保障物理、视觉与出生点数据同步更新。 +- 增强地图对象拾取与选择机制(`compiledMapPick`),支持点击穿透、空白取消与实时变换操纵。 + +## [0.8.0] - 2026-09-02 + +### 新增 + +- 增加基于 DeepSeek Agent 的强化学习奖励函数自调参系统(`training_server/tuning/`),结合训练曲线与固定评估指标提出受限参数 patch。 +- 支持全自动(Automatic)与逐轮审批(Approval)双调度模式,集成多阶段晋级(Successive-Halving)资源淘汰机制。 +- 新增固定多场景评估逻辑(`evaluate.py`),通过固定 seed 和组合运动指令产出与奖励权重无关的客观综合评分。 +- 新增 SQLite/WAL 存储管理,持久化保留调参会话、Trial 状态、参数护栏、审计记录与 TensorBoard 标量数据。 +- 新增独立的自调参 Web 工作台(`web_platform/tuning.html`,`web_platform/src/tuning/`),包含实时曲线可视化组件(`ScalarChart.tsx`)。 +- 训练服务支持将最佳配置保存为不可变预设并在普通训练中复用,支持导出最佳 `policy.onnx`。 +- 本地训练面板(`LocalTrainingPanel.tsx`)与训练客户端集成自调参能力探测与快速跳转入口。 + ## [0.7.3] - 2026-09-01 ### 新增 @@ -62,3 +138,23 @@ - 优化响应式工作区、可访问性、首屏加载、纹理兼容性和视口交互。 - 增加碰撞体、坐标系、关节轴、质心和惯量辅助可视化。 + +### 未发布 — Obstacle DeepSeek 自调参 + +- 新增Obstacle专属四标量schema与真实reward/导航command应用,保持Flat与护栏CAS语义;导出/浏览器执行配套目标速度。 +- 增加固定3seed/1000步、首terminal前捕获的客观导航评价,固定权重和成功/跌倒硬门槛;custom保持权威地图并诚实标注;真实checkpoint/观测统计加载、逐seed进程隔离和失败关闭。 +- 自调参界面支持任务选择、专属参数/评估展示、terrain/sensor/custom配置交接及最佳Obstacle策略事务导入;增加mock Agent、打分算例、CAS、argv配置与GPU评估验证。 + +### 未发布 — 已训练基础策略迁移与持续导航 + +- 普通训练与自调参增加已验证基础策略选择、checkpoint/SHA和兼容观测展示;管理员本地注册来源,内容寻址只读快照,拒绝客户端任意路径、symlink、unsafe反序列化及失效来源。 +- 严格迁移47维actor到47/81/97:新增列零初始化,保留源归一化统计/count;新trial重置critic/optimizer/iteration,同trial晋级只resume自己的checkpoint,导出保留无路径来源metadata。 +- 服务重启后的非终态session等待显式恢复,审计原状态并使旧pending proposal失效重提;保留来源、约束和Approval模式,防重复resume worker。 +- 移除浏览器Obstacle交互导航20秒截止,保留训练/固定三seed评估协议及跌倒/越界安全停止;设点不自动启用策略,明确显示暂停/未启用/错误状态。 +- 新增来源安全、SHA快照、trial续训/重启、双面板和真实WASM/ORT持续换目标回归;4env×1iteration仅证明初始化与优化,空旷地图实测不代表复杂障碍收敛。 + +### 未发布 — 基础策略目录刷新安全修复 + +- 修复普通训练在刷新服务目录后清掉失效来源、静默降级随机初始化的问题;保留已选内容ID,不按同别名自动替换新内容。 +- 两种训练面板统一对目录缺项、验证失败及任务不兼容来源显示失效提示,禁用启动并在handler再次拦截;必须用户明确从头训练或选择有效来源。显式切换任务仍清空选择。 +- 增加两面板A→B、not-ready、缺项及兼容性变更的真实组件刷新回归,验证失效时零启动请求及用户明确选择后的请求体。 diff --git a/package-lock.json b/package-lock.json index 1a8b8740..bee50c1d 100644 --- a/package-lock.json +++ b/package-lock.json @@ -1,12 +1,12 @@ { "name": "mujoco-web-platform", - "version": "0.8.3", + "version": "0.9.1", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "mujoco-web-platform", - "version": "0.8.3", + "version": "0.9.1", "license": "Apache-2.0", "dependencies": { "@monaco-editor/react": "^4.7.0", diff --git a/package.json b/package.json index 8eb467d6..d9c46c3d 100644 --- a/package.json +++ b/package.json @@ -1,6 +1,6 @@ { "name": "mujoco-web-platform", - "version": "0.8.3", + "version": "0.9.1", "description": "基于 MuJoCo WebAssembly 的本地机器人仿真与控制平台", "private": true, "type": "module", diff --git a/training_server/OBSTACLE_AVOIDANCE.md b/training_server/OBSTACLE_AVOIDANCE.md new file mode 100644 index 00000000..c67fa0ec --- /dev/null +++ b/training_server/OBSTACLE_AVOIDANCE.md @@ -0,0 +1,188 @@ +# Go2 前视射线避障部署契约 v1 + +## API + +默认任务白名单:`Unitree-Go2-Flat`、`Unitree-Go2-Rough`、`Unitree-Go2-ObstacleAvoidance`。 +`GET /api/training/health` 的 `tasks` 保持字符串列表,新增 `taskMetadata`:每项含 `id/name/browserCompatible/terrainPresets/terrainParameters/sensorTypes/sensorParameters/mapSyncScope`。参数元数据含 `min/max/default`,计数含 `integer:true`。 + +`POST /api/training/jobs` 保留原参数,并接受: + +```json +{ + "taskId": "Unitree-Go2-ObstacleAvoidance", + "numEnvs": 1024, + "maxIterations": 1000, + "seed": 42, + "runName": "obstacle-navigation", + "device": "gpu", + "gpuIds": [0], + "wandbMode": "offline", + "terrainPreset": "discrete_obstacles", + "terrainParams": { + "size": 12, + "obstacle_count": 24, + "obstacle_height_min": 0.2, + "obstacle_height_max": 0.6, + "spacing": 1.2, + "friction": 0.8 + }, + "sensorType": "raycast", + "sensorCfg": { "fov": 90, "maxDistance": 4, "safetyDistance": 0.5, "avoidanceWeight": 2 } +} +``` + +只实现 `raycast`,**不是 RGB 或真实深度相机**;`camera_depth` 明确拒绝。`sensorCfg.type` 可选但只能是 `raycast`。默认观测数量32;显式multi模式固定48,不接受任意数量(见文末多层契约)。预设与字段严格白名单校验,布尔、非有限、超范围值拒绝。地形 `size` 8–24m、障碍物数1–100,实际数量受 `spacing` 容量限制。全部可选参数与范围见 `task_config.py`。普通训练的奖励 preset 仍仅允许 Flat;Obstacle专属自调参见文末。 + +创建和查询 job 返回 `deployment`,即该作业的确定性部署配置;列表中的 job 也携带此字段。不新增端点;ONNX 下载仍为 `/api/training/jobs/{id}/artifacts/policy.onnx`。自定义配置由服务写入作业目录 `training_config.json`,以参数数组 `--task-config ` 传入训练器,绝不使用 shell。训练器再次校验配置与seed。 + +`policy.onnx` 的 `platform_deployment` metadata 字符串是同一对象的JSON编码,旁边同时导出 `deployment.json`。ONNX中包含Actor观测归一化,不要在前端再套一层运行均值/方差。其他标准 mjlab metadata 保留。旧 Rough 实际 Actor 为234维(47+17×11 downward height scan),`browserCompatible:false`;前端必须禁用一键加载,不能按 Flat 处理。没有新配置的旧 Flat/tuning 保持原环境行为。 + +## 地图:精确碰撞布局,不是同名编辑器高度场 + +`deployment.terrain` 为 `boxes-v1`:`size/friction/boxes/spawn/spawnQuaternion/target/actualObstacleCount/approximation`。 + +- 单块正方形地图,世界原点在地图中心,+X前、+Y左、+Z上,米/弧度。每个 box 的 `pos` 是世界中心,`size` **为半尺寸**,`yaw=0`。底板占 `z=[-0.2,0]`。无额外无限地面/边墙。 +- 直接用导出的boxes构造碰撞几何,**不要重新运行浏览器的同名地形生成器**。`rough/wave` 为16×16有界box离散化,`pyramid_stairs` 为四层box;`approximation:true` 必须在UI标识训练专用近似布局。 +- 障碍物坐标由 `random.Random(seed)` 对确定性格点洗牌;宽深均0.4m、高度按范围随机,实际数量受spacing与地图容量限制。两端保留平坦出生/目标条带。无需在JS复刻Python RNG,布局本身是权威结果。最多257个boxes。 +- `spawn=[-size/2+1,0,0.32]`、四元数 `wxyz=[1,0,0,0]`、目标 `[size/2-1,0]` 是部署及固定评估的参考起终点。训练时每个episode在同一连通自由区域内重新采样起终点(障碍/边界净空>0.55m、间距>=2m)并随机化初始yaw,初始速度为0;不会跨不可通行墙体抽取目标。MuJoCo地形原点仍为参考spawn,随机复位直接写入世界坐标。 +- mjlab采用1×1 patch,同一张地图用于多个互相独立的Warp环境世界;env origin就是出生点而非地图中心。目标在每个env中用 `env_origin.xy+[size-2,0]`,不是错误地再加地图中心。 +- terrain摩擦为 `[friction,0.005,0.0001]`,priority=1、condim=3;足端startup滑动摩擦固定为相同值。浏览器应保持匹配的摩擦及足端接触配置。 +- “同步当前场景地图”导出已应用场景的 `custom_boxes` 权威布局,细则见下节;预设生成模式仍保持原行为。 + +## 已应用场景 → custom_boxes + +请求使用 `terrainPreset: "custom_boxes"`、`customTerrainBoxes: <完整boxes-v1对象>`;`terrainParams` 可省略,若提供则只能是与布局完全一致的 `{size, friction}`。服务和训练入口均重新验证,不读取客户端路径或XML,不运行地形RNG、不降级预设。作业目录task JSON、deployment JSON及ONNX metadata中的terrain使用同一布局。health的terrainPresets公布能力。 + +- 浏览器从成功编译的全部sceneMaps静态碰撞 `geom_xpos/geom_xmat/geom_size` 提取,包括工程地图路径、重复实例和嵌套body变换;不使用Three装饰包围盒、编辑器RNG或不完整XML。动态机器人、无碰撞装饰排除;未声明地图身份的其他静态碰撞报错,不能漏障碍。 +- box世界半尺寸为 `abs(R)*halfsize`;sphere、capsule、ellipsoid、cylinder使用解析精确世界AABB。**旋转box与非box转换会扩大碰撞占据区域**。mesh/hfield明确拒绝;项目地图原有include导入限制不变,已经编译模型的include变换不重新解析。 +- 固定floor为中心 `[0,0,-.1]`、半尺寸 `[size/2,size/2,.1]`。仅明确命名floor/ground/flat的朝上水平z=0 plane,以及顶面z≈0的水平支撑box归并。地下/坑底/非零高度或倾斜plane拒绝,不能静默填坑。其他障碍原坐标保留。底板覆盖整个正方形范围,可能填补支撑面间空白;UI明确说明标准化,不保证与原场景几何同构。 +- 最多256障碍+1底板,超量报错不截断;世界原点不变,尺寸8–24m,取覆盖应用地图/几何的正方形,不clamp或平移几何。XY必须在floor内,Z界限[-.2,12],所有半尺寸严格>0(包括拒绝负零),有限数值。摩擦必须统一为 `[f,.005,.0001]` 且f为.2–2,混合值报错要求用户先显式统一,不能静默覆盖。 +- `approximation` 对custom必须为true(AABB及标准底板转换);`actualObstacleCount` 必须等于boxes数减一。字段严格白名单,不信任声明值。floor必须标准。出生z固定.32,默认XY/四元数取成功加载时真实初始浮动基座姿态,不取跑动后的pose;目标建议为初始朝向前方3m,不自动寻找安全点。面板提供最小出生/目标XY配置,修改后需重新同步。 +- 出生及目标中心必须距边界>=.5m;与每个非底板障碍XY AABB的最近距离必须严格>.5m(圆形安全区,包括切触拒绝),不能删除或移动障碍以修复。四元数必须有限且归一化。后端env origin、初始化姿态、目标偏移和越界中心按自定义spawn计算。 +- 同步/提交均拒绝未应用草稿、训练部署替换后的场景、加载中或已过时场景;提交前再次比较当前编译布局。HTTP `MAX_REQUEST_BYTES` 与训练JSON上限均128KiB,足够257个double-precision boxes且有界;ONNX部署metadata上限仍100KB。 + +## 81维观测/12维动作 + +Actor顺序(float32): + +| slice | 内容 | +| ----- | --------------------------------------------------------- | +| 0:3 | `robot/imu_ang_vel`,本体坐标角速度rad/s | +| 3:6 | 本体坐标重力单位向量(静止 `[0,0,-1]`) | +| 6:9 | 下述目标导航速度命令 `[vx,0,wz]` | +| 9:11 | sin/cos(2π×episodeTime/0.6);命令范数<0.1时均为0 | +| 11:23 | 相对默认关节位置 | +| 23:35 | 关节速度,无缩放 | +| 35:47 | 上一原始12维动作,不是力矩或位置目标 | +| 47:79 | 32条前向测距,miss→1,否则clamp(distance/maxDistance,0,1) | +| 79:81 | `[headingError/π, clamp(goalDistance/size,0,1)]` | + +关节顺序FL/FR/RL/RR,每腿hip/thigh/calf;默认角 `[-.1,.9,-1.8, .1,.9,-1.8, -.1,.9,-1.8, .1,.9,-1.8]`。 +位置目标 `defaultJointPosition+0.25*rawAction`;kp每腿 `[20,20,40]`,kd `[1,1,2]`,力矩限幅 `[23.5,23.5,45]`。后端物理dt=.005,decimation=4,策略50Hz;浏览器仍可更小物理dt并采用single-in-flight/held-action。Go2-W轮式模型不是此任务的同构训练机器人,不应声称动力学严格等同。 + +### 射线 + +安装在 `base_link`。对 i=0..31: + +- `angle=(-fov/2+i*fov/31)*π/180`,局部direction=`[cos(angle),sin(angle),0]`。 +- 局部offset=`[.3,0,.05]`,origin=`baseWorldPosition+Rbase*offset`,directionWorld=`Rbase*direction`。采用**完整body姿态**,不可只用yaw。采样从右到左且包括两端;无恰好0°的中间射线。 +- 后端地形group0、机器人visual group2/collision group3,include_geom_groups=(0,)排除整个机器人。浏览器必须语义等价地仅射向地形(如其group2),不能照搬group0或只排除base_link。包含地面;倾斜时地面命中也构成感知数据。miss或超过range归一化为1。 +- 前向扇形只能在一个高度切片感知;矮障碍物/跌落/盲区并不保证可检测,不能称安全导航保证。 + +### 导航及episode + +`delta=target.xy-basePosition.xy`;distance=hypot(delta);yaw由base wxyz计算;headingError=`atan2(sin(atan2(dy,dx)-yaw),cos(...))`。未到达时 `vx=navigation.speed*max(0,cos(headingError))`(缺省.6,调参可.3–1.2)、vy=0、wz=clamp(headingError,-1,1)。distance<.5m到达后command全0、headingError=0,distance observation仍保留实际距离;不重采样新目标,不因到达施加失败惩罚。移动机器人离开半径后自然恢复导航。 + +训练episode20s,超时重置;跌倒、非足端接触force>10N、机器人中心越过地图边界内0.3m终止并重置。训练复位随机选择同一连通自由区域的起终点和初始yaw,恢复初始关节姿态及零速度,并清phase/上一动作;固定评估仍使用deployment声明的参考起终点以保持版本间可比。浏览器episode若不自动重置必须明确自己的评测行为;到达至少要停命令而非无限推进。奖励保留速度跟踪/平滑/姿态/非法接触终止惩罚,并增加朝目标世界速度投影及归一化近障平方惩罚。 + +## 验证 + +```bash +.venv/bin/python -m unittest discover -s training_server/tests +.venv/bin/python -m ruff check training_server +GO2_RUN_MJLAB_SMOKE=1 MUJOCO_GL=egl .venv/bin/python -m unittest discover -s training_server/tests -p test_obstacle_env.py +WANDB_MODE=disabled .venv/bin/python training_server/rl/scripts/train.py Unitree-Go2-ObstacleAvoidance --env.scene.num-envs=4 --agent.max-iterations=1 --gpu-ids '[0]' --output-dir /tmp/go2-obstacle-train-smoke +``` + +GPU smoke对比真实mjlab/Warp与MuJoCo `mj_ray`距离、检查多env出生点/关节顺序/81维有限obs。1iteration只验训练/导出管线,**不证明策略已学会避障或Sim2Sim行为收敛**。 + +## Obstacle DeepSeek 自调参(obstacle-v1) + +`POST /api/tuning/sessions` 新支持 `taskId=Unitree-Go2-ObstacleAvoidance`。 +沿用 automatic/approval、运行时双向切换、稀疏patch、最多4项/单轮0.5–2倍、工程护栏与revision CAS;Flat完整schema、奖励应用与相对基线六维评分不变。 +可携带与训练相同的顶层terrain/sensor字段,或单独`taskConfig`对象(不能混用);场景/seed仅创建时确定。所有字段/范围、NaN/Inf、跨任务patch/护栏拒绝;旧revision返回409且不修改配置/审计/待审批项。 + +| JSON路径 | 范围;默认 | 真实环境映射 | +| --------------------------- | -------------------------------- | -------------------------------------------------------------------------- | +| `weights.avoidance_weight` | .5–5;2(创建时可取sensorCfg值) | `obstacle_proximity.weight=-value` | +| `params.target_velocity` | .3–1.2;.6 m/s | `NavigationCommandCfg.speed`,**不是奖励权重** | +| `weights.collision_penalty` | -10–-.5;-5 | `obstacle_collision`:同illegal_contact函数/参数,非足端地形接触>10N指示量 | +| `weights.action_smoothness` | -.05–-.001;-.05 | `action_rate_l2.weight` | + +保留原is_terminated等非白名单项,collision_penalty额外独立惩罚非法接触,不用它替代客观碰撞统计。服务写每trial独立reward/task JSON,以`--reward-config/--task-config` argv传入真实训练器和评估器;训练入口先应用场景再应用调参。导出ONNX/deployment的navigation.speed与sensorCfg.avoidanceWeight反映实际配置,浏览器仍81→12且使用配套速度。旧无调参模型保持.6。奖励Preset在普通训练面板仍仅Flat可选;Obstacle可从调参工作台直接导入最佳带metadata策略,不把Obstacle preset混入Flat训练。 + +### 固定客观评估与安全门槛 + +- 固定seed **101/202/303**、每seed **1000策略步×.02s=20s**;每env仅统计reset后的**第一个episode**。提前终止后仍运行完整horizon,但不将自动reset后新episode混入首episode;剩余步数不能带来平滑/净距奖励。各seed先按env等权均值,再三seed等权。evalNumEnvs创建时固定(所有trial相同),Agent不能改seed、horizon、权重、场景、传感器、起终点。 +- 固定评估按三个seed分别构造布局,非随机预设不保证三张不同地图;评估关闭训练期随机起终点。custom严格保持上传的同一地图及验证过的参考spawn/target,标记`fixed-custom-map`,**不是三个随机地形**。保留任务原有reset_joint/startup随机化并用seed复现;若轨迹相同仅是确定性重复试验,不能当作泛化证据。 +- 读取真实checkpoint Actor网络及其`obs_normalizer`统计,strict加载;没有零动作/假策略fallback,不读取训练reward作为指标。不同地图CUDA编译隔离:三个seed顺序启动同一Python解释器子进程(非fork/非并行),共享不可变checkpoint快照并核对SHA256;父进程使用独立临时目录,worker继承父进程组,现有session取消会终止worker,每seed超时1200s即整轮失败。任意seed失败、缺失、错身份、样本不足、非有限或协议不一致整轮fail closed,不用部分结果凑均值。 +- hook在termination manager计算后、`_reset_idx`之前捕获包括**首terminal**的物理量。与mjlab终止判断一致,派生pose最多滞后一个.005s物理子步;接触force history覆盖全部4子步。每步读取目标水平距离、前视射线命中比例、非底板障碍净距、实际action差分、接触/跌倒标志。前视ray_hit可包含地板,仅作为感知诊断,不直接评分。 +- **擦碰/碰撞**:任一机器人(含足端)与非底板障碍的实际接触力幅值>1N,或非足端与任意地形(含地面)>10N。使用编译terrain_1…几何的独立contact sensor(maxforce+4子步history),排除terrain_0标准底板;不是用ray命中推测碰撞。1N以下触碰不计,不能称零接触安全保证。 +- **跌倒**:base高度<.12m或姿态倾角>70°(projected_gravity.z > -cos70°)。**到达**:水平距离<.5m。**成功**:episode内到达,且整段首episode无碰撞、无跌倒;先到达后擦碰/跌倒也失败。未成功时间项必为0,不因提前跌倒取短时间高分。 + +固定五项0–1分量:`success`为成功指示;`time`成功时`1-firstArrivalStep/1000`,否则0;`clearance`每步`clamp((baseXY到非底板box XY AABB距离-.3m)/.5m,0,1)`,按1000步求均值(终止后补0,任何跌倒整项0,无障碍时每有效步1);`smooth`每步`1-clamp(mean((action_t-action_t-1)^2),0,1)`,初始action=0,按1000步求均值(终止后补0);`no_fall`为无跌倒指示。净距是保守圆形机身足迹代理,并非全机身mesh最近距离,地板不作为障碍、躺地不当安全。 + +绝对总分=`.4*success+.2*time+.2*clearance+.1*smooth+.1*no_fall`。例如两步简化fixture,第二步无碰撞到达,净距均.25m、action差分MSE均.5,得分`.4+0+.1+.05+.1=.65`。固定1000步实际协议不接受fixture horizon。 + +相对初始baseline的硬门槛:`fall_rate<=baseline+.02`且`success>=baseline-.02`(边界包含;1e-12浮点容差)。不通过eligible=false且score=-1,不能用其他高分绕过。baseline也用完整客观绝对评分,不默认指标为0。result记录完整协议/场景布局、seed结果、样本数、checkpoint SHA256、各项均值,candidate必须与baseline协议完全一致。 + +## 显式多层射线(阶段5;不是局部高程图) + +`sensorCfg.sensorMode` 是唯一模式字段:缺省/`single_ring_raycast` 保持旧32ray/81obs及旧proximity奖励行为;显式 `multi_ring_raycast` 为3×16=48ray/97obs。`sensorType/type`仍为`raycast`。health新增`sensorModes`,UI可选择;不是任意自定义扫描图案。没有局部高程图网格、用途或输入定义,本阶段**未实现高程图**。 + +任务JSON、deployment JSON及ONNX metadata规范化导出以下字段(FOV仍叫`fov`): + +```json +{ + "sensorMode": "multi_ring_raycast", + "rayCount": 48, + "pitchAngles": [0, -20, -45], + "yawCount": 16, + "yawAngles": [-45, -39, -33, -27, -21, -15, -9, -3, 3, 9, 15, 21, 27, 33, 39, 45], + "fov": 90, + "angleUnit": "deg", + "rayOrder": "layer-major" +} +``` + +旧single规范化为pitch `[0]`、yawCount/rayCount=32。所有yaw为`-fov/2+i*fov/(yawCount-1)`含两端,从右到左,无0度中心ray。只允许两种固定组合;显式矛盾字段/未知字段/不支持的模式拒绝,不静默覆盖。旧metadata缺新增字段可规范化后与新job语义匹配;旧81 graph绝不冒充97。 + +局部direction=`[cos(pitch)*cos(yaw),cos(pitch)*sin(yaw),sin(pitch)]`,先一层16yaw再下一pitch层;origin仍为`basePosition+Rbase*[.3,0,.05]`,direction乘**完整base姿态**。后端实际RayCastSensorCfg使用同一pattern;浏览器预计算局部direction并缓存每个已编译静态box世界变换,解析slab求交;只对精确单位旋转用轴对齐快速路径,无变换近似。 + +97维slice:基础`0:47`完全不变,真实测距`47:95`,目标误差`95:97`仍为原有符号的heading/π与距离归一化。miss/超range→1,地板命中保留真实距离;调参`target_velocity`继续传入真实command与导出navigation.speed。原single-in-flight、held-action、事务导入和资源释放协议不变。三seed评估仍是顺序独立解释器;advisor按实际32/48模式说明盲区。 + +### 地板与奖励 + +仅multi的`obstacle_proximity`排除标准底板顶面:首次计算用CPU编译模型验证**唯一**与首box相符的静态world-weld/group0 box(世界中心`[0,0,-.1]`、半尺寸`[size/2,size/2,.1]`、单位世界旋转),再核实生成器名称`terrain_0`,失败关闭;不是只根据名称。真实编译terrain为固定body而非bodyid=0,使用weld身份。 + +当前mjlab RayCastData没有hit geom ID,因此是**几何分类过滤**:必须finite实际命中,hit世界XY在floor范围(容差1e-5m),`abs(hit.z)<=1e-5m`,normal与+Z逐分量误差<=1e-5。每个Warp环境是独立同坐标地图,hit不额外加env_origin。容差为float32的10微米,不把5cm低障碍抹除。低box顶面、侧面不滤;全floor/miss惩罚0且finite,其余仍用最近有效距离的原平方归一化。观测数据不原地修改。与底板顶面共面的几何无法凭此分类区分,不能声称等价命中geom ID;台阶/非零高度表面不是标准地板。 + +下倾层改善部分低障碍/有界地板边缘感知,但层间、侧后方、遮挡和有限range仍有盲区;标准boxes-v1底板本身会填补内部空洞,不能声称已解决坑探测或安全导航。 + +### 阶段5复现 + +```bash +.venv/bin/python training_server/tests/generate_multi_ring_golden.py +GO2_RUN_MULTI_SMOKE=1 MUJOCO_GL=egl .venv/bin/python -m unittest discover -s training_server/tests -p test_multi_ring.py +WANDB_MODE=disabled MUJOCO_GL=egl .venv/bin/python training_server/rl/scripts/train.py Unitree-Go2-ObstacleAvoidance --env.scene.num-envs=4 --agent.max-iterations=1 --gpu-ids '[0]' --task-config /tmp/go2-multi-ring-stage5/task.json --output-dir /tmp/go2-multi-ring-stage5/train +``` + +真实4env/1iteration新97策略仅验训练、归一化Actor导出和浏览器ORT链路,不是旧81 checkpoint改metadata,不证明导航成功。CLI子进程`--help`回归分别验证Flat/Rough/Obstacle的任务注册;修复前直接CLI仅注册mjlab内建任务,原错误日志保留。 + +事务重导入回归确认97→97、97→81→97合法graph/metadata成功;伪97metadata+81graph与缺失ORT算子均在预期阶段报具体错误并保留旧会话。曾看到“模型编译失败”是App只转抛diagnostic.summary遮蔽了原detail,并非实际MJCF编译失败;已改为detail优先,未改事务资源/rollback。 + +### 复审修复:奖励preset任务归属 + +公共preset仍可保存Obstacle调参结果,但`save/get/list`均从来源session的`config.taskId`恢复权威任务归属并验证完整reward schema;列表和单项返回`taskId`。历史session仅缺`taskId`时按既有Flat默认解释,且必须通过完整Flat schema。来源缺失/损坏、显式未知或空taskId、损坏preset均失败关闭(列表不返回部分可信结果);不相信客户端声明,也不凭奖励字段猜任务。 + +普通本地训练的`rewardPresetId`入口仍仅支持Flat;resolver核对来源任务与请求任务,服务再次验证完整schema,跨任务/损坏preset返回400且不创建job。UI仅展示服务明确标记Flat的preset,旧服务无身份字段时不展示,而不是默认为Flat。Obstacle preset UI未扩展;Obstacle调参最佳策略配套导入、Approval/Automatic均保持原协议。 diff --git a/training_server/PRETRAINED.md b/training_server/PRETRAINED.md new file mode 100644 index 00000000..82f6e112 --- /dev/null +++ b/training_server/PRETRAINED.md @@ -0,0 +1,178 @@ +# Go2 基础策略迁移(普通训练、自调参、CLI) + +普通训练和自调参均可选择经服务验证的基础策略。浏览器持续导航不再因20秒截止而停用;训练/三seed评估协议不变。不新增高程图,不声称短训已学会复杂障碍导航。 + +## 前端直接上传单个文件(推荐) + +已安装训练依赖后,在仓库根目录执行一条默认启动命令: + +```bash +.venv/bin/python training_server/server.py +``` + +1. 打开普通训练面板或自调参工作台,将终端显示的令牌填入访问令牌,连接`http://127.0.0.1:8765`。 +2. 在“基础策略”勾选“确认Go2 legacy47模板”,点击“选择基础策略文件”,选择单个`.pt`或`.onnx`;无需注册JSON、服务器文件路径、YAML或其他sidecars。 +3. 等待上传/验证结束;完成后自动选中该文件的内容ID,检查文件名、格式、原文件SHA及初始化说明。上传中无法启动;可取消等待,同一文件可再次选择重试。 +4. 普通训练选择Flat47或Obstacle81/97,再点击“发起本地训练”;自调参保留逐轮审批/全自动选择,点击“启动自调参”。无DeepSeek配置时必须明确勾选Optuna fallback才能启动自调参;上传本身不调用Agent。 +5. 失败、取消、切换任务或更换连接不会把旧基础策略默默清空;失效/不兼容选择会阻止启动。若要从头训练,明确选择“不选择(随机初始化)”。 + +.pt继承actor及其归一化/探索参数;ONNX只继承确定性推理网络,缺失的count采用合成1,000,000、探索std采用新训练默认1.0,两者critic/optimizer均全新、iteration0。不是恢复原PPO训练。旧管理员注册功能仅为可选高级入口,见下文。 + +### HTTP与安全契约 + +默认启动服务即可上传,**不需要管理员JSON、服务器路径或邻接文件**。旧注册来源仍兼容,见下文“旧CLI/注册来源”。HTTP接口沿用Host/Origin检查及Bearer令牌: + +`POST /api/training/pretrained-sources/upload?format=pt|onnx&template=go2-legacy47-v1&name=` + +- body为原始文件字节,`Content-Type: application/octet-stream`,唯一`Content-Length`。不是JSON/base64/multipart/ZIP;普通JSON请求仍限制128KiB。 +- 必须显式确认模板,不能只凭47维shape自动证明物理语义。模板按Go2 FL/FR/RL/RR各hip/thigh/calf解释关节,以及本文legacy47观测、50Hz、PD/default/action_scale。文件缺失的坐标系/噪声等语义来自用户确认,不声称验证了源env.yaml。 +- `.pt`上限256MiB、ONNX上限64MiB;单个上传/解析槽,接收总限60秒、单次阻塞最多10秒;累计快照最多32个/2GiB(包括旧来源),额外预留16MiB派生产物。临时目录由服务随机生成,显示名去路径/控制字符。接收中断/超时/校验失败清理临时文件,不登记可选来源。崩溃遗留临时目录不列入目录且计入容量,需本机维护者清理。 +- CPU子进程解析:60秒墙钟/40秒CPU、8GiB虚拟地址、16MiB输出文件、64文件描述符、无core dump、CPU单线程。`.pt`只用`weights_only=True`,检查actor完整键/shape/dtype/有限值、normalizer var/std/count、探索std及已知嵌入语义metadata;缺actor的ZIP/任意checkpoint拒绝。不是针对本机同账户恶意写入或原生库漏洞的完整沙箱。 +- 单ONNX只支持opset17/18、自包含float32 `[1,47]→[1,12]` 的精确9节点链:Sub、Div、Gemm/ELU/Gemm/ELU/Gemm/ELU/Gemm。校验每条连边、算子属性、initializer、I/O、关节/观测/PD/action metadata,拒绝外部data/custom op/支路/重排/额外算子;随后CPU ORT对48个随机+物理probe对照重建actor(atol=rtol=2e-5)。不承诺任意ONNX可恢复PPO。 +- `.pt`保留原count和探索std;ONNX只有确定性网络/mean/denominator,以已知模板`epsilon=.01`推导`std=denominator-.01>0`及`var=std²`,**合成count=1,000,000**,探索std使用经检查的目标默认1.0。两者critic/optimizer fresh、iteration0;ONNX source_iteration未知,不能称为恢复原.pt训练。 +- 成功201返回与`health.pretrainedSources[]`同形的单个source record。ID为SHA256(`template:format:原文件SHA256`),文件名不参与;同时保存原文件SHA及服务派生`actor.pt` SHA。manifest明确`sourceFormat`、`derived_fields`、`template_confirmation`与`verified_facts`,不把派生artifact伪装成用户原checkpoint。 +- 后续普通job和tuning session继续只传`pretrainedSourceId`,复用原binding/immutable SHA/CLI初始化协议。受控`actor.pt`加`--pretrained-upload-manifest`加载,不扫描邻接文件;81/97首层新增列为零、统计0/1/1且可学习。rung仍只resume自身checkpoint,保存的已更新count不会重置为1e6。 +- 上传来源及manifest原子目录发布、只读快照;无注册参数重启仍能列出。重复上传返回旧record,不改label/artifact或旧job绑定;失效快照保持显式不可用,不退回随机初始化。不要删除仍被历史session引用的快照。 + +两面板共用PretrainedSourceSelect/LocalTrainingClient及上传控件,`accept=".pt,.onnx"`;成功后合并服务返回的已验证record并选择内容ID,连接/任务epoch变化会取消旧请求并忽略晚到响应,不能把旧endpoint结果写入新连接。上传状态显示上传并验证中(非百分比进度);取消只停止浏览器等待,服务可能已经完成验证,可刷新查看。 + +### 本轮真实CPU证据 + +两个真实文件分别在空临时根以单文件上传成功,均未要求或读取相邻`.pt`/ONNX/YAML。双方重建actor对原ONNX的48组eval最大误差均为`1.9073486e-6`。真实ONNX导入及81/97扩展、非零有限梯度、保存/恢复已更新统计通过;不是PPO收敛或安全行走证明。 + +| 来源 | 更新batch | 更新率 | 物理probe最大动作漂移 | +| ------------------- | -------------: | ---------: | --------------------: | +| pt,count=983138304 | 4096(UI默认) | 4.16623e-6 | 2.62260e-6 | +| pt | 16384(上限) | 1.66647e-5 | 1.04904e-5 | +| onnx,合成count=1e6 | 4096 | .00407929 | .00254631 | +| onnx | 16384 | .01611989 | .01006627 | + +以上为一次正常统计更新的实测,后续PPO可能退化,不能用更新前数值等价保证训练稳定性。 + +双面板×真实.pt/ONNX均经独立默认服务空上传根的浏览器文件选择验证。额外CPU真实runner使用4环境:两来源初始化对MuJoCo实际观测的原ONNX最大误差1.19209e-7;ONNX-derived仅运行1次PPO iteration,新增81维输入列更新为非零,导出对runner误差4.76837e-7,同trial恢复完整actor/optimizer且count保持1000096,未重设为1e6。该检查仅证明链路,不声称导航收敛。 + +## 旧CLI/注册来源使用(仍需配套文件;不适用于单文件上传) + +先激活仓库 `.venv`,安装 `training_server/rl/requirements.txt`。CPU 身份校验依赖 `onnxruntime==1.29.0`。 + +```bash +# SOURCE 是管理员认可的本地训练产物目录,不是浏览器传入的任意路径。 +python training_server/rl/scripts/train.py Unitree-Go2-ObstacleAvoidance \ + --pretrained-checkpoint "$SOURCE/model_10000.pt" \ + --pretrained-allowed-roots "[\"$SOURCE\"]" \ + --agent.logger tensorboard \ + --output-dir "$OUTPUT" +``` + +目录必须有 `params/env.yaml`、`params/agent.yaml`、`policy.onnx` 和明确选定的 `.pt`。可用 `--pretrained-onnx` 指定允许根内的配对 ONNX;没有对应 `.pt` 拒绝,不把 ONNX 当成 PPO resume。**不扫描/按 mtime 自动认定 model_10000 对应 ONNX**;当前实际对应关系由 CPU ORT 数值验证建立。 + +- `--pretrained-checkpoint` 与 `--resume-checkpoint` / `--agent.resume` 互斥。 +- warm-start:新 trial 的 iteration=0;只迁移 actor,critic 和 optimizer 保持新初始化。预算仍是新训练预算。 +- resume:既有 runner.load 和 explicit-resume 的 `current+1`、`max_iterations-current` 行为完全不变;同 trial successive-halving 不应再次传 pretrained 参数。 +- 原 Flat legacy47 支持迁移;81/97 支持扩展。Rough 普通训练/原 resume 不受影响,但本轮**不支持从 legacy47 warm-start 到含 height scan 的 Rough**,会明确拒绝。 +- 校验发生在模拟器分配前;创建环境后另核对编译后的 joint 顺序、PD、默认位姿、观测顺序与动作 scale。成功初始化时在输出目录写 `initialization.json`。 + +## 严格迁移契约 + +源 actor 必须是 RSL `MLPModel`:47→512→256→128→12、ELU、经验归一化、Gaussian scalar std。所有 state key、dtype、shape、有限值及 std/var 一致性均验证;没有 `strict=False`。观测语义顺序为: + +`base_ang_vel(3), projected_gravity(3), command(3), phase(2), joint_pos(12), joint_vel(12), actions(12)`。 + +校验 term 函数名、参数、噪声/缩放/裁剪/延迟/history、动作配置、机器人 spec 名、默认关节状态、actuator/armature/collision 配置、decimation=4 与 dt=.005。观测 corruption 开关和命令采样分布可以因导航任务改变,但基础观测的坐标/单位/顺序不变。phase 为 .6s sin/cos,命令 norm<.1 时清零;本地与源 phase 实现只差尾空行,robot XML 逐字节相同(本次只读比对)。并非宣称训练分布和导航地形完全相同。 + +迁移覆盖: + +| 参数 | 行为 | +| ---------------------------------------------------- | ----------------------------------------------- | +| `mlp.0.weight[:, :47]` | 原值复制 | +| `mlp.0.weight[:, 47:]` | 零初始化,仍 requires_grad,可由 PPO 更新 | +| 所有 MLP bias、其余 weight、`distribution.std_param` | 原值复制 | +| normalizer mean/var/std 前47维及标量 count | 逐值复制 | +| normalizer 新维 | mean=0、var=1、std=1 | +| 源 critic/optimizer/iter | 不加载,保持全新目标 critic/optimizer/iteration | + +归一化策略经批准为 **`preserve-source-count/unit-new-features`**。真实源 count=983138304,RSL 更新率为 batch_size / 更新后 count。清 count 会让首批覆盖旧47统计,故禁止清零。新 ray∈[0,1]、heading∈[-1,1]、distance∈[0,1] 原本有界;使用近 identity 初始化(实际除以 std+.01),继承原算法缓慢更新,**不是完全冻结**。不引入 split normalizer、不改变 checkpoint 格式。 + +初始化等价在 eval / 未更新统计时验证。统计更新后不是要求动作数学上不变,而是必须与原47 actor执行相同统计更新后的动作一致,且变化无突变。48组含分布外随机 gravity 的探针实测最大变化9.06e-5;真实4env首batch变化0。 + +## 可选旧管理员注册来源并在两种面板使用 + +注册JSON由管理员在本机配置,不是HTTP请求。以下直接使用本次用户提供的源目录(只读,不改原文件): + +```bash +source .venv/bin/activate +export SOURCE=/home/cen/Embodied_Workspace/unitree_rl_mjlab/logs/rsl_rl/go2_velocity/2026-08-25_Go2 +export SOURCE_CONFIG=/tmp/go2-pretrained-sources.json +python - <<'PYCONFIG' +import json, os +from pathlib import Path +root = Path(os.environ["SOURCE"]) +Path(os.environ["SOURCE_CONFIG"]).write_text(json.dumps({ + "allowedRoots": [str(root)], + "sources": [{"id": "go2-base", "label": "Go2已训练基础行走策略", + "checkpoint": str(root / "model_10000.pt"), + "onnx": str(root / "policy.onnx")}] +}, ensure_ascii=False, indent=2)) +PYCONFIG +python training_server/server.py --trainer-python "$VIRTUAL_ENV/bin/python" \ + --pretrained-sources "$SOURCE_CONFIG" \ + --tuning-data-root "$HOME/.local/share/mujoco-go2-training" +``` + +1. 复制服务显示的访问令牌,在普通训练面板连接服务,或在自调参工作台连接同一服务。 +2. 在“基础策略”选择已验证条目,检查`model_10000.pt`及checkpoint/ONNX完整SHA、兼容任务和观测维度。两面板默认不选,保留原随机训练行为。 +3. 选择Flat47或Obstacle81/97;Rough显示不兼容并由服务拒绝。不兼容/失效来源须显式处理,不可静默清空后退回随机初始化。 +4. 普通训练点击“发起本地训练”;自调参选择Approval/Automatic后启动。基础策略不是LLM可调字段;每个新trial(含baseline)使用同一快照与相同session seed的新critic/optimizer初始化协议。后续rung只加载该trial自己的checkpoint,不覆盖成原基础权重。 + +允许根只能来自管理员JSON,不能来自浏览器;最多32注册条目。服务对路径各级使用dir_fd/O_NOFOLLOW,拒绝symlink、穿越、FIFO、目录与超限文件(pt256MiB/ONNX64MiB/env2MiB/agent128KiB)。CPU子进程只使用受限YAML与`torch.load(weights_only=True)`验证形状、语义和ONNX数值身份。未知object/apply、重复键、不匹配配置或缺.pt均拒绝;不逆向ONNX恢复PPO,也不按mtime猜checkpoint。 + +验证后原始四件产物复制到`/pretrained_sources//`内容寻址只读快照。浏览器提交的`pretrainedSourceId`是目录返回的内容SHA身份,不是管理员可读别名或文件路径;同一别名重新注册了不同文件时,旧UI的内容ID会拒绝,必须刷新。注册之后原文件变化不改变已有job/session;每次启动trial都重新核验快照SHA,CLI还核验期望source_id。勿删除有历史session依赖的快照。管理员仍须保护数据根;这是常规文件/反序列化防护,不是任意ONNX计算或同账户恶意写入的完整沙箱。 + +session的SQLite config持久化完整无路径来源描述,旧session无该字段仍走原默认。服务重启会将queued/running/evaluating/paused/awaiting_approval标记interrupted并记录原state,不自动训练或调用Agent;必须用户显式恢复。恢复只清理不完整trial,保留完整评估/来源/约束版本/mode。遗留pending proposal以`service_restart/recovery_invalidated`审计拒绝(非用户拒绝),重新提案后Approval仍需审批;已批准历史不改。来源快照丢失/损坏时恢复排队前拒绝。连续resume请求不会启动两个worker。 + +源manifest/迁移覆盖写入`initialization.json`;source-bound rung要求父checkpoint的来源记录匹配。最佳ONNX的`pretrained_initialization`metadata仅含SHA/source ID、源迭代、normalizer协议与initialIteration=0,不含本地路径或任意checkpoint infos。 + +## 持续点击导航 + +仅Obstacle81/97启用配套地图点击导航,不给普通Flat47增加非同构导航地图模式。浏览器没有20秒交互截止,但跌倒/越界仍安全停止;需用户重置并重新启用。设定新目标不重新加载策略、不改变单inflight/held-action机制,也不会自动启用已停止策略。面板明确显示策略未启用、仿真暂停或安全停止。训练与固定三seed/1000步(20秒)评估不变。 + +真实源warm-start81在WASM/ORT中21.02秒仍启用:从x=-5走到2.566;换目标使command从vx=.598/yaw=-.077变为vx=0/yaw=1。25.04秒仍启用,距新目标从2.828m降到2.243m;无teleport、脚本轨迹或constant actor。此为空旷plane单例,不证明复杂障碍泛化或导航收敛。 + +## 验证 + +```bash +python -m unittest discover -s training_server/tests -p test_pretrained.py -v +# 可选真实源;不硬编码用户个人目录到源码: +GO2_PRETRAINED_SOURCE="$SOURCE" \ +GO2_PRETRAINED_EVIDENCE_DIR="$EVIDENCE" \ +python -m unittest discover -s training_server/tests -p test_pretrained.py -v +# 可选仅4env×1iteration GPU初始化/优化/导出/重载烟测: +GO2_PRETRAINED_GPU_SMOKE=1 MUJOCO_GL=egl \ +GO2_PRETRAINED_SOURCE="$SOURCE" GO2_PRETRAINED_EVIDENCE_DIR="$EVIDENCE" \ +python -m unittest discover -s training_server/tests -p test_pretrained.py -v +``` + +本次真实源:checkpoint SHA `d94eebd23be8b3cc998b493fe056a9e60ceaf4705900320015f30387412a4ffb`;ONNX SHA `80150119e93ce2f656625fc3048ece43ec1281fcdfe9109e07a0157438d0df7c`。 + +- 源 ORT vs checkpoint 48 probe 最大误差 `1.9073486328125e-6`。 +- 47→47/81/97、不同ray/goal:CPU初始化最大误差0;扩展81/97的CPU ONNX导出最大误差 `1.9073486328125e-6`。 +- 真实4env源ONNX vs 扩展actor最大误差 `3.427267074584961e-7`;首batch基础mean变化 `4.190951585769653e-9`,动作变化0。 +- 新列梯度最大 `.0273208`;单iteration PPO后新列weight最大 `.00291905`,critic/optimizer得到更新。 +- 新runner加载保存checkpoint后全部actor张量精确恢复,count=983138404;导出ONNX最大误差 `4.76837158203125e-7`。 +- 安全失败例:越界路径/逃逸symlink、unsafe YAML、未知观测语义/phase/scale/armature、坏shape/NaN/std、错误actor身份、warmstart+resume混用。 + +这些是初始化和优化步骤证据,**不等于导航收敛/到达率验证**。完整命令日志和源manifest在本次交接的 `/tmp/go2-pretrained-core-*/` 工件目录。 + +### 浏览器实产物复核(不再次训练) + +本次保留的真实warm-start产物为 `/tmp/go2-pretrained-integration-Yqhm9K/train/policy.onnx`,source SHA及新迭代0可从同目录`initialization.json`核对。下列测试需Vite开发服务(动态导入真实PhysicsAdapter);另一个终端启动后运行: + +```bash +npm run dev -- --host 127.0.0.1 --port 4174 +# 在另一个终端: +GO2_PRETRAINED_DEV_URL=http://127.0.0.1:4174 \ +GO2_PRETRAINED_NAV_POLICY=/tmp/go2-pretrained-integration-Yqhm9K/train/policy.onnx \ +npx playwright test -c web_platform/playwright.config.ts web_platform/e2e/pretrainedNavigation.spec.ts +``` + +原始47维Flat UI兼容测试还可设置`GO2_PRETRAINED_FLAT_POLICY`和`GO2_PRETRAINED_FLAT_XML`。后者必须是带12个有效actuator、采用既有`FL_hip`等绑定名称的Go2模型;裸训练XML本来不含actuator,会明确拒绝。此次`/tmp/go2-pretrained-integration-Yqhm9K/go2-flat.xml`由本地`Entity(get_go2_robot_cfg())`导出,保留源PD/armature,仅把actuator标识改为`actuator.target.removesuffix('_joint')`以匹配既有浏览器命名;没有修改生产模型或为Flat新增点击导航。Flat此测试只证明加载/推理兼容;持续行走/换目标证据来自上述真实81维warm-start产物。 diff --git a/training_server/README.md b/training_server/README.md index 377721ac..98a0a936 100644 --- a/training_server/README.md +++ b/training_server/README.md @@ -4,6 +4,10 @@ 仓库已在 [`rl/`](rl/) 内置 `Unitree-Go2-Flat` 所需的 PPO 训练代码、Go2 模型资产和 ONNX 导出逻辑,不再要求另外克隆 `unitree_rl_mjlab`。`mjlab`、PyTorch 等大型运行依赖仍需安装在本机训练环境中。 +## 自定义任务与地形 + +内置新增 `Unitree-Go2-ObstacleAvoidance`(81维前视射线导航),并放行 `Unitree-Go2-Rough` 训练。健康接口提供可配置参数元数据,job 的 `deployment` 返回精确地图布局、传感器及策略契约。首版支持 `plane/discrete_obstacles/rough/pyramid_stairs/wave` 的训练专用box布局;不是任意场景导入,也不实现真实深度相机。旧 Rough 的234维actor不能在当前浏览器一键部署。完整字段、坐标系、观测动作及复现方式见 [避障部署契约](OBSTACLE_AVOIDANCE.md)。 + ## 准备训练环境 使用仓库已有的 `.venv` 项目虚拟环境: @@ -19,8 +23,7 @@ python -m pip install -r training_server/requirements.txt 使用已安装训练依赖的 Python 解释器启动服务: ```bash -npm run training-server -- \ - --trainer-python "$PWD/.venv/bin/python" +.venv/bin/python training_server/server.py ``` 服务启动时会在终端显示一个随机访问令牌。将该令牌填入前端“访问令牌”字段后再连接。令牌只保存在当前浏览器标签页的 `sessionStorage` 中。自动化启动时可固定令牌: @@ -39,7 +42,15 @@ export MUJOCO_TUNING_AGENT_MODEL='deepseek-v4-flash' # 可省略 npm run training-server -- --trainer-python "$PWD/.venv/bin/python" ``` -普通训练不要求 DeepSeek key。未配置时健康接口会把 tuning 标记为不可用;只有创建 session 时明确勾选 fallback,才允许 Agent 失败后使用 Optuna 候选,不会静默降级。 +普通训练与基础策略上传不要求 DeepSeek key。未配置时健康接口会把 tuning 标记为不可用;只有创建 session 时明确勾选 fallback,才允许使用 Optuna 候选,不会静默降级。 + +### 从浏览器选择已有策略(推荐) + +普通训练和自调参均可在“基础策略”直接选择单个`.pt`或`.onnx`文件:连接默认服务 → 勾选Go2 legacy47观测/关节语义确认 → 点击“选择基础策略文件” → 等待验证成功后自动选择内容ID并检查文件名/格式/原文件SHA → 发起训练或自调参。不需要`--pretrained-sources`、管理员JSON、服务器路径或sidecars;旧管理员注册是可选高级功能。 + +上传期间禁止启动;取消、验证失败、任务/连接改变会保留旧选择,旧请求结果不会污染新连接。相同文件可重选重试,失效来源不会静默退回随机;从头训练须明确选择“不选择(随机初始化)”。支持47维Go2 actor扩展到81/97新输入,不接受未知机器人/任意ONNX/97维源checkpoint。 + +**只继承actor,不是完整PPO resume**:两格式critic/optimizer全新、iteration0。`.pt`保留原normalizer count和探索std;ONNX仅继承确定性网络,count合成1,000,000、std/var按受限模板推导、探索std使用目标默认1.0,源迭代未知。详细上限、安全边界及点击步骤见[基础策略上传与迁移](PRETRAINED.md)。 默认训练工程是仓库内的 `training_server/rl`。如需使用包含其他已注册任务的外部训练工程,仍可通过 `--trainer-root /path/to/trainer` 或 `UNITREE_RL_MJLAB_ROOT` 覆盖。默认端口是 `8765`。如果前端不是从 `localhost` 或 `127.0.0.1` 提供,可显式添加来源: @@ -64,6 +75,7 @@ python training_server/server.py \ ## 接口 - `GET /api/training/health`:运行环境、允许的任务和活动任务; +- `POST /api/training/pretrained-sources/upload?format=pt|onnx&template=go2-legacy47-v1&name=显示名`:认证有界二进制单文件上传,返回持久化内容ID; - `POST /api/training/jobs`:发起训练; - `GET /api/training/jobs/{id}`:状态、迭代进度和最近日志; - `DELETE /api/training/jobs/{id}`:停止训练; @@ -80,7 +92,7 @@ python training_server/server.py \ - `GET /api/tuning/sessions/{id}/artifacts/best/policy.onnx`:下载最佳策略; - `GET /api/tuning/presets`:列出可供普通训练复用的最佳奖励 preset。 -普通任务状态在服务重启后丢失,但日志、checkpoint 和 ONNX 保留在 `training_server/rl/logs/rsl_rl/`;调参状态、调度令牌、参数护栏、回滚基准及产物持久化在 `logs/auto_tuning/`。候选配置数可在 1–100 间设置(包含基线配置,仍受连续无提升早停约束)。API 只接收 32 位资源 ID,不接收客户端文件路径;奖励 patch 受到名称、符号、上下界、Session 护栏、每轮最多 4 项及 `0.5×–2×` 变化率校验。 +普通任务状态在服务重启后丢失,但日志、checkpoint 和 ONNX 保留在 `training_server/rl/logs/rsl_rl/`;调参状态、调度令牌、参数护栏、回滚基准及产物持久化在 `logs/auto_tuning/`。候选配置数可在 1–100 间设置(包含基线配置,仍受连续无提升早停约束)。作业/session API 接收32位资源ID,基础策略选择使用64位内容ID;上传仅接收文件字节,不接收客户端服务器文件路径;奖励 patch 受到名称、符号、上下界、Session 护栏、每轮最多 4 项及 `0.5×–2×` 变化率校验。 ## 测试 @@ -90,3 +102,7 @@ python -m pip install -r requirements-dev.txt npm run lint:python npm run test:training-server ``` + +## 使用已训练基础策略 + +普通训练与自调参面板均提供“基础策略”选择。管理员通过 `--pretrained-sources` 注册只读本地ONNX及数值匹配的.pt;服务创建SHA绑定快照,浏览器只能选择内容ID,不能提交文件路径。完整可直接运行的注册/启动命令、兼容47/81/97范围、重启恢复规则与验证方法见 [基础策略迁移](PRETRAINED.md)。缺.pt、Rough不匹配或快照失效均明确拒绝,绝不静默随机初始化。 diff --git a/training_server/pretrained.py b/training_server/pretrained.py new file mode 100644 index 00000000..85fae5c3 --- /dev/null +++ b/training_server/pretrained.py @@ -0,0 +1,502 @@ +"""Validated Go2 velocity actor warm-start; never an optimizer/iteration resume. + +Callers supply administrator-owned allowed roots, not roots from an HTTP request. +Legacy YAML tags are decoded as inert data: no import, eval or object construction. +""" + +import hashlib +import io +import json +from dataclasses import asdict, dataclass +from enum import Enum +from importlib.metadata import version +from pathlib import Path + +import torch +import yaml + +BASE_TERMS = [ + "base_ang_vel", + "projected_gravity", + "command", + "phase", + "joint_pos", + "joint_vel", + "actions", +] +JOINTS = [ + f"{leg}_{joint}_joint" for leg in ("FL", "FR", "RL", "RR") for joint in ("hip", "thigh", "calf") +] +NORMALIZATION_POLICY = "preserve-source-count/unit-new-features" + + +class PretrainedError(ValueError): + """The supplied source does not prove the supported transfer contract.""" + + +class _DataLoader(yaml.SafeLoader): + def construct_mapping(self, node, deep=False): + keys = [self.construct_object(key, deep=deep) for key, _ in node.value] + if any(type(key) not in (str, int) for key in keys) or len(set(keys)) != len(keys): + raise PretrainedError("Configuration mappings require unique string/integer keys") + return super().construct_mapping(node, deep=deep) + + +def _symbol(loader, suffix, node): + if loader.construct_scalar(node) != "": + raise PretrainedError("Python name tags must have an empty value") + return {"symbol": suffix} + + +def _enum(loader, node): + values = loader.construct_sequence(node) + if len(values) != 1 or not isinstance(values[0], (str, int)): + raise PretrainedError("Invalid legacy enum data") + return values[0] + + +_DataLoader.add_constructor("tag:yaml.org,2002:python/tuple", _DataLoader.construct_yaml_seq) +_DataLoader.add_multi_constructor("tag:yaml.org,2002:python/name:", _symbol) +for _name in ("mjlab.actuator.actuator.TransmissionType", "mjlab.viewer.viewer_config.OriginType"): + _DataLoader.add_constructor("tag:yaml.org,2002:python/object/apply:" + _name, _enum) +_DataLoader.add_constructor( + "tag:yaml.org,2002:python/object/apply:builtins.slice", + lambda loader, node: {"slice": loader.construct_sequence(node)}, +) + + +def _plain(value): + if isinstance(value, Enum): + return value.value + if callable(value): + return {"symbol": value.__module__ + "." + value.__qualname__} + if isinstance(value, dict): + return {k: _plain(v) for k, v in value.items()} + if isinstance(value, (list, tuple)): + return [_plain(v) for v in value] + return value + + +def _require(ok, message): + if not ok: + raise PretrainedError(message) + + +def _read_allowed(path, roots, limit): + path = Path(path).expanduser().resolve(strict=True) + _require( + any(path.is_relative_to(root) for root in roots), + "Source is outside configured allowed roots", + ) + _require(path.is_file() and path.stat().st_size <= limit, "Source file missing or oversized") + # Hash and deserialize the same bytes, not a second pathname lookup. + with path.open("rb") as stream: + data = stream.read(limit + 1) + _require(len(data) <= limit, "Source file oversized") + return path, data + + +def _yaml_data(data): + try: + result = yaml.load(data, Loader=_DataLoader) + _require(isinstance(result, dict), "Expected YAML mapping") + return result + except yaml.YAMLError as exc: + raise PretrainedError("Unsupported or invalid configuration YAML") from exc + + +def _actor_config(config): + actor = dict(config["actor"]) + distribution = dict(actor["distribution_cfg"]) + # Old RSL config omitted this default; the tensor contract is checked too. + distribution.setdefault("class_name", "GaussianDistribution") + actor["distribution_cfg"] = distribution + actor.setdefault("class_name", "MLPModel") + actor.setdefault("cnn_cfg", None) + return actor + + +def validate_semantics(source_env, source_agent, target_env, target_agent): + """Compare source and target against the repository's supported legacy47 contract.""" + 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 + + reference = _plain(asdict(unitree_go2_flat_env_cfg())) + agent = _plain(asdict(unitree_go2_ppo_runner_cfg())) + target_env, target_agent = _plain(target_env), _plain(target_agent) + try: + for label, env in (("source", source_env), ("target", target_env)): + group = env["observations"]["actor"] + names = list(group["terms"]) + _require(names[:7] == BASE_TERMS, f"{label}: unknown base observation order") + _require( + names == BASE_TERMS + if label == "source" + else names in (BASE_TERMS, BASE_TERMS + ["forward_depth", "target_error"]), + f"{label}: unsupported observation suffix", + ) + for name in BASE_TERMS: + _require( + group["terms"][name] == reference["observations"]["actor"]["terms"][name], + f"{label}: incompatible observation {name}", + ) + for key in ( + "concatenate_terms", + "concatenate_dim", + "history_length", + "flatten_history_dim", + ): + _require( + group[key] == reference["observations"]["actor"][key], + f"{label}: incompatible observation {key}", + ) + _require( + env["actions"] == reference["actions"], f"{label}: incompatible action transform" + ) + robot, ref_robot = ( + env["scene"]["entities"]["robot"], + reference["scene"]["entities"]["robot"], + ) + for key in ("spec_fn", "articulation", "sort_actuators", "collisions"): + _require(robot[key] == ref_robot[key], f"{label}: incompatible robot {key}") + for key in ("joint_pos", "joint_vel"): + _require( + robot["init_state"][key] == ref_robot["init_state"][key], + f"{label}: incompatible default {key}", + ) + _require( + env["decimation"] == 4 and env["sim"]["mujoco"]["timestep"] == 0.005, + f"{label}: policy must run at 50 Hz", + ) + if len(target_env["observations"]["actor"]["terms"]) > 7: + from src.tasks.obstacle_avoidance.env_cfg import unitree_go2_obstacle_env_cfg + + suffix = _plain(asdict(unitree_go2_obstacle_env_cfg()))["observations"]["actor"][ + "terms" + ] + terms = target_env["observations"]["actor"]["terms"] + depth = dict(terms["forward_depth"]) + _require( + set(depth["params"]) == {"max_distance"} + and 0 < depth["params"]["max_distance"] < float("inf"), + "Invalid ray normalization", + ) + depth["params"] = suffix["forward_depth"]["params"] + _require( + depth == suffix["forward_depth"] + and terms["target_error"] == suffix["target_error"], + "Unknown ray/goal observation semantics", + ) + for label, cfg in (("source", source_agent), ("target", target_agent)): + _require( + _actor_config(cfg) == _actor_config(agent), + f"{label}: incompatible actor architecture/activation/std", + ) + _require( + cfg["obs_groups"]["actor"] == ["actor"] and cfg["clip_actions"] is None, + f"{label}: incompatible actor groups/action clipping", + ) + except (KeyError, TypeError) as exc: + raise PretrainedError(f"Missing or invalid semantic configuration: {exc}") from exc + + +def make_reference_actor(dim=47): + from rsl_rl.models import MLPModel + from tensordict import TensorDict + + return MLPModel( + TensorDict({"actor": torch.zeros(1, dim)}, batch_size=[1]), + {"actor": ["actor"]}, + "actor", + 12, + hidden_dims=(512, 256, 128), + activation="elu", + obs_normalization=True, + distribution_cfg={ + "class_name": "GaussianDistribution", + "init_std": 1.0, + "std_type": "scalar", + }, + ) + + +def validate_actor_state(state, dim=47): + expected = make_reference_actor(dim).state_dict() + _require( + isinstance(state, dict) and state.keys() == expected.keys(), "Unexpected actor state keys" + ) + for key, tensor in state.items(): + _require( + isinstance(tensor, torch.Tensor) + and tensor.shape == expected[key].shape + and tensor.dtype == expected[key].dtype, + f"Incompatible actor tensor: {key}", + ) + _require(bool(torch.isfinite(tensor).all()), f"Nonfinite actor tensor: {key}") + _require(state["obs_normalizer.count"].item() >= 0, "Negative normalizer count") + _require( + bool((state["obs_normalizer._var"] >= 0).all()) + and bool((state["obs_normalizer._std"] >= 0).all()), + "Negative normalization variance/std", + ) + _require( + torch.allclose( + state["obs_normalizer._std"].square(), + state["obs_normalizer._var"], + atol=1e-5, + rtol=1e-5, + ), + "Inconsistent normalization variance/std", + ) + _require(bool((state["distribution.std_param"] > 0).all()), "Nonpositive exploration std") + + +def comparison_observations(): + """Deterministic random + physically plausible standing/walking probe vectors.""" + g = torch.Generator().manual_seed(20260825) + random = torch.randn(32, 47, generator=g) + physical = torch.zeros(16, 47) + physical[:, 5] = -1 # body-frame projected gravity + physical[:, 6] = torch.linspace(0, 1, 16) + phase = torch.linspace(0, 2 * torch.pi, 16) + physical[:, 9], physical[:, 10] = phase.sin(), phase.cos() + physical[0, 9:11] = 0 # standing phase is masked + return torch.cat((random, physical)) + + +def verify_onnx(onnx_bytes, actor): + import numpy as np + import onnx + import onnxruntime as ort + + graph = onnx.load_model_from_string(onnx_bytes) + _require( + all(t.data_location != onnx.TensorProto.EXTERNAL for t in graph.graph.initializer), + "External ONNX tensors are not allowed", + ) + options = ort.SessionOptions() + options.intra_op_num_threads = 1 + options.inter_op_num_threads = 1 + session = ort.InferenceSession( + onnx_bytes, sess_options=options, providers=["CPUExecutionProvider"] + ) + inputs, outputs = session.get_inputs(), session.get_outputs() + _require( + len(inputs) == len(outputs) == 1 + and inputs[0].shape == [1, 47] + and outputs[0].shape == [1, 12] + and inputs[0].type == outputs[0].type == "tensor(float)", + "ONNX must be float32 [1,47] -> [1,12]", + ) + metadata = session.get_modelmeta().custom_metadata_map + _require( + metadata.get("joint_names", "").split(",") == JOINTS + and metadata.get("observation_names", "").split(",") == BASE_TERMS + and metadata.get("command_names") == "twist", + "ONNX joint/observation/command semantics mismatch", + ) + for key, expected in { + "action_scale": [0.25], + "joint_stiffness": [20, 20, 40] * 4, + "joint_damping": [1, 1, 2] * 4, + "default_joint_pos": [-0.1, 0.9, -1.8, 0.1, 0.9, -1.8] * 2, + }.items(): + try: + actual = [float(v) for v in metadata[key].split(",")] + _require( + len(actual) == len(expected) and np.allclose(actual, expected, rtol=0, atol=1e-6), + f"ONNX {key} mismatch", + ) + except (KeyError, ValueError) as exc: + raise PretrainedError(f"Invalid ONNX {key}") from exc + actor.eval() + obs = comparison_observations() + with torch.inference_mode(): + expected = actor.mlp(actor.obs_normalizer(obs)).numpy() + actual = np.concatenate( + [session.run(None, {inputs[0].name: row[None].numpy()})[0] for row in obs] + ) + _require( + np.isfinite(actual).all() and np.allclose(actual, expected, atol=2e-5, rtol=2e-5), + "ONNX does not match the selected checkpoint actor", + ) + return { + "provider": "CPUExecutionProvider", + "probe_count": len(obs), + "max_abs_error": float(np.max(np.abs(actual - expected))), + "atol": 2e-5, + "rtol": 2e-5, + } + + +@dataclass +class ValidatedSource: + actor_state: dict + manifest: dict + + +def read_pretrained_source(checkpoint, *, allowed_roots, target_env, target_agent, onnx_path=None): + """Read an explicit .pt plus paired ONNX/config. Never guess checkpoint by mtime. + + Future services should resolve opaque source IDs to these server-owned paths. + An ONNX selection without an explicit corresponding .pt is not a training source. + """ + roots = [Path(p).expanduser().resolve(strict=True) for p in allowed_roots] + _require( + bool(roots) and all(p.is_dir() for p in roots), + "Configure at least one local allowed source root", + ) + checkpoint, checkpoint_bytes = _read_allowed(checkpoint, roots, 256 * 1024 * 1024) + _require( + checkpoint.suffix == ".pt", + "Warm-start requires a corresponding .pt checkpoint; ONNX cannot resume PPO", + ) + files = {"checkpoint": (checkpoint, checkpoint_bytes)} + for name, path, limit in ( + ("onnx", onnx_path or checkpoint.parent / "policy.onnx", 64 * 1024 * 1024), + ("env", checkpoint.parent / "params/env.yaml", 2 * 1024 * 1024), + ("agent", checkpoint.parent / "params/agent.yaml", 128 * 1024), + ): + files[name] = _read_allowed(path, roots, limit) + source_env, source_agent = _yaml_data(files["env"][1]), _yaml_data(files["agent"][1]) + validate_semantics(source_env, source_agent, target_env, target_agent) + try: + state = torch.load(io.BytesIO(checkpoint_bytes), map_location="cpu", weights_only=True) + _require( + isinstance(state, dict) and isinstance(state.get("iter"), int), + "Invalid checkpoint iteration", + ) + actor_state = state["actor_state_dict"] + validate_actor_state(actor_state) + except (KeyError, RuntimeError) as exc: + raise PretrainedError("Unsupported checkpoint") from exc + actor = make_reference_actor() + actor.load_state_dict(actor_state, strict=True) + identity = verify_onnx(files["onnx"][1], actor) + artifacts = { + name: {"name": path.name, "sha256": hashlib.sha256(data).hexdigest(), "bytes": len(data)} + for name, (path, data) in files.items() + } + manifest = { + "schema_version": 1, + "mode": "pretrained-warm-start", + "contract": "go2-legacy47-v1", + "artifacts": artifacts, + "source_iteration": state["iter"], + "source_actor_dim": 47, + "source_normalizer_count": actor_state["obs_normalizer.count"].item(), + "verification_versions": { + name: version(name) + for name in ("torch", "rsl-rl-lib", "mjlab", "onnxruntime", "PyYAML") + }, + "normalization": NORMALIZATION_POLICY, + "onnx_identity": identity, + "critic": "fresh-target-initialization", + "optimizer": "fresh", + "iteration": 0, + "base_observation_terms": BASE_TERMS, + "joint_names": JOINTS, + } + manifest["source_id"] = hashlib.sha256( + json.dumps(artifacts, sort_keys=True).encode() + ).hexdigest() + return ValidatedSource(actor_state, manifest) + + +def warm_start_actor(actor, source): + """Strictly transfer a validated legacy actor into a fresh 47/81/97 actor.""" + dim = actor.obs_dim + _require(dim in (47, 81, 97), "Only legacy47 and ray/goal 81/97 actors are supported") + reference = make_reference_actor(dim) + _require( + type(actor) is type(reference) + and repr(actor.mlp) == repr(reference.mlp) + and type(actor.distribution) is type(reference.distribution) + and actor.distribution.std_type == "scalar" + and actor.obs_normalization + and type(actor.obs_normalizer) is type(reference.obs_normalizer) + and actor.obs_normalizer.eps == reference.obs_normalizer.eps + and actor.obs_normalizer.until is None + and list(actor.obs_groups) == ["actor"], + "Target actor runtime architecture mismatch", + ) + validate_actor_state(actor.state_dict(), dim) + validate_actor_state(source.actor_state) + migrated = {} + for key, tensor in source.actor_state.items(): + if key == "mlp.0.weight": + value = tensor.new_zeros((512, dim)) + value[:, :47] = tensor + elif key in ("obs_normalizer._mean", "obs_normalizer._var", "obs_normalizer._std"): + value = tensor.new_full((1, dim), 0 if key.endswith("_mean") else 1) + value[:, :47] = tensor + else: + value = tensor.clone() + migrated[key] = value + actor.load_state_dict(migrated, strict=True) + return { + **source.manifest, + "target_actor_dim": dim, + "copied_tensors": list(source.actor_state), + "zero_initialized_input_columns": [47, dim], + "new_feature_statistics": {"mean": 0, "var": 1, "std": 1}, + } + + +def validate_runtime_contract(env): + """Confirm compiled joint order/PD/defaults, not only configuration names.""" + from mjlab.rl.exporter_utils import get_base_metadata + + metadata = get_base_metadata(env, "pretrained-validation") + _require(metadata["joint_names"] == JOINTS, "Compiled robot joint order mismatch") + _require( + list(metadata["observation_names"]) + in (BASE_TERMS, BASE_TERMS + ["forward_depth", "target_error"]), + "Compiled observation order mismatch", + ) + _require(metadata["command_names"] == ["twist"], "Compiled command mismatch") + for key, expected in { + "action_scale": [0.25], + "joint_stiffness": [20, 20, 40] * 4, + "joint_damping": [1, 1, 2] * 4, + "default_joint_pos": [-0.1, 0.9, -1.8, 0.1, 0.9, -1.8] * 2, + }.items(): + actual = torch.as_tensor(metadata[key], dtype=torch.float64).reshape(-1) + target = torch.tensor(expected, dtype=torch.float64) + _require( + actual.shape == target.shape and torch.allclose(actual, target, atol=1e-6, rtol=0), + f"Compiled {key} mismatch", + ) + return metadata + + +def public_initialization_metadata(initialization): + """Export provenance without local paths, including original single-file identity.""" + artifacts = initialization["artifacts"] + result = { + "sourceId": initialization["source_id"], + "checkpointSha256": artifacts["checkpoint"]["sha256"], + "normalization": initialization["normalization"], + "sourceIteration": initialization["source_iteration"], + "initialIteration": 0, + } + if "sourceFormat" in initialization: + result.update( + sourceFormat=initialization["sourceFormat"], + uploadSha256=artifacts["upload"]["sha256"], + templateId=initialization["contract"], + derivedFields=initialization["derived_fields"], + templateConfirmation=initialization["template_confirmation"], + ) + else: + result["onnxSha256"] = artifacts["onnx"]["sha256"] + return result + + +def initialize_runner(runner, source): + _require( + runner.current_learning_iteration == 0 and not runner.alg.optimizer.state, + "Warm-start requires a fresh runner, never a resumed trial", + ) + # Do not load critic/optimizer/iteration from the source, even if shapes match. + return warm_start_actor(runner.alg.actor, source) diff --git a/training_server/pretrained_sources.py b/training_server/pretrained_sources.py new file mode 100644 index 00000000..9ff3c527 --- /dev/null +++ b/training_server/pretrained_sources.py @@ -0,0 +1,447 @@ +"""Administrator-registered, content-bound local training sources (no client paths).""" + +from __future__ import annotations + +import hashlib +import json +import os +import re +import shutil +import stat +import subprocess +import tempfile +import threading +import time +from contextlib import contextmanager +from copy import deepcopy +from pathlib import Path + +SOURCE_ID = re.compile(r"^[A-Za-z0-9_-]{1,64}$") +DIGEST = re.compile(r"^[0-9a-f]{64}$") +LIMITS = { + "checkpoint": 256 * 1024**2, + "onnx": 64 * 1024**2, + "env": 2 * 1024**2, + "agent": 128 * 1024, +} +TASKS = ["Unitree-Go2-Flat", "Unitree-Go2-ObstacleAvoidance"] + + +class SourceError(ValueError): + pass + + +def regular_bytes(path: Path, root: Path, limit: int) -> bytes: + """Walk with dir_fd/O_NOFOLLOW, so symlink substitution cannot escape the root.""" + if ".." in path.parts or not path.is_absolute(): + raise SourceError("基础策略路径必须是允许根内的绝对路径,不能含 ..") + try: + parts = path.relative_to(root).parts + except ValueError as error: + raise SourceError("基础策略文件不在管理员配置的允许根内") from error + if not parts: + raise SourceError("基础策略必须是常规文件") + try: + fd = os.open(root, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW) + except OSError as error: + raise SourceError("基础策略快照目录缺失或无效,拒绝训练") from error + try: + for part in parts[:-1]: + next_fd = os.open(part, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW, dir_fd=fd) + os.close(fd) + fd = next_fd + leaf = os.open(parts[-1], os.O_RDONLY | os.O_NOFOLLOW | os.O_NONBLOCK, dir_fd=fd) + with os.fdopen(leaf, "rb") as stream: + info = os.fstat(stream.fileno()) + if not stat.S_ISREG(info.st_mode) or info.st_size > limit: + raise SourceError("基础策略文件必须是常规文件且不能超过大小上限") + data = stream.read(limit + 1) + if len(data) > limit: + raise SourceError("基础策略文件超过大小上限") + return data + except OSError as error: + raise SourceError( + "基础策略文件缺失或含symlink;请提供.pt、policy.onnx和params配置常规文件" + ) from error + finally: + os.close(fd) + + +class PretrainedSources: + def __init__(self, config_path: Path | None, store: Path, python: str, trainer_root: Path): + self.store = store.expanduser().resolve() + self.store.mkdir(parents=True, exist_ok=True, mode=0o700) + self.python = python + self.trainer_root = trainer_root.resolve() + self.entries: dict[str, dict] = {} + self._upload_lock = threading.Lock() + self._catalog_lock = threading.RLock() + self._restore_uploads() + if config_path is None: + return + try: + if config_path.stat().st_size > 128 * 1024: + raise SourceError("基础策略注册配置不能超过128KiB") + config = json.loads(config_path.read_text()) + if not isinstance(config, dict) or set(config) != {"allowedRoots", "sources"}: + raise SourceError("注册配置只接受allowedRoots和sources") + if ( + not isinstance(config["allowedRoots"], list) + or not isinstance(config["sources"], list) + or len(config["sources"]) > 32 + ): + raise SourceError("允许根/sources必须为数组,最多32个注册条目") + roots = [Path(p).expanduser().resolve(strict=True) for p in config["allowedRoots"]] + if not roots or any(not p.is_dir() for p in roots): + raise SourceError("请配置存在的本地允许根目录") + for entry in config["sources"]: + if not isinstance(entry, dict) or set(entry) != { + "id", + "label", + "checkpoint", + "onnx", + }: + raise SourceError("注册条目必须包含id/label/checkpoint/onnx") + key = entry["id"] + if not isinstance(key, str) or not SOURCE_ID.fullmatch(key) or key in self.entries: + raise SourceError("注册id无效或重复") + if not isinstance(entry["label"], str) or not 1 <= len(entry["label"]) <= 100: + raise SourceError("基础策略label必须为1–100字符") + record = {"id": key, "label": entry["label"], "ready": False, "compatibleTasks": []} + self.entries[key] = {"public": record} + try: + checkpoint, onnx = ( + Path(entry["checkpoint"]).expanduser(), + Path(entry["onnx"]).expanduser(), + ) + if not re.fullmatch(r"[A-Za-z0-9_.-]+\.pt", checkpoint.name): + raise SourceError("请选择ONNX对应的.pt训练checkpoint,ONNX不能直接续训") + files = { + "checkpoint": checkpoint, + "onnx": onnx, + "env": checkpoint.parent / "params/env.yaml", + "agent": checkpoint.parent / "params/agent.yaml", + } + data = {} + for name, path in files.items(): + root = next((r for r in roots if path.is_relative_to(r)), None) + if root is None: + raise SourceError("注册文件不在允许根内") + data[name] = regular_bytes(path, root, LIMITS[name]) + names = { + "checkpoint": checkpoint.name, + "onnx": "policy.onnx", + "env": "params/env.yaml", + "agent": "params/agent.yaml", + } + with tempfile.TemporaryDirectory(dir=self.store, prefix="import-") as temporary: + directory = Path(temporary) + for name, content in data.items(): + destination = directory / names[name] + destination.parent.mkdir(exist_ok=True) + destination.write_bytes(content) + manifest = self._validate(directory, checkpoint.name, TASKS[0], None) + digest = manifest["source_id"] + bound = { + "sourceId": digest, + "registeredId": key, + "label": entry["label"], + "manifest": manifest, + } + destination = self.store / digest + if not destination.exists(): + shutil.copytree(directory, destination) + for file in destination.rglob("*"): + file.chmod(0o555 if file.is_dir() else 0o444) + destination.chmod(0o555) + self.verify(bound) + record.update( + id=digest, + ready=True, + compatibleTasks=TASKS, + observationSizes=[47, 81, 97], + initialization=bound, + ) + except (SourceError, OSError, ValueError, TypeError) as error: + record["error"] = f"{error};请修正管理员注册配置并重启服务,不会退回随机初始化" + except (OSError, TypeError, ValueError) as error: + raise SourceError(f"基础策略注册配置无效:{error}") from error + + def _restore_uploads(self): + # Only committed content directories are catalogued. Partial uploads are never sources. + for directory in self.store.iterdir(): + if not DIGEST.fullmatch(directory.name) or not (directory / "upload.json").is_file(): + continue + try: + manifest = json.loads( + regular_bytes(directory / "upload.json", self.store, 64 * 1024) + ) + label = json.loads(regular_bytes(directory / "label.json", self.store, 1024))[ + "label" + ] + bound = { + "sourceId": directory.name, + "registeredId": directory.name, + "label": self._display_name(label), + "manifest": manifest, + } + if manifest["source_id"] != directory.name: + raise SourceError("上传来源身份损坏") + self.verify(bound) + self._publish_upload(bound) + except (SourceError, ValueError, KeyError, TypeError, OSError): + self.entries[directory.name] = { + "public": { + "id": directory.name, + "label": "失效的已上传策略", + "ready": False, + "compatibleTasks": [], + "error": "上传快照损坏,拒绝随机初始化;请重新上传原文件或恢复快照", + } + } + + @staticmethod + def _display_name(name): + if not isinstance(name, str): + return "uploaded-policy" + name = name.replace("\\", "/").split("/")[-1] + return ( + re.sub(r"[^\w .()\-]", "_", name, flags=re.UNICODE)[:100].strip(" .") + or "uploaded-policy" + ) + + def _publish_upload(self, bound): + record = { + "id": bound["sourceId"], + "label": bound["label"], + "ready": True, + "compatibleTasks": TASKS, + "observationSizes": [47, 81, 97], + "initialization": bound, + } + with self._catalog_lock: + self.entries[bound["sourceId"]] = {"public": record} + return deepcopy(record) + + @contextmanager + def upload_slot(self, length): + # At most one receiving/decoding upload, 32 sources and 2GiB cumulative store. + if not self._upload_lock.acquire(blocking=False): + raise SourceError("已有文件正在上传/验证,请稍后重试") + try: + files = list(self.store.rglob("*")) + used = sum(p.lstat().st_size for p in files if not p.is_symlink() and p.is_file()) + count = sum(1 for p in self.store.iterdir() if DIGEST.fullmatch(p.name)) + if count >= 32 or used + length + 16 * 1024**2 > 2 * 1024**3: + raise SourceError("基础策略存储已达32个来源或2GiB上限,请由本机维护者释放空间") + with tempfile.TemporaryDirectory(dir=self.store, prefix="upload-") as temporary: + yield Path(temporary) + finally: + self._upload_lock.release() + + def receive_upload(self, stream, length, fmt, template, display_name, *, set_timeout=None): + if template != "go2-legacy47-v1": + raise SourceError("必须明确确认go2-legacy47-v1观测与关节模板") + limit = {"pt": 256 * 1024**2, "onnx": 64 * 1024**2}.get(fmt) + if limit is None: + raise SourceError("仅支持单个.pt或.onnx,不接受ZIP/路径/配套目录") + if type(length) is not int or not 0 < length <= limit: + raise SourceError("上传文件为空或超过.pt 256MiB / ONNX 64MiB上限") + with self.upload_slot(length) as directory: + path = directory / f"upload.{fmt}" + deadline = time.monotonic() + 60 + remaining = length + try: + with path.open("xb") as output: + while remaining: + timeout = deadline - time.monotonic() + if timeout <= 0: + raise SourceError("上传超时,请重试") + try: + if set_timeout is not None: + set_timeout(min(10, timeout)) + # BufferedReader.read(n) may perform many recv calls whose socket + # timeouts reset with each trickled byte. read1 returns after one + # raw read, allowing the total deadline to be checked every time. + read_chunk = getattr(stream, "read1", stream.read) + chunk = read_chunk(min(1024 * 1024, remaining)) + except OSError as error: + raise SourceError("上传超时/连接中断,临时文件已清理") from error + if not chunk or len(chunk) > remaining: + raise SourceError("上传连接中断或实际长度不符") + output.write(chunk) + remaining -= len(chunk) + finally: + if set_timeout is not None: + set_timeout(10) + request = { + "path": str(path), + "directory": str(directory), + "format": fmt, + "template": template, + } + env = { + **os.environ, + "CUDA_VISIBLE_DEVICES": "", + "OMP_NUM_THREADS": "1", + "OPENBLAS_NUM_THREADS": "1", + "MKL_NUM_THREADS": "1", + } + try: + result = subprocess.run( + [self.python, str(Path(__file__).parent / "rl/scripts/validate_upload.py")], + input=json.dumps(request), + capture_output=True, + text=True, + cwd=self.trainer_root, + env=env, + timeout=60, + ) + manifest = json.loads(result.stdout) + except (OSError, subprocess.TimeoutExpired, ValueError) as error: + raise SourceError("模型验证超时/资源超限或训练Python依赖不可用") from error + if result.returncode: + raise SourceError(manifest.get("error", "不支持的模型文件")) + digest = manifest["source_id"] + if not DIGEST.fullmatch(digest): + raise SourceError("模型验证返回无效身份") + label = self._display_name(display_name) + (directory / "upload.json").write_text(json.dumps(manifest), encoding="utf-8") + (directory / "label.json").write_text(json.dumps({"label": label}), encoding="utf-8") + destination = self.store / digest + if destination.exists(): + # Dedup never replaces a committed artifact or changes an old job binding. + old = next((r for r in self.catalog() if r["id"] == digest and r["ready"]), None) + if old is None: + raise SourceError("同ID旧快照已损坏,请恢复快照后重试;不会覆盖旧任务来源") + self.verify(old["initialization"]) + return old + for file in directory.iterdir(): + file.chmod(0o444) + directory.rename(destination) + destination.chmod(0o555) + bound = { + "sourceId": digest, + "registeredId": digest, + "label": label, + "manifest": manifest, + } + self.verify(bound) + return self._publish_upload(bound) + + def _validate(self, directory: Path, checkpoint: str, task_id: str, task_config: dict | None): + command = [self.python, str(Path(__file__).parent / "rl/scripts/validate_pretrained.py")] + request = { + "directory": str(directory), + "checkpoint": checkpoint, + "taskId": task_id, + "taskConfig": task_config, + "uploaded": (directory / "upload.json").is_file(), + } + try: + result = subprocess.run( + command, + input=json.dumps(request), + cwd=self.trainer_root, + capture_output=True, + text=True, + timeout=60, + ) + except (subprocess.TimeoutExpired, OSError) as error: + raise SourceError("基础策略验证器超时或不可用,请检查训练Python环境") from error + if result.returncode: + # Validator emits only an actionable error, never a traceback or file contents. + raise SourceError( + result.stderr.strip()[-1500:] or "基础策略验证器失败,请检查训练Python依赖" + ) + try: + return json.loads(result.stdout) + except ValueError as error: + raise SourceError("基础策略验证器未返回有效身份描述") from error + + def catalog(self): + with self._catalog_lock: + sources = {} + for entry in self.entries.values(): + sources.setdefault(entry["public"]["id"], deepcopy(entry["public"])) + return list(sources.values()) + + def bind(self, source_id, task_id: str, task_config=None): + # Selection is content-addressed, so a stale UI cannot silently bind a newly + # registered model under the same administrator-friendly registration name. + entry = next( + (entry for entry in self.catalog() if entry["id"] == source_id), + None, + ) + if not isinstance(source_id, str) or entry is None: + raise SourceError("基础策略内容ID不存在或已变化;请刷新服务连接,不接受文件路径") + if not entry["ready"]: + raise SourceError(entry["error"]) + if task_id not in entry["compatibleTasks"]: + raise SourceError("该基础策略仅兼容Flat47/Obstacle81或97;Rough高度扫描不支持迁移") + bound = deepcopy(entry["initialization"]) + directory = self.verify(bound) + manifest = self._validate( + directory, bound["manifest"]["artifacts"]["checkpoint"]["name"], task_id, task_config + ) + if manifest["source_id"] != bound["sourceId"]: + raise SourceError("基础策略快照身份变化,拒绝初始化") + return bound + + def verify(self, bound: dict) -> Path: + try: + digest = bound["sourceId"] + if not isinstance(digest, str) or not DIGEST.fullmatch(digest): + raise SourceError("基础策略快照身份无效") + directory = self.store / digest + artifacts = bound["manifest"]["artifacts"] + if "sourceFormat" in bound["manifest"]: + fmt = bound["manifest"]["sourceFormat"] + if fmt not in ("pt", "onnx"): + raise SourceError("上传格式无效") + stored = json.loads(regular_bytes(directory / "upload.json", self.store, 64 * 1024)) + if stored != bound["manifest"]: + raise SourceError("上传manifest SHA绑定变化,拒绝初始化") + for name, relative, limit in ( + ( + "upload", + f"upload.{fmt}", + LIMITS["checkpoint"] if fmt == "pt" else LIMITS["onnx"], + ), + ("checkpoint", "actor.pt", 16 * 1024**2), + ): + content = regular_bytes(directory / relative, self.store, limit) + if ( + artifacts[name]["name"] != relative + or hashlib.sha256(content).hexdigest() != artifacts[name]["sha256"] + ): + raise SourceError("上传快照SHA不匹配,拒绝训练") + return directory + for name, relative in { + "checkpoint": artifacts["checkpoint"]["name"], + "onnx": "policy.onnx", + "env": "params/env.yaml", + "agent": "params/agent.yaml", + }.items(): + content = regular_bytes(directory / relative, self.store, LIMITS[name]) + if hashlib.sha256(content).hexdigest() != artifacts[name]["sha256"]: + raise SourceError("基础策略快照SHA不匹配,拒绝训练;请恢复原快照") + return directory + except (KeyError, TypeError) as error: + raise SourceError("持久化基础策略描述损坏,拒绝训练") from error + + def arguments(self, bound: dict) -> list[str]: + directory = self.verify(bound) + upload_args = ( + ["--pretrained-upload-manifest", str(directory / "upload.json")] + if "sourceFormat" in bound["manifest"] + else [] + ) + return upload_args + [ + "--pretrained-checkpoint", + str(directory / bound["manifest"]["artifacts"]["checkpoint"]["name"]), + "--pretrained-allowed-roots", + json.dumps([str(directory)]), + "--pretrained-source-id", + bound["sourceId"], + ] diff --git a/training_server/pretrained_upload.py b/training_server/pretrained_upload.py new file mode 100644 index 00000000..522b8a69 --- /dev/null +++ b/training_server/pretrained_upload.py @@ -0,0 +1,336 @@ +"""Single-file Go2 legacy47 import. No adjacent files or arbitrary ONNX conversion.""" + +import hashlib +import io +import json +from dataclasses import asdict +from pathlib import Path + +import torch +from pretrained import ( + BASE_TERMS, + JOINTS, + NORMALIZATION_POLICY, + PretrainedError, + ValidatedSource, + _require, + make_reference_actor, + validate_actor_state, + validate_semantics, + verify_onnx, +) + +TEMPLATE = "go2-legacy47-v1" +SYNTHETIC_COUNT = 1_000_000 +UPLOAD_LIMITS = {"pt": 256 * 1024**2, "onnx": 64 * 1024**2} +SEMANTICS = { + "joint_names": JOINTS, + "observation_names": BASE_TERMS, + "command_names": ["twist"], + "action_scale": [0.25], + "joint_stiffness": [20, 20, 40] * 4, + "joint_damping": [1, 1, 2] * 4, + "default_joint_pos": [-0.1, 0.9, -1.8, 0.1, 0.9, -1.8] * 2, +} + + +def validate_metadata(metadata, *, required=False): + """Check recognized embedded semantics; absence remains a user template assumption.""" + import numpy as np + + _require(isinstance(metadata, dict), "模型metadata必须是对象") + checked = [] + for key, expected in SEMANTICS.items(): + if key not in metadata: + _require(not required, f"ONNX缺少语义metadata: {key}") + continue + actual = metadata[key] + if isinstance(actual, str): + actual = actual.split(",") + if isinstance(expected[0], str): + _require(actual == expected, f"模型metadata冲突: {key}") + else: + try: + actual = np.asarray(actual, dtype=np.float64) + _require( + actual.shape == (len(expected),) + and np.isfinite(actual).all() + and np.allclose(actual, expected, atol=1e-6, rtol=0), + f"模型metadata冲突: {key}", + ) + except (TypeError, ValueError) as error: + raise PretrainedError(f"模型metadata无效: {key}") from error + checked.append(key) + for key in ("contract", "templateId"): + if key in metadata: + _require(metadata[key] == TEMPLATE, "模型模板metadata冲突") + checked.append(key) + return checked + + +def _checkpoint(data): + try: + checkpoint = torch.load(io.BytesIO(data), map_location="cpu", weights_only=True) + except Exception as error: + raise PretrainedError( + "不支持或不安全的.pt;仅支持weights_only Go2 legacy47 actor checkpoint" + ) from error + _require( + isinstance(checkpoint, dict) and "actor_state_dict" in checkpoint, + "请选择含actor_state_dict的Go2 legacy47 .pt,不接受ZIP工程或任意模型", + ) + state = checkpoint["actor_state_dict"] + if isinstance(state, dict) and isinstance(state.get("mlp.0.weight"), torch.Tensor): + _require( + tuple(state["mlp.0.weight"].shape) == (512, 47), + "当前仅支持Go2 legacy47输入;81/97 checkpoint请使用原trial续训,而不是基础策略上传", + ) + validate_actor_state(state) + iteration = checkpoint.get("iter") + _require( + iteration is None or (type(iteration) is int and iteration >= 0), "无效checkpoint iteration" + ) + checked = validate_metadata(checkpoint) + for key in ("metadata", "infos"): + if key in checkpoint: + checked.extend(validate_metadata(checkpoint[key])) + nested = checkpoint[key].get("metadata") + if nested is not None: + checked.extend(validate_metadata(nested)) + return state, iteration, sorted(set(checked)) + + +def _onnx(data): + import onnx + from onnx import helper, numpy_helper + + model = onnx.load_model_from_string(data) + graph = model.graph + _require( + not model.functions + and not model.training_info + and len(model.opset_import) == 1 + and model.opset_import[0].domain == "" + and model.opset_import[0].version in (17, 18), + "仅支持标准opset17/18受限MLP ONNX", + ) + _require( + not graph.sparse_initializer and not graph.quantization_annotation, "不支持稀疏或量化ONNX" + ) + _require(len(graph.input) == len(graph.output) == 1, "ONNX必须单输入单输出") + for value, shape in ((graph.input[0], [1, 47]), (graph.output[0], [1, 12])): + tensor = value.type.tensor_type + _require( + tensor.elem_type == onnx.TensorProto.FLOAT + and [d.dim_value for d in tensor.shape.dim] == shape + and all(not d.dim_param for d in tensor.shape.dim), + "ONNX仅支持float32 [1,47] -> [1,12];其他输入请使用对应训练器", + ) + expected_shapes = {"obs_normalizer._mean": (1, 47), "onnx::Div_24": (1, 47)} + for i, shape in zip((0, 2, 4, 6), ((512, 47), (256, 512), (128, 256), (12, 128)), strict=True): + expected_shapes[f"mlp.{i}.weight"] = shape + expected_shapes[f"mlp.{i}.bias"] = (shape[0],) + _require( + len(graph.initializer) == len(expected_shapes) + and {t.name for t in graph.initializer} == set(expected_shapes), + "ONNX initializer不符合受支持MLP", + ) + tensors = {} + for tensor in graph.initializer: + _require( + tensor.data_location == onnx.TensorProto.DEFAULT and not tensor.external_data, + "拒绝ONNX external data;必须是单个自包含文件", + ) + _require( + tensor.data_type == onnx.TensorProto.FLOAT + and tuple(tensor.dims) == expected_shapes[tensor.name], + "ONNX tensor类型/shape不支持", + ) + tensors[tensor.name] = torch.from_numpy(numpy_helper.to_array(tensor).copy()) + _require(bool(torch.isfinite(tensors[tensor.name]).all()), "ONNX tensor含非有限值") + # Exact dataflow, not merely op/tensor names: no branch, reorder, alias or extra op. + expected_ops = ["Sub", "Div", "Gemm", "Elu", "Gemm", "Elu", "Gemm", "Elu", "Gemm"] + _require( + [n.op_type for n in graph.node] == expected_ops, "ONNX必须是Sub/Div及4层Gemm+3层ELU精确链路" + ) + previous = graph.input[0].name + seen = set(tensors) | {previous} + _require(len(seen) == len(tensors) + 1, "ONNX输入与initializer重名") + layer = 0 + for node in graph.node: + _require(node.domain == "" and not node.overload, "拒绝ONNX custom op") + attributes = {a.name: helper.get_attribute_value(a) for a in node.attribute} + _require(len(attributes) == len(node.attribute), "重复ONNX属性") + if node.op_type in ("Sub", "Div"): + inputs = [previous, "obs_normalizer._mean" if node.op_type == "Sub" else "onnx::Div_24"] + _require(not attributes, "不支持normalizer算子属性") + elif node.op_type == "Gemm": + inputs = [previous, f"mlp.{layer}.weight", f"mlp.{layer}.bias"] + layer += 2 + _require( + set(attributes) <= {"alpha", "beta", "transA", "transB"} + and attributes.get("alpha", 1.0) == 1.0 + and attributes.get("beta", 1.0) == 1.0 + and attributes.get("transA", 0) == 0 + and attributes.get("transB", 0) == 1, + "不支持Gemm缩放/转置属性", + ) + else: + inputs = [previous] + _require( + set(attributes) <= {"alpha"} and attributes.get("alpha", 1.0) == 1.0, + "仅支持ELU alpha=1", + ) + _require( + list(node.input) == inputs + and len(node.output) == 1 + and node.output[0] + and node.output[0] not in seen, + "ONNX实际连边/输出不符合受支持MLP", + ) + previous = node.output[0] + seen.add(previous) + _require(previous == graph.output[0].name, "ONNX输出必须是最后Gemm结果") + metadata = {p.key: p.value for p in model.metadata_props} + _require(len(metadata) == len(model.metadata_props), "重复ONNX metadata") + checked = validate_metadata(metadata, required=True) + onnx.checker.check_model(model, full_check=True) + actor = make_reference_actor() + _require( + actor.obs_normalizer.eps == 0.01 + and bool((actor.state_dict()["distribution.std_param"] == 1).all()), + "目标训练默认normalizer/exploration已变化,需要新模板", + ) + state = actor.state_dict() + for key in state: + if key in tensors: + state[key] = tensors[key] + std = tensors["onnx::Div_24"] - 0.01 + _require(bool((std > 0).all()), "ONNX denominator必须大于模板epsilon=.01") + state["obs_normalizer._std"] = std + state["obs_normalizer._var"] = std.square() + state["obs_normalizer.count"].fill_(SYNTHETIC_COUNT) + validate_actor_state(state) + actor.load_state_dict(state) + identity = verify_onnx(data, actor) + return state, checked, identity + + +def source_identity(fmt, digest): + return hashlib.sha256(f"{TEMPLATE}:{fmt}:{digest}".encode()).hexdigest() + + +def import_upload(path, fmt, template, directory): + """Run only inside the resource-limited validator process.""" + _require(template == TEMPLATE, "必须明确确认go2-legacy47-v1模板") + _require(fmt in UPLOAD_LIMITS, "仅支持单个.pt或.onnx,不支持ZIP") + data = Path(path).read_bytes() + _require(0 < len(data) <= UPLOAD_LIMITS[fmt], "上传文件为空或过大") + if fmt == "pt": + state, iteration, checked = _checkpoint(data) + identity = None + else: + state, checked, identity = _onnx(data) + iteration = None + from pretrained import comparison_observations + + actor = make_reference_actor().eval() + actor.load_state_dict(state) + with torch.inference_mode(): + output = actor.mlp(actor.obs_normalizer(comparison_observations())) + _require(bool(torch.isfinite(output).all()), "actor在随机/物理probe上产生非有限动作") + directory = Path(directory) + actor_path = directory / "actor.pt" + torch.save({"actor_state_dict": state}, actor_path) + digest = hashlib.sha256(data).hexdigest() + manifest = { + "schema_version": 1, + "mode": "pretrained-warm-start", + "contract": TEMPLATE, + "sourceFormat": fmt, + "source_id": source_identity(fmt, digest), + "source_iteration": iteration, + "source_actor_dim": 47, + "source_normalizer_count": state["obs_normalizer.count"].item(), + "normalization": NORMALIZATION_POLICY + if fmt == "pt" + else "synthetic-count/unit-new-features", + "template_confirmation": { + "id": TEMPLATE, + "confirmed_by": "user", + "assumptions": ( + "Go2 legacy47 observation physics, ordering, 50Hz and action semantics; " + "not verified source env.yaml" + ), + }, + "verified_facts": { + "actor_tensor_shapes": "47-512-256-128-12/float32", + "metadata_fields": checked, + "finite_output_probes": 48, + "activation": "graph-verified-ELU" if fmt == "onnx" else "user-template-assumed-ELU", + }, + "derived_fields": {} + if fmt == "pt" + else { + "normalizer_count": {"policy": "synthetic", "value": SYNTHETIC_COUNT}, + "normalizer_std": "denominator - 0.01", + "normalizer_var": "std squared", + "epsilon": 0.01, + "exploration_std": {"policy": "fresh-target-default", "value": 1.0}, + }, + "onnx_identity": identity, + "critic": "fresh-target-initialization", + "optimizer": "fresh", + "iteration": 0, + "base_observation_terms": BASE_TERMS, + "joint_names": JOINTS, + "artifacts": { + "upload": {"name": f"upload.{fmt}", "sha256": digest, "bytes": len(data)}, + "checkpoint": { + "name": "actor.pt", + "sha256": hashlib.sha256(actor_path.read_bytes()).hexdigest(), + "bytes": actor_path.stat().st_size, + "origin": "service-derived-actor-only", + }, + }, + } + return manifest + + +def read_uploaded_source(checkpoint, *, allowed_roots, manifest_path, target_env, target_agent): + """Load only the service-derived actor and bound provenance, never adjacent sidecars.""" + from pretrained_sources import regular_bytes + + roots = [Path(p).resolve(strict=True) for p in allowed_roots] + checkpoint, manifest_path = Path(checkpoint), Path(manifest_path) + root = next( + (r for r in roots if checkpoint.is_relative_to(r) and manifest_path.is_relative_to(r)), None + ) + _require(root is not None, "上传artifact不在受控根内") + manifest = json.loads(regular_bytes(manifest_path, root, 64 * 1024)) + _require( + manifest.get("contract") == TEMPLATE and manifest.get("sourceFormat") in UPLOAD_LIMITS, + "无效上传manifest模板/格式", + ) + artifacts = manifest["artifacts"] + _require( + manifest["source_id"] + == source_identity(manifest["sourceFormat"], artifacts["upload"]["sha256"]), + "上传原始SHA身份不匹配", + ) + data = regular_bytes(checkpoint, root, UPLOAD_LIMITS["pt"]) + _require( + hashlib.sha256(data).hexdigest() == artifacts["checkpoint"]["sha256"], "上传actor SHA不匹配" + ) + from pretrained import _plain + 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 + + validate_semantics( + _plain(asdict(unitree_go2_flat_env_cfg())), + _plain(asdict(unitree_go2_ppo_runner_cfg())), + target_env, + target_agent, + ) + state, _, _ = _checkpoint(data) + return ValidatedSource(state, manifest) diff --git a/training_server/rl/requirements.txt b/training_server/rl/requirements.txt index ad49b6b6..e1dd03f4 100644 --- a/training_server/rl/requirements.txt +++ b/training_server/rl/requirements.txt @@ -7,3 +7,7 @@ mujoco-warp==3.5.0 warp-lang==1.15.0 # RSL-RL 5.0.1 仍向 wandb.Settings 传递 start_method,0.29 已删除该字段。 wandb==0.28.2 +# 基础策略迁移依赖已验证的 actor/normalizer/checkpoint 契约。 +rsl-rl-lib==5.0.1 +# 基础策略身份校验只使用 CPU ONNX Runtime。 +onnxruntime==1.29.0 diff --git a/training_server/rl/scripts/evaluate.py b/training_server/rl/scripts/evaluate.py index 6c78bc0b..862e59f6 100644 --- a/training_server/rl/scripts/evaluate.py +++ b/training_server/rl/scripts/evaluate.py @@ -32,7 +32,6 @@ from mjlab.tasks.registry import list_tasks, load_env_cfg, load_rl_cfg, load_run from mjlab.tasks.velocity.mdp import UniformVelocityCommandCfg from mjlab.utils.torch import configure_torch_backends from torch.utils.tensorboard import SummaryWriter - from tuning.schema import apply_reward_configuration, validate_configuration SCENARIOS = ( @@ -61,6 +60,7 @@ class EvaluateConfig: checkpoint: str output: str reward_config: str | None = None + task_config: str | None = None num_envs: int = 256 steps_per_seed: int = 1000 seeds: tuple[int, ...] = field(default_factory=lambda: (101, 202, 303)) @@ -163,31 +163,41 @@ def run_evaluation(task_id: str, cfg: EvaluateConfig) -> dict: checkpoint = Path(cfg.checkpoint).expanduser().resolve(strict=True) output = Path(cfg.output).expanduser().resolve() output.parent.mkdir(parents=True, exist_ok=True) - per_seed = [_evaluate_seed(task_id, cfg, seed) for seed in cfg.seeds] - metrics = { - key: fmean(seed_metrics[key] for seed_metrics in per_seed) - for key in (*METRIC_NAMES, "fall_rate") - } - deviations = { - key: pstdev(seed_metrics[key] for seed_metrics in per_seed) - for key in (*METRIC_NAMES, "fall_rate") - } - result = { - "protocolVersion": 1, - "taskId": task_id, - "checkpoint": checkpoint.name, - "checkpointSha256": _sha256(checkpoint), - "seeds": list(cfg.seeds), - "numEnvs": cfg.num_envs, - "stepsPerSeed": cfg.steps_per_seed, - "scenarios": [list(value) for value in SCENARIOS], - "metrics": metrics, - "metricStd": deviations, - "seedMetrics": [ - {"seed": seed, "metrics": values} - for seed, values in zip(cfg.seeds, per_seed, strict=True) - ], - } + if task_id == "Unitree-Go2-ObstacleAvoidance": + from scripts.evaluate_obstacle import run_obstacle_evaluation + checkpoint_hash = _sha256(checkpoint) + result = run_obstacle_evaluation(task_id, cfg) + if _sha256(checkpoint) != checkpoint_hash: + raise ValueError("Checkpoint changed during evaluation") + result.update({"taskId": task_id, "checkpoint": checkpoint.name, + "checkpointSha256": checkpoint_hash}) + metrics = result["metrics"] + else: + per_seed = [_evaluate_seed(task_id, cfg, seed) for seed in cfg.seeds] + metrics = { + key: fmean(seed_metrics[key] for seed_metrics in per_seed) + for key in (*METRIC_NAMES, "fall_rate") + } + deviations = { + key: pstdev(seed_metrics[key] for seed_metrics in per_seed) + for key in (*METRIC_NAMES, "fall_rate") + } + result = { + "protocolVersion": 1, + "taskId": task_id, + "checkpoint": checkpoint.name, + "checkpointSha256": _sha256(checkpoint), + "seeds": list(cfg.seeds), + "numEnvs": cfg.num_envs, + "stepsPerSeed": cfg.steps_per_seed, + "scenarios": [list(value) for value in SCENARIOS], + "metrics": metrics, + "metricStd": deviations, + "seedMetrics": [ + {"seed": seed, "metrics": values} + for seed, values in zip(cfg.seeds, per_seed, strict=True) + ], + } output.write_text(json.dumps(result, ensure_ascii=False, indent=2) + "\n", encoding="utf-8") writer = SummaryWriter(log_dir=str(output.parent / "evaluation-events")) try: diff --git a/training_server/rl/scripts/evaluate_obstacle.py b/training_server/rl/scripts/evaluate_obstacle.py new file mode 100644 index 00000000..143354d1 --- /dev/null +++ b/training_server/rl/scripts/evaluate_obstacle.py @@ -0,0 +1,241 @@ +"""Actual checkpoint inference, capturing first-terminal physics before mjlab resets.""" + +import hashlib +import json +import shutil +import subprocess +import sys +import tempfile +from dataclasses import asdict +from pathlib import Path +from statistics import fmean +from types import SimpleNamespace + +# Also executable as a fresh-interpreter seed worker (never fork a CUDA context). +for source in (Path(__file__).resolve().parents[1], Path(__file__).resolve().parents[2]): + if str(source) not in sys.path: + sys.path.insert(0, str(source)) + +import torch +from mjlab.envs import ManagerBasedRlEnv +from mjlab.rl import MjlabOnPolicyRunner, RslRlVecEnvWrapper +from mjlab.sensor import ContactMatch, ContactSensorCfg +from mjlab.tasks.registry import load_env_cfg, load_rl_cfg, load_runner_cls +from scripts.train import _load_reward_config, _load_task_config +from src.tasks.obstacle_avoidance.env_cfg import apply_obstacle_configuration +from task_config import OBSTACLE_TASK +from tuning.obstacle_scoring import METRICS, STEPS, protocol, score_trajectory +from tuning.schema import apply_reward_configuration + + +class FirstEpisodeRecorder: + """The termination hook runs post-physics, before _reset_idx (including first terminal).""" + + def __init__(self, env, layout): + self.env = env + self.layout = layout + self.samples = [[] for _ in range(env.num_envs)] + self.previous = torch.zeros((env.num_envs, 12), device=env.device) + self.calls = 0 + self.compute = env.termination_manager.compute + env.termination_manager.compute = self.capture + + def capture(self): + env = self.env + terminal = self.compute() + robot = env.scene["robot"].data + # Like mjlab termination itself, derived pose is one physics substep old. + position = robot.root_link_pos_w + local_xy = position[:, :2] - env.scene.env_origins[:, :2] + xy = local_xy + torch.tensor(self.layout["spawn"][:2], device=env.device) + boxes = self.layout["boxes"][1:] # Never include the support floor. + clearance = torch.full((env.num_envs,), 0.5, device=env.device) + if boxes: + centers = torch.tensor([b["pos"][:2] for b in boxes], device=env.device) + sizes = torch.tensor([b["size"][:2] for b in boxes], device=env.device) + clearance = ( + ((xy[:, None, :] - centers).abs() - sizes).clamp(min=0).norm(dim=2).amin(dim=1) + ) + clearance = (clearance - 0.3).clamp(min=0) # Conservative body footprint radius. + forces = env.scene["evaluation_obstacles"].data.force_history + collision = forces.norm(dim=-1).flatten(1).amax(dim=1) > 1.0 + else: + collision = torch.zeros(env.num_envs, dtype=torch.bool, device=env.device) + nonfoot = env.scene["nonfoot_ground_touch"].data.force_history + collision |= nonfoot.norm(dim=-1).flatten(1).amax(dim=1) > 10.0 + fall = (robot.projected_gravity_b[:, 2] > -0.3420201433) | (position[:, 2] < 0.12) + actions = env.action_manager.action + delta = (actions - self.previous).square().mean(dim=1) + self.previous.copy_(actions) + rays = env.scene["forward_scan"].data.distances + hits = ( + ((rays >= 0) & (rays <= env.scene["forward_scan"].cfg.max_distance)).float().mean(dim=1) + ) + distance = env.command_manager.get_term("twist").errors()[1] + values = torch.stack((distance, clearance, delta, hits, collision, fall, terminal), dim=1) + if not torch.isfinite(values).all(): + raise ValueError("Nonfinite rollout measurements") + keys = ("distance", "clearance", "action_delta", "ray_hit", "collision", "fall", "terminal") + for samples, row in zip(self.samples, values.cpu().tolist(), strict=True): + if not samples or not samples[-1]["terminal"]: + samples.append(dict(zip(keys, row, strict=True))) + self.calls += 1 + return terminal + + def metrics(self, horizon=STEPS): + if self.calls != horizon: + raise ValueError("Incomplete rollout horizon") + metrics = [score_trajectory(samples, horizon) for samples in self.samples] + return {key: fmean(item[key] for item in metrics) for key in METRICS} + + +def configure_seed(task_id, cfg, scenario, reward): + env_cfg = load_env_cfg(task_id, play=False) + env_cfg.seed = scenario["seed"] + env_cfg.scene.num_envs = cfg.num_envs + # Keep the benchmark's declared pair fixed; training itself randomizes every reset. + apply_obstacle_configuration(env_cfg, scenario["taskConfig"], randomize_navigation=False) + apply_reward_configuration(env_cfg, reward, task_id) + env_cfg.curriculum = {} + env_cfg.observations["actor"].enable_corruption = False + env_cfg.events.pop("push_robot", None) + # Primary geom names are literal compiled static terrain geoms, not a regex. + names = tuple(f"terrain_{i}" for i in range(1, len(scenario["terrain"]["boxes"]))) + if names: + env_cfg.scene.sensors += ( + ContactSensorCfg( + name="evaluation_obstacles", + primary=ContactMatch(mode="geom", pattern=names), + secondary=ContactMatch(mode="subtree", pattern="base_link", entity="robot"), + fields=("force",), + reduce="maxforce", + history_length=env_cfg.decimation, + ), + ) + return env_cfg + + +def evaluate_seed(task_id, cfg, scenario, reward): + torch.manual_seed(scenario["seed"]) + env_cfg = configure_seed(task_id, cfg, scenario, reward) + agent_cfg = load_rl_cfg(task_id) + env = ManagerBasedRlEnv(env_cfg, device=cfg.device) + wrapped = RslRlVecEnvWrapper(env, clip_actions=agent_cfg.clip_actions) + try: + runner_cls = load_runner_cls(task_id) or MjlabOnPolicyRunner + runner = runner_cls(wrapped, asdict(agent_cfg), log_dir=None, device=wrapped.device) + # rsl_rl actor state_dict includes observation_normalizer buffers. Strict load + # rejects incompatible architecture/statistics; no fixture or zero-action fallback. + runner.load( + str(Path(cfg.checkpoint).resolve(strict=True)), + load_cfg={"actor": True}, + strict=True, + map_location=str(wrapped.device), + ) + policy = runner.get_inference_policy(device=str(wrapped.device)) + obs, _ = env.reset(seed=scenario["seed"]) + recorder = FirstEpisodeRecorder(env, scenario["terrain"]) + with torch.inference_mode(): + for _ in range(STEPS): + actions = policy(obs) + if not torch.isfinite(actions).all(): + raise ValueError("Nonfinite policy actions") + obs, _, _, _ = wrapped.step(actions) + return { + "seed": scenario["seed"], + "metrics": recorder.metrics(), + "episodes": cfg.num_envs, + "rolloutSteps": STEPS, + } + finally: + wrapped.close() + + +def run_obstacle_evaluation(task_id, cfg): + from tuning.obstacle_scoring import SEEDS + + if tuple(cfg.seeds) != SEEDS or cfg.steps_per_seed != STEPS: + raise ValueError("Obstacle evaluation seeds/horizon are fixed") + if not cfg.task_config or not cfg.reward_config: + raise ValueError("Obstacle evaluation requires task and parameter configurations") + raw = json.loads(Path(cfg.task_config).read_text()) + custom = _load_task_config(task_id, cfg.task_config, raw["seed"]) + reward = _load_reward_config(cfg.reward_config, None, OBSTACLE_TASK) + fixed = protocol(custom, cfg.num_envs) + seeds = evaluate_isolated_seeds(task_id, cfg, fixed, reward) + return { + "protocol": fixed, + "seedMetrics": seeds, + "metrics": {key: fmean(item["metrics"][key] for item in seeds) for key in METRICS}, + } + + +def file_hash(path): + digest = hashlib.sha256() + with Path(path).open("rb") as stream: + for chunk in iter(lambda: stream.read(1024 * 1024), b""): + digest.update(chunk) + return digest.hexdigest() + + +def evaluate_isolated_seeds(task_id, cfg, fixed, reward): + from tuning.obstacle_scoring import validate_evaluation + + seeds = [] + # Inherit the evaluation parent's process group: manager cancellation kills + # both parent and the active seed. subprocess.run kills/reaps on timeout too. + with tempfile.TemporaryDirectory(prefix="go2-obstacle-eval-") as directory: + root = Path(directory) + checkpoint = root / "checkpoint.pt" + source_hash = file_hash(cfg.checkpoint) + shutil.copyfile(cfg.checkpoint, checkpoint) + if file_hash(checkpoint) != source_hash: + raise ValueError("Checkpoint changed during snapshot") + for scenario in fixed["scenarios"]: + request = root / f"request-{scenario['seed']}.json" + output = root / f"result-{scenario['seed']}.json" + request.write_text( + json.dumps( + { + "taskId": task_id, + "config": {**asdict(cfg), "checkpoint": str(checkpoint)}, + "scenario": scenario, + "reward": reward, + "checkpointSha256": source_hash, + }, + allow_nan=False, + ) + ) + subprocess.run( + [sys.executable, str(Path(__file__).resolve()), str(request), str(output)], + check=True, + timeout=1200, + shell=False, + ) + result = json.loads(output.read_text()) + if ( + result.pop("checkpointSha256", None) != source_hash + or file_hash(checkpoint) != source_hash + ): + raise ValueError("Seed checkpoint identity mismatch") + seeds.append(result) + candidate = { + "protocol": fixed, + "seedMetrics": seeds, + "metrics": {key: fmean(item["metrics"][key] for item in seeds) for key in METRICS}, + } + validate_evaluation(candidate, fixed) + return seeds + + +if __name__ == "__main__": + import src.tasks # noqa: F401 + + request_path, output_path = map(Path, sys.argv[1:]) + request = json.loads(request_path.read_text()) + config = SimpleNamespace(**request["config"]) + if file_hash(config.checkpoint) != request["checkpointSha256"]: + raise ValueError("Worker checkpoint identity mismatch") + result = evaluate_seed(request["taskId"], config, request["scenario"], request["reward"]) + result["checkpointSha256"] = file_hash(config.checkpoint) + output_path.write_text(json.dumps(result, allow_nan=False)) diff --git a/training_server/rl/scripts/train.py b/training_server/rl/scripts/train.py index 5fc7f1e1..35ce47cf 100644 --- a/training_server/rl/scripts/train.py +++ b/training_server/rl/scripts/train.py @@ -7,7 +7,7 @@ import sys from dataclasses import asdict, dataclass, field from datetime import datetime from pathlib import Path -from typing import Literal, cast +from typing import Literal # 训练器作为仓库内置子集直接从 scripts/ 启动,不要求额外执行 pip install -e。 TRAINER_ROOT = Path(__file__).resolve().parents[1] @@ -34,7 +34,7 @@ from mjlab.utils.gpu import select_gpus from mjlab.utils.os import dump_yaml, get_checkpoint_path from mjlab.utils.torch import configure_torch_backends from mjlab.utils.wrappers import VideoRecorder - +from task_config import OBSTACLE_TASK, deployment_metadata, validate_task_config from tuning.schema import apply_reward_configuration, validate_configuration @@ -51,8 +51,14 @@ class TrainConfig: gpu_ids: list[int] | Literal["all"] | None = field(default_factory=lambda: [0]) output_dir: str | None = None resume_checkpoint: str | None = None + pretrained_checkpoint: str | None = None + pretrained_upload_manifest: str | None = None + pretrained_onnx: str | None = None + pretrained_source_id: str | None = None + pretrained_allowed_roots: list[str] = field(default_factory=list) reward_config: str | None = None reward_config_json: str | None = None + task_config: str | None = None @staticmethod def from_task(task_id: str) -> "TrainConfig": @@ -61,27 +67,96 @@ class TrainConfig: return TrainConfig(env=env_cfg, agent=agent_cfg) -def _load_reward_config(path: str | None, inline: str | None) -> dict | None: +def _load_reward_config(path: str | None, inline: str | None, task_id="Unitree-Go2-Flat") -> dict | None: if path is not None and inline is not None: raise ValueError("Use only one of reward_config and reward_config_json") if inline is not None: if len(inline.encode("utf-8")) > 64 * 1024: raise ValueError("Reward configuration is larger than 64 KiB") - return validate_configuration(json.loads(inline)) + return validate_configuration(json.loads(inline), task_id) if path is None: return None source = Path(path).expanduser().resolve(strict=True) if source.stat().st_size > 64 * 1024: raise ValueError("Reward configuration is larger than 64 KiB") with source.open(encoding="utf-8") as stream: - return validate_configuration(json.load(stream)) + return validate_configuration(json.load(stream), task_id) + + +def _load_task_config(task_id: str, path: str | None, seed: int) -> dict | None: + if path is None: + return validate_task_config(task_id, {}, seed) + source = Path(path).expanduser().resolve(strict=True) + if source.stat().st_size > 128 * 1024: + raise ValueError("Task configuration is larger than 128 KiB") + payload = json.loads(source.read_text(encoding="utf-8")) + allowed = {"terrainPreset", "terrainParams", "sensorCfg", "seed", "customTerrainBoxes"} + if not isinstance(payload, dict) or payload.keys() - allowed: + raise ValueError("Unknown task configuration fields") + if ( + isinstance(payload.get("seed"), bool) + or not isinstance(payload.get("seed"), int) + or payload.get("seed") != seed + ): + raise ValueError("Task configuration seed must equal the agent seed") + if payload.get("sensorCfg") is None: + payload.pop("sensorCfg", None) + return validate_task_config(task_id, payload, seed) + + +def _configure_task_and_rewards(task_id: str, cfg: TrainConfig): + custom = _load_task_config(task_id, cfg.task_config, cfg.agent.seed) + if custom is not None: + from src.tasks.obstacle_avoidance.env_cfg import apply_obstacle_configuration + from src.tasks.obstacle_avoidance.terrain import apply_terrain_configuration + if task_id == OBSTACLE_TASK: + apply_obstacle_configuration(cfg.env, custom) + else: + apply_terrain_configuration(cfg.env, custom) + deployment = deployment_metadata(task_id, custom, cfg.agent.seed) + reward_config = _load_reward_config(cfg.reward_config, cfg.reward_config_json, task_id) + if reward_config is not None: + apply_reward_configuration(cfg.env, reward_config, task_id) + if task_id == OBSTACLE_TASK: + deployment["navigation"]["speed"] = reward_config["params"]["target_velocity"] + deployment["sensorCfg"]["avoidanceWeight"] = reward_config["weights"]["avoidance_weight"] + + return deployment, reward_config + + +def _load_pretrained(cfg: TrainConfig): + if cfg.pretrained_checkpoint is None: + if (cfg.pretrained_onnx is not None or cfg.pretrained_allowed_roots + or cfg.pretrained_source_id or cfg.pretrained_upload_manifest): + raise ValueError("Pretrained options require --pretrained-checkpoint (.pt), not ONNX alone") + return None + if cfg.resume_checkpoint is not None or cfg.agent.resume: + raise ValueError("Pretrained warm-start and resume are mutually exclusive") + from pretrained import read_pretrained_source + + options = {"onnx_path": cfg.pretrained_onnx} + if cfg.pretrained_upload_manifest is not None: + if cfg.pretrained_onnx is not None: + raise ValueError("Uploaded actor cannot use adjacent ONNX sidecars") + from pretrained_upload import read_uploaded_source + read_pretrained_source = read_uploaded_source + options = {"manifest_path": cfg.pretrained_upload_manifest} + source = read_pretrained_source( + cfg.pretrained_checkpoint, + allowed_roots=cfg.pretrained_allowed_roots, + **options, + target_env=asdict(cfg.env), + target_agent=asdict(cfg.agent), + ) + if cfg.pretrained_source_id is not None and source.manifest["source_id"] != cfg.pretrained_source_id: + raise ValueError("基础策略SHA身份变化,拒绝初始化") + return source def run_train(task_id: str, cfg: TrainConfig, log_dir: Path) -> None: - reward_config = _load_reward_config(cfg.reward_config, cfg.reward_config_json) - if reward_config is not None: - apply_reward_configuration(cfg.env, reward_config) - + deployment, reward_config = _configure_task_and_rewards(task_id, cfg) + # Validate source identity/semantics before allocating a simulation or optimizer. + pretrained = _load_pretrained(cfg) cuda_visible = os.environ.get("CUDA_VISIBLE_DEVICES", "") if cuda_visible == "": device = "cpu" @@ -135,6 +210,17 @@ def run_train(task_id: str, cfg: TrainConfig, log_dir: Path) -> None: cfg=cfg.env, device=device, render_mode="rgb_array" if cfg.video else None ) + if pretrained is not None: + from pretrained import validate_runtime_contract + + try: + validate_runtime_contract(env) + except Exception: + env.close() + raise + + # The ONNX runner exports this exact snapshot beside and inside policy.onnx. + env.platform_deployment = deployment log_root_path = log_dir.parent # Go up from specific run dir to experiment dir. resume_path: Path | None = None @@ -171,9 +257,30 @@ def run_train(task_id: str, cfg: TrainConfig, log_dir: Path) -> None: runner = runner_cls(env, agent_cfg, str(log_dir), device, **runner_kwargs) runner.add_git_repo_to_log(__file__) + if pretrained is not None: + from pretrained import initialize_runner + + initialization = initialize_runner(runner, pretrained) + env.unwrapped.platform_initialization = initialization + if rank == 0: + (log_dir / "initialization.json").write_text( + json.dumps(initialization, ensure_ascii=False, indent=2) + "\n", encoding="utf-8" + ) + print(f"[INFO] Warm-start actor from source {initialization['source_id']}; fresh critic/optimizer, iteration=0") if resume_path is not None: print(f"[INFO]: Loading model checkpoint from: {resume_path}") runner.load(str(resume_path), map_location=device) + # Promotion copies the origin manifest, never the original actor weights. + origin_path = resume_path.parent / "initialization.json" + if origin_path.is_file(): + if origin_path.stat().st_size > 64 * 1024: + raise ValueError("Initialization provenance exceeds 64 KiB") + origin = json.loads(origin_path.read_text(encoding="utf-8")) + if not isinstance(origin, dict) or origin.get("mode") != "pretrained-warm-start": + raise ValueError("Invalid initialization provenance") + env.unwrapped.platform_initialization = origin + if rank == 0: + (log_dir / "initialization.json").write_text(json.dumps(origin, ensure_ascii=False, indent=2) + "\n", encoding="utf-8") if explicit_resume: # RSL-RL stores the last completed zero-based iteration and otherwise # repeats it after load. Explicit tuning promotion uses an absolute target. @@ -260,7 +367,7 @@ def main(): # Parse first argument to choose the task. # Import tasks to populate the registry. import mjlab.tasks # noqa: F401 - import src.tasks + import src.tasks # noqa: F401 all_tasks = list_tasks() chosen_task, remaining_args = tyro.cli( diff --git a/training_server/rl/scripts/validate_pretrained.py b/training_server/rl/scripts/validate_pretrained.py new file mode 100644 index 00000000..f7f8f3d1 --- /dev/null +++ b/training_server/rl/scripts/validate_pretrained.py @@ -0,0 +1,48 @@ +"""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) diff --git a/training_server/rl/scripts/validate_upload.py b/training_server/rl/scripts/validate_upload.py new file mode 100644 index 00000000..e42995f6 --- /dev/null +++ b/training_server/rl/scripts/validate_upload.py @@ -0,0 +1,43 @@ +"""Resource-limited CPU single-file decoder; no network or neighboring model lookup.""" + +import contextlib +import json +import os +import resource +import sys +from pathlib import Path + +# Apply limits before importing tensor/protobuf runtimes. Thread env also set by parent. +resource.setrlimit(resource.RLIMIT_CPU, (40, 40)) +resource.setrlimit(resource.RLIMIT_AS, (8 * 1024**3, 8 * 1024**3)) +resource.setrlimit(resource.RLIMIT_FSIZE, (16 * 1024**2, 16 * 1024**2)) +resource.setrlimit(resource.RLIMIT_NOFILE, (64, 64)) +resource.setrlimit(resource.RLIMIT_CORE, (0, 0)) +os.environ["CUDA_VISIBLE_DEVICES"] = "" +ROOT = Path(__file__).resolve().parents[1] +for root in (ROOT, ROOT.parent): + sys.path.insert(0, str(root)) + +if __name__ == "__main__": + try: + payload = json.loads(sys.stdin.read(4096)) + with contextlib.redirect_stdout(sys.stderr): + import torch + from pretrained_upload import import_upload + + torch.set_num_threads(1) + manifest = import_upload( + payload["path"], payload["format"], payload["template"], payload["directory"] + ) + print(json.dumps(manifest)) + except Exception as error: + # Do not echo untrusted pickle/tensor names or absolute local paths. + from pretrained import PretrainedError + + message = ( + str(error) + if isinstance(error, PretrainedError) + else "文件解析失败;请提供受支持的自包含Go2 legacy47 actor" + ) + print(json.dumps({"error": message[:500]})) + sys.exit(1) diff --git a/training_server/rl/src/tasks/obstacle_avoidance/__init__.py b/training_server/rl/src/tasks/obstacle_avoidance/__init__.py new file mode 100644 index 00000000..02880aa4 --- /dev/null +++ b/training_server/rl/src/tasks/obstacle_avoidance/__init__.py @@ -0,0 +1,17 @@ +from mjlab.tasks.registry import register_mjlab_task +from task_config import OBSTACLE_TASK + +from src.tasks.velocity.config.go2.rl_cfg import unitree_go2_ppo_runner_cfg +from src.tasks.velocity.rl import VelocityOnPolicyRunner + +from .env_cfg import unitree_go2_obstacle_env_cfg + +_rl_cfg = unitree_go2_ppo_runner_cfg() +_rl_cfg.experiment_name = "go2_obstacle_avoidance" +register_mjlab_task( + task_id=OBSTACLE_TASK, + env_cfg=unitree_go2_obstacle_env_cfg(), + play_env_cfg=unitree_go2_obstacle_env_cfg(play=True), + rl_cfg=_rl_cfg, + runner_cls=VelocityOnPolicyRunner, +) diff --git a/training_server/rl/src/tasks/obstacle_avoidance/env_cfg.py b/training_server/rl/src/tasks/obstacle_avoidance/env_cfg.py new file mode 100644 index 00000000..7fd28023 --- /dev/null +++ b/training_server/rl/src/tasks/obstacle_avoidance/env_cfg.py @@ -0,0 +1,89 @@ +"""81/97-to-12 forward-ray navigation, retaining the Go2 50 Hz position action contract.""" + +from mjlab.managers import ObservationTermCfg, RewardTermCfg, TerminationTermCfg +from mjlab.managers.event_manager import EventTermCfg +from mjlab.sensor import ObjRef, RayCastSensorCfg +from task_config import ( + OBSTACLE_TASK, + build_terrain_layout, + navigation_candidates, + validate_task_config, +) + +from src.tasks.velocity.config.go2.env_cfgs import unitree_go2_flat_env_cfg + +from . import mdp +from .terrain import apply_terrain_configuration + + +def apply_obstacle_configuration(cfg, custom, randomize_navigation=True): + apply_terrain_configuration(cfg, custom) + sensor = custom["sensorCfg"] + cfg.scene.sensors = tuple(s for s in cfg.scene.sensors if s.name != "forward_scan") + ( + RayCastSensorCfg( + name="forward_scan", + frame=ObjRef(type="body", name="base_link", entity="robot"), + pattern=mdp.ForwardFanPatternCfg(fov=sensor["fov"], sensor_mode=sensor["sensorMode"]), + ray_alignment="base", + max_distance=sensor["maxDistance"], + include_geom_groups=(0,), + exclude_parent_body=True, + ), + ) + for group in ("actor", "critic"): + cfg.observations[group].terms["forward_depth"] = ObservationTermCfg( + func=mdp.forward_depth, params={"max_distance": sensor["maxDistance"]} + ) + cfg.observations[group].terms["target_error"] = ObservationTermCfg(func=mdp.target_error) + layout = build_terrain_layout(custom) + size = layout["size"] + candidates = navigation_candidates(layout) if randomize_navigation else None + cfg.commands["twist"] = mdp.NavigationCommandCfg( + resampling_time_range=(1e9, 1e9), + goal_offset=tuple(layout["target"][i] - layout["spawn"][i] for i in range(2)), + distance_scale=size, + **( + { + "navigation_points": tuple(tuple(point) for point in candidates["points"]), + "component_starts": tuple(candidates["componentStarts"]), + "component_counts": tuple(candidates["componentCounts"]), + "fallback_pairs": tuple( + (tuple(pair[0]), tuple(pair[1])) for pair in candidates["fallbackPairs"] + ), + "min_goal_distance": candidates["minDistance"], + } + if candidates + else {} + ), + ) + if randomize_navigation: + cfg.events["reset_base"] = EventTermCfg(func=mdp.reset_navigation_episode, mode="reset") + cfg.rewards["goal_progress"] = RewardTermCfg(func=mdp.goal_progress, weight=2.0) + cfg.rewards["obstacle_proximity"] = RewardTermCfg( + func=mdp.obstacle_proximity, + weight=-sensor["avoidanceWeight"], + params={ + "max_distance": sensor["maxDistance"], + "safety_distance": sensor["safetyDistance"], + **({"floor_size": size} if sensor["sensorMode"] == "multi_ring_raycast" else {}), + }, + ) + cfg.terminations["outside_map"] = TerminationTermCfg( + func=mdp.outside_map, params={"size": size} + ) + + +def unitree_go2_obstacle_env_cfg(play=False): + cfg = unitree_go2_flat_env_cfg(play=False) + cfg.sim.nconmax = 128 + cfg.sim.njmax = 1500 + cfg.sim.contact_sensor_maxmatch = 500 + cfg.curriculum = {} + for event in ("push_robot", "encoder_bias", "base_com"): + cfg.events.pop(event, None) + cfg.observations["actor"].enable_corruption = False + custom = validate_task_config(OBSTACLE_TASK, {}, 42) + apply_obstacle_configuration(cfg, custom) + if play: + cfg.episode_length_s = 20.0 + return cfg diff --git a/training_server/rl/src/tasks/obstacle_avoidance/mdp.py b/training_server/rl/src/tasks/obstacle_avoidance/mdp.py new file mode 100644 index 00000000..73cba966 --- /dev/null +++ b/training_server/rl/src/tasks/obstacle_avoidance/mdp.py @@ -0,0 +1,231 @@ +"""Go2 navigation observations: body-frame +X forward, +Y left, +Z up.""" + +from dataclasses import dataclass + +import torch +from mjlab.managers.command_manager import CommandTerm, CommandTermCfg +from mjlab.sensor import GridPatternCfg + + +@dataclass +class ForwardFanPatternCfg(GridPatternCfg): + """Inclusive yaw, layer-major pitch; base-aligned with one shared offset.""" + + fov: float = 90.0 + sensor_mode: str = "single_ring_raycast" + + def generate_rays(self, mj_model, device): + from task_config import sensor_pattern + + pattern = sensor_pattern(self.sensor_mode, self.fov) + yaw = torch.tensor(pattern["yawAngles"], device=device) * torch.pi / 180 + pitch = torch.tensor(pattern["pitchAngles"], device=device) * torch.pi / 180 + p, y = torch.meshgrid(pitch, yaw, indexing="ij") + directions = torch.stack((p.cos() * y.cos(), p.cos() * y.sin(), p.sin()), dim=-1).reshape( + -1, 3 + ) + offsets = torch.tensor([0.3, 0, 0.05], device=device).repeat(pattern["rayCount"], 1) + return offsets, directions + + +@dataclass(kw_only=True) +class NavigationCommandCfg(CommandTermCfg): + goal_offset: tuple[float, float] = (10.0, 0.0) + distance_scale: float = 12.0 + speed: float = 0.6 + arrival_radius: float = 0.5 + navigation_points: tuple[tuple[float, float], ...] = () + component_starts: tuple[int, ...] = () + component_counts: tuple[int, ...] = () + fallback_pairs: tuple[tuple[tuple[float, float], tuple[float, float]], ...] = () + min_goal_distance: float = 2.0 + + def build(self, env): + return NavigationCommand(self, env) + + +class NavigationCommand(CommandTerm): + cfg: NavigationCommandCfg + + def __init__(self, cfg, env): + super().__init__(cfg, env) + self.metrics["distance_to_goal"] = torch.zeros(self.num_envs, device=self.device) + self.goals_w = self._env.scene.env_origins[:, :2] + torch.tensor( + self.cfg.goal_offset, device=self.device + ) + self._points = ( + torch.tensor(self.cfg.navigation_points, dtype=torch.float32, device=self.device) + if self.cfg.navigation_points + else None + ) + self._component_starts = torch.tensor( + self.cfg.component_starts, dtype=torch.long, device=self.device + ) + self._component_counts = torch.tensor( + self.cfg.component_counts, dtype=torch.long, device=self.device + ) + self._fallback_pairs = ( + torch.tensor(self.cfg.fallback_pairs, dtype=torch.float32, device=self.device) + if self.cfg.fallback_pairs + else None + ) + + def errors(self): + robot = self._env.scene["robot"] + delta = self.goals_w - robot.data.root_link_pos_w[:, :2] + q = robot.data.root_link_quat_w + yaw = torch.atan2( + 2 * (q[:, 0] * q[:, 3] + q[:, 1] * q[:, 2]), + 1 - 2 * (q[:, 2].square() + q[:, 3].square()), + ) + heading = torch.atan2(delta[:, 1], delta[:, 0]) - yaw + heading = torch.atan2(heading.sin(), heading.cos()) + distance = delta.norm(dim=1) + heading = torch.where(distance < self.cfg.arrival_radius, 0.0, heading) + return heading, distance, delta + + @property + def command(self): + # Compute from current physics, not a one-step-old command manager cache. + heading, distance, _ = self.errors() + moving = distance >= self.cfg.arrival_radius + vx = self.cfg.speed * heading.cos().clamp(min=0) + return torch.stack( + ( + torch.where(moving, vx, 0.0), + torch.zeros_like(vx), + torch.where(moving, heading.clamp(-1, 1), 0.0), + ), + dim=1, + ) + + def _update_metrics(self): + self.metrics["distance_to_goal"][:] = self.errors()[1] + + def sample_episode(self, env_ids): + """Randomize a reachable free-space pair and the initial heading.""" + if self._points is None or len(env_ids) == 0: + return + count = len(env_ids) + component = torch.randint(len(self._component_starts), (count,), device=self.device) + starts = self._component_starts[component] + counts = self._component_counts[component] + + def indices(): + return starts + (torch.rand(count, device=self.device) * counts).long() + + spawn = self._points[indices()] + target = self._points[indices()] + valid = (target - spawn).norm(dim=1) >= self.cfg.min_goal_distance + # Keep this branch-free on CUDA; unresolved samples use a guaranteed component pair. + for _ in range(16): + candidate = self._points[indices()] + target = torch.where(valid[:, None], target, candidate) + valid = (target - spawn).norm(dim=1) >= self.cfg.min_goal_distance + fallback = self._fallback_pairs[component] + spawn = torch.where(valid[:, None], spawn, fallback[:, 0]) + target = torch.where(valid[:, None], target, fallback[:, 1]) + + robot = self._env.scene["robot"] + default = robot.data.default_root_state[env_ids].clone() + position = torch.cat((spawn, default[:, 2:3]), dim=1) + yaw = torch.rand(count, device=self.device) * (2 * torch.pi) - torch.pi + half = yaw / 2 + quaternion = torch.stack( + (half.cos(), torch.zeros_like(half), torch.zeros_like(half), half.sin()), dim=1 + ) + robot.write_root_link_pose_to_sim(torch.cat((position, quaternion), dim=1), env_ids=env_ids) + robot.write_root_link_velocity_to_sim(torch.zeros_like(default[:, 7:13]), env_ids=env_ids) + self.goals_w[env_ids] = target + + def _resample_command(self, env_ids): + # Episode pairs are sampled by the reset event before manager resets. + self.metrics["distance_to_goal"][env_ids] = self.errors()[1][env_ids] + + def _update_command(self): + # No cached command: the property is evaluated after every physics update. + self._update_metrics() + + +def reset_navigation_episode(env, env_ids): + if env_ids is None: + env_ids = torch.arange(env.num_envs, device=env.device, dtype=torch.int) + env.command_manager.get_term("twist").sample_episode(env_ids) + + +def forward_depth(env, sensor_name="forward_scan", max_distance=4.0): + distances = env.scene[sensor_name].data.distances + distances = torch.where(distances < 0, max_distance, distances) + return distances.clamp(0, max_distance) / max_distance + + +def target_error(env): + command = env.command_manager.get_term("twist") + heading, distance, _ = command.errors() + return torch.stack( + (heading / torch.pi, (distance / command.cfg.distance_scale).clamp(0, 1)), dim=1 + ) + + +def standard_floor_id(model, size): + """Fail closed: match the compiled first layout box, not a name alone.""" + import mujoco + import numpy as np + + data = mujoco.MjData(model) + mujoco.mj_forward(model, data) + ids = [ + i + for i in range(model.ngeom) + if model.geom_type[i] == 6 + and model.body_weldid[model.geom_bodyid[i]] == 0 + and model.geom_group[i] == 0 + and np.allclose(data.geom_xpos[i], [0, 0, -0.1], atol=1e-7, rtol=0) + and np.allclose(model.geom_size[i], [size / 2, size / 2, 0.1], atol=1e-7, rtol=0) + and np.allclose(data.geom_xmat[i], [1, 0, 0, 0, 1, 0, 0, 0, 1], atol=1e-7, rtol=0) + ] + if len(ids) != 1 or model.geom(ids[0]).name != "terrain_0": + raise ValueError("multi-ring reward requires one verified static standard floor") + return ids[0] + + +def floor_top_hits(data, size): + # Warp ray hits use the same world-centered coordinates in each independent env. + # Float32 tolerance: 10 micrometres; never erase 0.05m low obstacles. + hit = data.hit_pos_w + normal = data.normals_w + return ( + (data.distances >= 0) + & torch.isfinite(data.distances) + & (hit[..., :2].abs() <= size / 2 + 1e-5).all(dim=-1) + & (hit[..., 2].abs() <= 1e-5) + & (normal[..., :2].abs() <= 1e-5).all(dim=-1) + & ((normal[..., 2] - 1).abs() <= 1e-5) + ) + + +def obstacle_proximity(env, max_distance=4.0, safety_distance=0.5, floor_size=None): + depths = forward_depth(env, max_distance=max_distance) + if floor_size is not None: + # Validate once per environment/model. Observation data remains untouched. + if not hasattr(env, "_multi_ring_floor_id"): + env._multi_ring_floor_id = standard_floor_id(env.sim.mj_model, floor_size) + depths = torch.where( + floor_top_hits(env.scene["forward_scan"].data, floor_size), 1.0, depths + ) + closest = depths.amin(dim=1) * max_distance + return ((safety_distance - closest).clamp(min=0) / safety_distance).square() + + +def goal_progress(env): + command = env.command_manager.get_term("twist") + _, distance, delta = command.errors() + velocity = env.scene["robot"].data.root_link_lin_vel_w[:, :2] + progress = (velocity * delta / distance.clamp(min=1e-6).unsqueeze(1)).sum(dim=1) + return torch.where(distance >= command.cfg.arrival_radius, progress, 0.0) + + +def outside_map(env, size=12.0): + # Every Warp environment is an independent copy of the same world-centered map. + position = env.scene["robot"].data.root_link_pos_w[:, :2] + return (position.abs() > size / 2 - 0.3).any(dim=1) diff --git a/training_server/rl/src/tasks/obstacle_avoidance/terrain.py b/training_server/rl/src/tasks/obstacle_avoidance/terrain.py new file mode 100644 index 00000000..7acd2143 --- /dev/null +++ b/training_server/rl/src/tasks/obstacle_avoidance/terrain.py @@ -0,0 +1,70 @@ +"""Install the exact bounded boxes-v1 deployment layout in mjlab's generator.""" + +from dataclasses import dataclass, field + +import mujoco +import numpy as np +from mjlab.envs import mdp +from mjlab.managers.event_manager import EventTermCfg +from mjlab.terrains import TerrainEntityCfg, TerrainGeneratorCfg +from mjlab.terrains.terrain_generator import SubTerrainCfg, TerrainGeometry, TerrainOutput +from task_config import build_terrain_layout + + +@dataclass +class DeploymentTerrainCfg(SubTerrainCfg): + layout: dict = field(default_factory=dict) + + def function(self, difficulty, spec, rng): + # Generator translates patch-local corner coordinates by (-size/2,-size/2). + half = self.layout["size"] / 2 + translation = np.array([half, half, 0]) + body = spec.body("terrain") + geometries = [] + for box in self.layout["boxes"]: + geom = body.add_geom( + type=mujoco.mjtGeom.mjGEOM_BOX, + pos=np.array(box["pos"]) + translation, + size=box["size"], + group=0, + contype=1, + conaffinity=1, + friction=[self.layout["friction"], 0.005, 0.0001], + priority=1, + condim=3, + ) + geometries.append(TerrainGeometry(geom=geom)) + spawn = np.array(self.layout["spawn"]) + spawn[2] = 0 # Robot init_state contributes the 0.32 m body height. + return TerrainOutput(origin=spawn + translation, geometries=geometries) + + +def apply_terrain_configuration(cfg, custom): + layout = build_terrain_layout(custom) + size = layout["size"] + cfg.scene.entities["robot"].init_state.pos = (0.0, 0.0, layout["spawn"][2]) + cfg.scene.entities["robot"].init_state.rot = tuple(layout["spawnQuaternion"]) + cfg.scene.terrain = TerrainEntityCfg( + terrain_type="generator", + max_init_terrain_level=0, + terrain_generator=TerrainGeneratorCfg( + seed=custom["seed"], + size=(size, size), + num_rows=1, + num_cols=1, + border_width=0, + curriculum=False, + color_scheme="none", + sub_terrains={"deployment": DeploymentTerrainCfg(layout=layout)}, + ), + ) + cfg.curriculum.pop("terrain_levels", None) + cfg.events.pop("randomize_terrain", None) + # Install the deterministic reset first; the navigation task may replace it with + # collision-free random point-goal sampling after the command term is configured. + cfg.events["reset_base"] = EventTermCfg( + func=mdp.reset_root_state_uniform, + mode="reset", + params={"pose_range": {}, "velocity_range": {}}, + ) + cfg.events["foot_friction"].params["ranges"] = (layout["friction"], layout["friction"]) diff --git a/training_server/rl/src/tasks/velocity/rl/runner.py b/training_server/rl/src/tasks/velocity/rl/runner.py index 4900af64..fc7b7d6a 100644 --- a/training_server/rl/src/tasks/velocity/rl/runner.py +++ b/training_server/rl/src/tasks/velocity/rl/runner.py @@ -1,4 +1,6 @@ +import json import os +from pathlib import Path import wandb @@ -23,6 +25,19 @@ class VelocityOnPolicyRunner(MjlabOnPolicyRunner): ) # type: ignore[assignment] onnx_path = os.path.join(policy_path, filename) metadata = get_base_metadata(self.env.unwrapped, run_name) + deployment = getattr(self.env.unwrapped, "platform_deployment", None) + if deployment is not None: + metadata["platform_deployment"] = json.dumps(deployment, separators=(",", ":")) + Path(policy_path, "deployment.json").write_text( + json.dumps(deployment, ensure_ascii=False, indent=2) + "\n", encoding="utf-8" + ) + initialization = getattr(self.env.unwrapped, "platform_initialization", None) + if initialization is not None: + # Only content identity and protocol, never local source paths or arbitrary infos. + from pretrained import public_initialization_metadata + metadata["pretrained_initialization"] = json.dumps( + public_initialization_metadata(initialization), separators=(",", ":") + ) attach_metadata_to_onnx(onnx_path, metadata) if self.logger.logger_type in ["wandb"]: wandb.save(policy_path + filename, base_path=os.path.dirname(policy_path)) diff --git a/training_server/server.py b/training_server/server.py index f12524cd..642f0613 100644 --- a/training_server/server.py +++ b/training_server/server.py @@ -25,14 +25,23 @@ from pathlib import Path from typing import Any from urllib.parse import parse_qs, unquote, urlsplit +from pretrained_sources import PretrainedSources, SourceError +from task_config import ( + OBSTACLE_TASK, + TaskConfigError, + deployment_metadata, + task_metadata, + validate_task_config, +) from tuning.manager import TuningError, TuningManager from tuning.process import GpuLease, ResourceBusyError -from tuning.schema import RewardConfigError +from tuning.schema import RewardConfigError, validate_configuration from tuning.scoring import EvaluationError VERSION = "0.4.0" -# 浏览器当前 ONNX 运行时只实现 Go2 的 47→12 部署契约;其他任务须由服务启动参数显式放行。 -DEFAULT_TASKS = ("Unitree-Go2-Flat",) +MAX_REQUEST_BYTES = 128 * 1024 # Bounded full boxes-v1 payload (<=257 boxes). +# Rough 可以训练,但其高度扫描 actor 不允许冒充浏览器 Flat 部署。 +DEFAULT_TASKS = ("Unitree-Go2-Flat", "Unitree-Go2-Rough", OBSTACLE_TASK) ACTIVE_STATES = {"queued", "running"} MAX_JOBS = 20 ANSI_ESCAPE = re.compile(r"\x1b\[[0-?]*[ -/]*[@-~]") @@ -71,6 +80,12 @@ class TrainingConfig: gpu_ids: list[int] wandb_mode: str reward_config: dict[str, Any] | None = None + terrain_preset: str | None = None + terrain_params: dict[str, Any] = field(default_factory=dict) + sensor_cfg: dict[str, Any] | None = None + task_config: dict[str, Any] | None = None + deployment: dict[str, Any] = field(default_factory=dict) + pretrained: dict[str, Any] | None = None @dataclass @@ -96,6 +111,8 @@ class TrainingJob: "id": self.id, "state": self.state, "taskId": self.config.task_id, + "deployment": self.config.deployment, + "pretrained": self.config.pretrained, "createdAt": self.created_at, "startedAt": self.started_at, "endedAt": self.ended_at, @@ -117,6 +134,7 @@ class TrainingManager: tasks: tuple[str, ...], check_environment: bool = True, lease: GpuLease | None = None, + sources: PretrainedSources | None = None, ): self.trainer_root = trainer_root.expanduser().resolve() self.python = str(Path(python).expanduser()) if os.sep in python else python @@ -127,6 +145,7 @@ class TrainingManager: self._environment_error: str | None | bool = False self.lease = lease or GpuLease() self.preset_resolver: Any = None + self.sources = sources def readiness_error(self) -> str | None: if not self.trainer_root.is_dir(): @@ -172,6 +191,13 @@ class TrainingManager: "trainerRoot": str(self.trainer_root), "python": self.python, "tasks": list(self.tasks), + "pretrainedSources": self.sources.catalog() if self.sources else [], + "pretrainedUpload": { + "enabled": self.sources is not None, "templateId": "go2-legacy47-v1", + "formats": {"pt": 256 * 1024**2, "onnx": 64 * 1024**2}, + "endpoint": "/api/training/pretrained-sources/upload", + }, + "taskMetadata": task_metadata(self.tasks), "activeJobId": self.active_job_id(), "error": error, } @@ -179,6 +205,25 @@ class TrainingManager: def parse_config(self, payload: Any) -> TrainingConfig: if not isinstance(payload, dict): raise ApiError(HTTPStatus.BAD_REQUEST, "请求体必须是 JSON 对象") + allowed = { + "taskId", + "numEnvs", + "maxIterations", + "seed", + "runName", + "device", + "gpuIds", + "wandbMode", + "rewardPresetId", + "pretrainedSourceId", + "terrainPreset", + "terrainParams", + "sensorCfg", + "sensorType", + "customTerrainBoxes", + } + if payload.keys() - allowed: + raise ApiError(HTTPStatus.BAD_REQUEST, "请求包含未知字段(不接受配置路径/MJCF)") task_id = payload.get("taskId") if task_id not in self.tasks: raise ApiError(HTTPStatus.BAD_REQUEST, f"不允许的训练任务:{task_id}") @@ -216,19 +261,46 @@ class TrainingManager: preset_id = payload.get("rewardPresetId") reward_config = None if preset_id is not None: + if task_id != "Unitree-Go2-Flat": + raise ApiError(HTTPStatus.BAD_REQUEST, "自动调参奖励 preset 仅支持 Flat 任务") if not isinstance(preset_id, str) or not re.fullmatch(r"[0-9a-f]{32}", preset_id): raise ApiError(HTTPStatus.BAD_REQUEST, "rewardPresetId 格式无效") if self.preset_resolver is None: raise ApiError(HTTPStatus.BAD_REQUEST, "奖励 preset 服务未就绪") try: - reward_config = self.preset_resolver(preset_id) + reward_config = validate_configuration( + self.preset_resolver(preset_id, task_id), task_id + ) except KeyError as error: raise ApiError(HTTPStatus.BAD_REQUEST, "奖励 preset 不存在") from error + except RewardConfigError as error: + raise ApiError(HTTPStatus.BAD_REQUEST, str(error)) from error + seed = integer("seed", 0, 2_147_483_647) + try: + custom = validate_task_config(task_id, payload, seed) + except TaskConfigError as error: + raise ApiError(HTTPStatus.BAD_REQUEST, str(error)) from error + pretrained = None + if "pretrainedSourceId" in payload: + if self.sources is None: + raise ApiError( + HTTPStatus.BAD_REQUEST, "服务尚未注册基础策略,请配置--pretrained-sources" + ) + try: + pretrained = self.sources.bind(payload["pretrainedSourceId"], task_id, custom) + except SourceError as error: + raise ApiError(HTTPStatus.BAD_REQUEST, str(error)) from error return TrainingConfig( + pretrained=pretrained, task_id=task_id, + terrain_preset=custom["terrainPreset"] if custom else None, + terrain_params=custom["terrainParams"] if custom else {}, + sensor_cfg=custom["sensorCfg"] if custom else None, + task_config=custom, + deployment=deployment_metadata(task_id, custom, seed), num_envs=integer("numEnvs", 1, 16384), max_iterations=integer("maxIterations", 1, 1_000_000), - seed=integer("seed", 0, 2_147_483_647), + seed=seed, run_name=run_name, device=device, gpu_ids=raw_gpu_ids, @@ -325,7 +397,9 @@ class TrainingManager: os.killpg(process.pid, signal.SIGKILL) process.wait(timeout=2) - def command_for(self, config: TrainingConfig) -> list[str]: + def command_for( + self, config: TrainingConfig, task_config_path: Path | None = None + ) -> list[str]: command = [ self.python, "-u", @@ -336,6 +410,12 @@ class TrainingManager: f"--agent.seed={config.seed}", f"--agent.run-name={config.run_name}", ] + if task_config_path is not None: + command.extend(("--task-config", str(task_config_path))) + if config.pretrained is not None: + if self.sources is None: + raise SourceError("基础策略快照服务未配置,拒绝随机初始化") + command.extend(self.sources.arguments(config.pretrained)) if config.reward_config is not None: command.extend( ( @@ -387,13 +467,21 @@ class TrainingManager: def _run(self, job: TrainingJob) -> None: before = self._artifact_snapshot() - command = self.command_for(job.config) environment = os.environ.copy() # 默认离线记录,保留本地 W&B 指标但不要求 API Key;只有前端明确选择 # online 时才允许 wandb 发起登录/联网。 environment["WANDB_MODE"] = job.config.wandb_mode environment["WANDB_SILENT"] = "true" try: + config_path = None + if job.config.task_config is not None: + job_dir = self.trainer_root / "logs" / "rsl_rl" / "web_jobs" / job.id + job_dir.mkdir(parents=True, exist_ok=True) + config_path = job_dir / "training_config.json" + config_path.write_text(json.dumps(job.config.task_config), encoding="utf-8") + command = self.command_for(job.config, config_path) + if config_path is not None: + command.extend(("--output-dir", str(config_path.parent))) # Popen 与 process 登记必须和取消检查处于同一个临界区:cancel() 要么在 # 创建前标记取消,要么在创建后取得进程并终止,不能落入二者之间。 with self.lock: @@ -500,7 +588,7 @@ class TrainingRequestHandler(BaseHTTPRequestHandler): self._json(HTTPStatus.NOT_FOUND, {"error": "调参 session、trial 或 proposal 不存在"}) elif isinstance(error, ResourceBusyError): self._json(HTTPStatus.CONFLICT, {"error": str(error)}) - elif isinstance(error, (TuningError, RewardConfigError, EvaluationError)): + elif isinstance(error, (TuningError, RewardConfigError, EvaluationError, SourceError)): self._json(HTTPStatus.BAD_REQUEST, {"error": str(error)}) else: self._json( @@ -523,15 +611,47 @@ class TrainingRequestHandler(BaseHTTPRequestHandler): length = int(self.headers.get("Content-Length", "0")) except ValueError as error: raise ApiError(HTTPStatus.BAD_REQUEST, "Content-Length 无效") from error - if length <= 0 or length > 32 * 1024: + if length <= 0 or length > MAX_REQUEST_BYTES: raise ApiError( - HTTPStatus.REQUEST_ENTITY_TOO_LARGE, "训练请求体不能为空且不能超过 32 KiB" + HTTPStatus.REQUEST_ENTITY_TOO_LARGE, "训练请求体不能为空且不能超过 128 KiB" ) try: return json.loads(self.rfile.read(length)) except (UnicodeDecodeError, json.JSONDecodeError) as error: raise ApiError(HTTPStatus.BAD_REQUEST, "训练请求不是有效 JSON") from error + def _upload(self): + # This dedicated binary route must never pass through the 128KiB JSON reader. + self.close_connection = True # no unread extra bytes may become a second request + if self.manager.sources is None: + raise ApiError(HTTPStatus.SERVICE_UNAVAILABLE, "基础策略上传服务不可用") + if self.headers.get("Transfer-Encoding") or self.headers.get("Content-Encoding"): + raise ApiError(HTTPStatus.BAD_REQUEST, "上传不接受chunked/压缩编码") + lengths = self.headers.get_all("Content-Length", []) + if len(lengths) != 1 or not re.fullmatch(r"[0-9]{1,10}", lengths[0]): + raise ApiError(HTTPStatus.BAD_REQUEST, "上传必须提供唯一Content-Length") + length = int(lengths[0]) + if self.headers.get("Content-Type") != "application/octet-stream": + raise ApiError(HTTPStatus.BAD_REQUEST, "上传Content-Type必须是application/octet-stream") + query = parse_qs(urlsplit(self.path).query, keep_blank_values=True) + if set(query) != {"format", "template", "name"} or any(len(v) != 1 for v in query.values()): + raise ApiError(HTTPStatus.BAD_REQUEST, "上传必须显式提供format/template/name") + fmt, template, name = (query[key][0] for key in ("format", "template", "name")) + limit = {"pt": 256 * 1024**2, "onnx": 64 * 1024**2}.get(fmt) + if limit is not None and (length == 0 or length > limit): + raise ApiError(HTTPStatus.REQUEST_ENTITY_TOO_LARGE, "上传为空或超过格式大小上限") + if len(name) > 255: + raise ApiError(HTTPStatus.BAD_REQUEST, "上传显示名称过长") + try: + result = self.manager.sources.receive_upload( + self.rfile, length, fmt, template, name, set_timeout=self.connection.settimeout, + ) + except OSError as error: + raise ApiError( + HTTPStatus.SERVICE_UNAVAILABLE, "上传存储写入失败,请检查本机磁盘空间/权限" + ) from error + self._json(HTTPStatus.CREATED, result) + @staticmethod def _route(path: str) -> tuple[str | None, bool]: match = re.fullmatch(r"/api/training/jobs/([0-9a-f]{32})(/artifacts/policy\.onnx)?", path) @@ -628,6 +748,9 @@ class TrainingRequestHandler(BaseHTTPRequestHandler): try: self._ensure_request() path = urlsplit(self.path).path + if path == "/api/training/pretrained-sources/upload": + self._upload() + return if path == "/api/training/jobs": self._json(HTTPStatus.ACCEPTED, self.manager.start(self._payload())) return @@ -742,6 +865,12 @@ def parse_args() -> argparse.Namespace: default=None, help="调参 SQLite 与 trial 产物目录;默认位于训练工程 logs/auto_tuning", ) + parser.add_argument( + "--pretrained-sources", + type=Path, + default=None, + help="可选旧管理员基础策略注册JSON;不配置也可直接上传单个.pt/.onnx", + ) parser.add_argument( "--task", action="append", dest="tasks", help="允许前端启动的任务 ID;可重复" ) @@ -762,14 +891,23 @@ def main() -> None: if len(token) < 16: raise SystemExit("训练服务访问令牌至少需要 16 个字符") lease = GpuLease() + tuning_root = args.tuning_data_root or (Path(args.trainer_root) / "logs" / "auto_tuning") + sources = PretrainedSources( + args.pretrained_sources, + tuning_root / "pretrained_sources", + args.trainer_python, + args.trainer_root, + ) manager = TrainingManager( args.trainer_root, args.trainer_python, tuple(args.tasks or DEFAULT_TASKS), lease=lease, + sources=sources, + ) + tuning_manager = TuningManager( + args.trainer_root, args.trainer_python, tuning_root, lease, sources=sources ) - tuning_root = args.tuning_data_root or (Path(args.trainer_root) / "logs" / "auto_tuning") - tuning_manager = TuningManager(args.trainer_root, args.trainer_python, tuning_root, lease) manager.preset_resolver = tuning_manager.preset_config TrainingRequestHandler.manager = manager TrainingRequestHandler.tuning_manager = tuning_manager diff --git a/training_server/task_config.py b/training_server/task_config.py new file mode 100644 index 00000000..9043c1da --- /dev/null +++ b/training_server/task_config.py @@ -0,0 +1,479 @@ +"""Validated, dependency-free browser training/deployment contract (version 1).""" + +from __future__ import annotations + +import copy +import math +import random +from typing import Any + +FLAT_TASK = "Unitree-Go2-Flat" +ROUGH_TASK = "Unitree-Go2-Rough" +OBSTACLE_TASK = "Unitree-Go2-ObstacleAvoidance" +TERRAIN_PRESETS = ("plane", "discrete_obstacles", "rough", "pyramid_stairs", "wave", "custom_boxes") +# Range metadata is also the sole validation source for browser configuration. +TERRAIN_PARAMETERS = { + "size": {"min": 8, "max": 24, "default": 12}, + "obstacle_count": {"min": 1, "max": 100, "default": 24, "integer": True}, + "obstacle_height_min": {"min": 0.05, "max": 1.5, "default": 0.2}, + "obstacle_height_max": {"min": 0.05, "max": 1.5, "default": 0.6}, + "spacing": {"min": 0.6, "max": 3, "default": 1.2}, + "friction": {"min": 0.2, "max": 2, "default": 0.8}, + "roughness": {"min": 0.01, "max": 0.2, "default": 0.06}, + "step_height": {"min": 0.03, "max": 0.2, "default": 0.08}, + "wave_amplitude": {"min": 0.01, "max": 0.2, "default": 0.08}, +} +SENSOR_PARAMETERS = { + "fov": {"min": 30, "max": 120, "default": 90}, + "maxDistance": {"min": 1, "max": 5, "default": 4}, + "safetyDistance": {"min": 0.1, "max": 1, "default": 0.5}, + "avoidanceWeight": {"min": 0, "max": 10, "default": 2}, +} +JOINT_NAMES = [ + f"{leg}_{joint}_joint" for leg in ("FL", "FR", "RL", "RR") for joint in ("hip", "thigh", "calf") +] +DEFAULT_JOINT_POSITION = [-0.1, 0.9, -1.8, 0.1, 0.9, -1.8] * 2 + + +class TaskConfigError(ValueError): + pass + + +def _parameters(value: Any, schema: dict, label: str) -> dict: + if not isinstance(value, dict) or value.keys() - schema.keys(): + raise TaskConfigError(f"{label} 包含未知参数或不是对象") + result = {} + for name, bounds in schema.items(): + number = value.get(name, bounds["default"]) + if ( + isinstance(number, bool) + or not isinstance(number, (int, float)) + or not bounds["min"] <= number <= bounds["max"] + or not math.isfinite(number) + or (bounds.get("integer") and not isinstance(number, int)) + ): + raise TaskConfigError(f"{label}.{name} 超出允许的有限数值范围") + result[name] = number + return result + + +def sensor_pattern(mode: str, fov: float) -> dict: + if mode not in ("single_ring_raycast", "multi_ring_raycast"): + raise TaskConfigError("不支持的 sensorMode") + multi = mode == "multi_ring_raycast" + count = 16 if multi else 32 + return { + "sensorMode": mode, + "rayCount": 48 if multi else 32, + "pitchAngles": [0, -20, -45] if multi else [0], + "yawCount": count, + "yawAngles": [-fov / 2 + i * fov / (count - 1) for i in range(count)], + "angleUnit": "deg", + "rayOrder": "layer-major", + } + + +def validate_sensor_config(raw: dict) -> dict: + pattern_keys = set(sensor_pattern("single_ring_raycast", 90)) + sensor = _parameters( + {k: v for k, v in raw.items() if k not in pattern_keys | {"type"}}, + SENSOR_PARAMETERS, + "sensorCfg", + ) + pattern = sensor_pattern(raw.get("sensorMode", "single_ring_raycast"), sensor["fov"]) + for key, expected in pattern.items(): + if key not in raw: + continue + actual = raw[key] + if isinstance(expected, list): + valid = ( + isinstance(actual, list) + and len(actual) == len(expected) + and all( + not isinstance(a, bool) + and isinstance(a, (int, float)) + and math.isfinite(a) + and abs(a - b) <= 1e-10 + for a, b in zip(actual, expected, strict=True) + ) + ) + else: + valid = type(actual) is type(expected) and actual == expected + if not valid: + raise TaskConfigError(f"sensorCfg.{key} 与mode/FOV矛盾") + return {**sensor, "type": "raycast", **pattern} + + +def validate_custom_terrain(value: Any) -> dict: + """Untrusted full layout: exact schema, no RNG, no geometry repair or clipping.""" + fields = { + "representation", + "approximation", + "size", + "friction", + "boxes", + "spawn", + "spawnQuaternion", + "target", + "actualObstacleCount", + } + if not isinstance(value, dict) or set(value) != fields: + raise TaskConfigError("customTerrainBoxes 字段缺失或未知(不接受路径/MJCF)") + + def number(v, lo, hi): + if ( + isinstance(v, bool) + or not isinstance(v, (int, float)) + or not lo <= v <= hi + or not math.isfinite(v) + ): + raise TaskConfigError("customTerrainBoxes 必须使用范围内有限数值") + return v + + def vector(v, n, lo=-12, hi=12): + if not isinstance(v, list) or len(v) != n: + raise TaskConfigError("customTerrainBoxes 向量长度无效") + return [number(x, lo, hi) for x in v] + + size = number(value["size"], 8, 24) + number(value["friction"], 0.2, 2) + if value["representation"] != "boxes-v1" or value["approximation"] is not True: + raise TaskConfigError( + "custom_boxes 必须声明 boxes-v1 和 approximation=true(AABB/底板标准化)" + ) + boxes = value["boxes"] + if not isinstance(boxes, list) or not 1 <= len(boxes) <= 257: + raise TaskConfigError("customTerrainBoxes 只能包含底板及最多256障碍") + for box in boxes: + if not isinstance(box, dict) or set(box) != {"pos", "size", "yaw"}: + raise TaskConfigError("box 字段缺失或未知") + center = vector(box["pos"], 3) + half = vector(box["size"], 3, 0, 12) + number(box["yaw"], 0, 0) + if any(x <= 0 for x in half): + raise TaskConfigError("box 半尺寸必须严格大于零") + if any(abs(center[i]) + half[i] > size / 2 + 1e-6 for i in range(2)): + raise TaskConfigError("box 超出世界地图边界") + if center[2] - half[2] < -0.2 - 1e-6 or center[2] + half[2] > 12: + raise TaskConfigError("box 高度超出边界") + if boxes[0] != {"pos": [0, 0, -0.1], "size": [size / 2, size / 2, 0.1], "yaw": 0}: + raise TaskConfigError("必须使用标准 floor z=[-0.2,0]") + count = value["actualObstacleCount"] + if isinstance(count, bool) or not isinstance(count, int) or count != len(boxes) - 1: + raise TaskConfigError("actualObstacleCount 与布局不一致") + spawn = vector(value["spawn"], 3) + target = vector(value["target"], 2) + if spawn[2] != 0.32: + raise TaskConfigError("出生高度必须为0.32") + quaternion = vector(value["spawnQuaternion"], 4, -1, 1) + if abs(sum(x * x for x in quaternion) - 1) > 1e-6: + raise TaskConfigError("出生四元数必须归一化") + for point in (spawn, target): + if any(abs(point[i]) > size / 2 - 0.5 for i in range(2)): + raise TaskConfigError("起终点0.5m安全区超出地图") + for box in boxes[1:]: + distance_sq = sum( + max(abs(point[i] - box["pos"][i]) - box["size"][i], 0) ** 2 for i in range(2) + ) + if distance_sq <= 0.5**2: + raise TaskConfigError("障碍物侵占起终点0.5m圆形安全区;请修改坐标,不会清除障碍") + return copy.deepcopy(value) + + +def validate_task_config(task_id: str, payload: dict, seed: int) -> dict | None: + """Validate preset parameters or an authoritative boxes-v1 layout, never paths/XML.""" + custom_keys = { + "terrainPreset", + "terrainParams", + "sensorCfg", + "sensorType", + "customTerrainBoxes", + } + if task_id != OBSTACLE_TASK and not custom_keys.intersection(payload): + return None + if task_id not in (FLAT_TASK, ROUGH_TASK, OBSTACLE_TASK): + raise TaskConfigError("该任务不支持自定义地形") + preset = payload.get( + "terrainPreset", "discrete_obstacles" if task_id == OBSTACLE_TASK else "plane" + ) + if not isinstance(preset, str) or preset not in TERRAIN_PRESETS: + raise TaskConfigError("不支持的 terrainPreset") + layout = None + if preset == "custom_boxes": + layout = validate_custom_terrain(payload.get("customTerrainBoxes")) + expected = {"size": layout["size"], "friction": layout["friction"]} + if "terrainParams" in payload and ( + not isinstance(payload["terrainParams"], dict) + or payload["terrainParams"] != expected + or any( + isinstance(v, bool) or not isinstance(v, (int, float)) + for v in payload["terrainParams"].values() + ) + ): + raise TaskConfigError("custom_boxes terrainParams 必须与布局size/friction完全一致") + terrain = expected + else: + if "customTerrainBoxes" in payload: + raise TaskConfigError("customTerrainBoxes 仅允许 custom_boxes,不能降级预设") + terrain = _parameters(payload.get("terrainParams", {}), TERRAIN_PARAMETERS, "terrainParams") + if not layout and terrain["obstacle_height_min"] > terrain["obstacle_height_max"]: + raise TaskConfigError("障碍物最小高度不能超过最大高度") + raw_sensor = payload.get("sensorCfg", {}) + if not isinstance(raw_sensor, dict): + raise TaskConfigError("sensorCfg 必须是对象") + sensor_type = payload.get("sensorType", raw_sensor.get("type", "raycast")) + if sensor_type != "raycast" or raw_sensor.get("type", "raycast") != "raycast": + raise TaskConfigError("首版只支持 raycast,未实现 camera_depth") + if task_id != OBSTACLE_TASK and (raw_sensor or "sensorType" in payload): + raise TaskConfigError("只有避障任务支持 sensorCfg") + sensor = None + if task_id == OBSTACLE_TASK: + sensor = validate_sensor_config(raw_sensor) + sensor["type"] = "raycast" + if sensor["safetyDistance"] >= sensor["maxDistance"]: + raise TaskConfigError("安全距离必须小于探测距离") + return { + "terrainPreset": preset, + "terrainParams": terrain, + "sensorCfg": sensor, + "seed": seed, + **({"customTerrainBoxes": layout} if layout else {}), + } + + +def navigation_candidates( + layout: dict, spacing: float = 0.5, clearance: float = 0.55, min_distance: float = 2.0 +) -> dict: + """Build connected, collision-free point-goal samples for episode resets. + + Obstacles are inflated by ``clearance`` and the map is sampled on a regular grid. + Only components containing a pair at least ``min_distance`` apart are retained, so + runtime sampling cannot place a goal across an impassable wall. + """ + size = layout["size"] + lo, hi = -size / 2 + clearance, size / 2 - clearance + count = max(1, int(math.floor((hi - lo) / spacing)) + 1) + axis = [lo + i * (hi - lo) / max(count - 1, 1) for i in range(count)] + obstacles = layout["boxes"][1:] + + def free(x: float, y: float) -> bool: + return all( + math.hypot( + max(abs(x - box["pos"][0]) - box["size"][0], 0), + max(abs(y - box["pos"][1]) - box["size"][1], 0), + ) + > clearance + for box in obstacles + ) + + cells = {(i, j) for i, x in enumerate(axis) for j, y in enumerate(axis) if free(x, y)} + components = [] + while cells: + pending = [cells.pop()] + component = [] + while pending: + cell = pending.pop() + component.append(cell) + i, j = cell + for neighbour in ((i - 1, j), (i + 1, j), (i, j - 1), (i, j + 1)): + if neighbour in cells: + cells.remove(neighbour) + pending.append(neighbour) + points = [(axis[i], axis[j]) for i, j in sorted(component)] + extrema = [ + ( + min(points, key=lambda point: point[axis_index]), + max(points, key=lambda point: point[axis_index]), + ) + for axis_index in (0, 1) + ] + first, second = max(extrema, key=lambda pair: math.dist(*pair)) + if math.dist(first, second) >= min_distance: + components.append((points, first, second)) + if not components: + raise TaskConfigError( + f"地图没有可用于随机起终点的连通自由区域(至少需要{min_distance:g}m间距)" + ) + flattened = [] + starts = [] + counts = [] + fallbacks = [] + for points, first, second in components: + starts.append(len(flattened)) + counts.append(len(points)) + flattened.extend(points) + fallbacks.append((first, second)) + return { + "points": flattened, + "componentStarts": starts, + "componentCounts": counts, + "fallbackPairs": fallbacks, + "minDistance": min_distance, + "clearance": clearance, + "spacing": spacing, + } + + +def build_terrain_layout(config: dict) -> dict: + """Emit exact world-centered boxes; consumers do not reimplement RNG/terrain presets. + + size values are MuJoCo half-extents, pos is the center, yaw is always zero. + TerrainGenerator uses one patch, shifted back to these coordinates. All parallel + environments are independent worlds with the same map and spawn, not a grid. + """ + if config["terrainPreset"] == "custom_boxes": + return validate_custom_terrain(config.get("customTerrainBoxes")) + p = config["terrainParams"] + size = p["size"] + rng = random.Random(config["seed"]) + boxes = [{"pos": [0, 0, -0.1], "size": [size / 2, size / 2, 0.1], "yaw": 0}] + preset = config["terrainPreset"] + + def box(x: float, y: float, sx: float, sy: float, height: float) -> None: + boxes.append({"pos": [x, y, height / 2], "size": [sx, sy, height / 2], "yaw": 0}) + + # Leave two flat end strips for the spawn and goal; no rejection sampling. + if preset == "discrete_obstacles": + spacing = p["spacing"] + nx = max(1, int((size - 4) / spacing)) + ny = max(1, int((size - 2) / spacing)) + cells = [ + ( + -size / 2 + 2 + (i + 0.5) * (size - 4) / nx, + -size / 2 + 1 + (j + 0.5) * (size - 2) / ny, + ) + for i in range(nx) + for j in range(ny) + ] + rng.shuffle(cells) + for x, y in cells[: p["obstacle_count"]]: + box(x, y, 0.2, 0.2, rng.uniform(p["obstacle_height_min"], p["obstacle_height_max"])) + elif preset in ("rough", "wave"): + # 16x16 bounded box approximation, deliberately not the editor heightfield. + nx = ny = 16 + sx, sy = (size - 4) / nx, size / ny + for i in range(nx): + for j in range(ny): + height = ( + rng.uniform(0.005, p["roughness"]) + if preset == "rough" + else 0.005 + p["wave_amplitude"] * (1 + math.sin(i * math.pi / 4)) / 2 + ) + box( + -size / 2 + 2 + (i + 0.5) * sx, + -size / 2 + (j + 0.5) * sy, + sx / 2, + sy / 2, + height, + ) + elif preset == "pyramid_stairs": + for i in range(4): + half = (size - 4) / 2 - i * (size - 4) / 10 + box(0, 0, half, half, p["step_height"] * (i + 1)) + return { + "representation": "boxes-v1", + "approximation": preset in ("rough", "wave", "pyramid_stairs"), + "size": size, + "friction": p["friction"], + "boxes": boxes, + "spawn": [-size / 2 + 1, 0, 0.32], + "spawnQuaternion": [1, 0, 0, 0], + "target": [size / 2 - 1, 0], + "actualObstacleCount": len(boxes) - 1, + } + + +def deployment_metadata(task_id: str, config: dict | None, seed: int) -> dict: + obstacle = task_id == OBSTACLE_TASK + metadata = { + "version": 1, + "taskId": task_id, + "browserCompatible": task_id in (FLAT_TASK, OBSTACLE_TASK), + "observationSize": (49 + config["sensorCfg"]["rayCount"]) + if obstacle + else (234 if task_id == ROUGH_TASK else 47), + "actionSize": 12, + "controlHz": 50, + "gaitPeriod": 0.6, + "jointNames": JOINT_NAMES, + "defaultJointPosition": DEFAULT_JOINT_POSITION, + "actionScale": [0.25] * 12, + "stiffness": [20, 20, 40] * 4, + "damping": [1, 1, 2] * 4, + "effortLimits": [23.5, 23.5, 45] * 4, + "observationTerms": [ + "base_ang_vel", + "projected_gravity", + "command", + "phase", + "joint_pos", + "joint_vel", + "actions", + ] + + ( + ["forward_depth", "target_error"] + if obstacle + else (["height_scan"] if task_id == ROUGH_TASK else []) + ), + "seed": seed, + } + if task_id == ROUGH_TASK: + metadata["incompatibilityReason"] = "旧 Rough actor 含向下高度扫描;浏览器未实现此部署契约" + if config is not None: + metadata.update( + { + "terrainPreset": config["terrainPreset"], + "terrainParams": config["terrainParams"], + "terrain": build_terrain_layout(config), + } + ) + if obstacle: + assert config is not None + metadata["sensorCfg"] = { + **config["sensorCfg"], + "offset": [0.3, 0, 0.05], + "alignment": "base", + "terrainOnly": True, + "includeGround": True, + } + metadata["navigation"] = { + "speed": 0.6, + "arrivalRadius": 0.5, + "distanceScale": config["terrainParams"]["size"], + "headingScale": math.pi, + "yawGain": 1.0, + "maxYawRate": 1.0, + "episodeSeconds": 20, + "onArrival": "stop", + "onReset": "respawn", + "trainingReset": "random-connected-free-pair", + "trainingMinGoalDistance": 2.0, + "trainingClearance": 0.55, + } + return metadata + + +def task_metadata(tasks: tuple[str, ...]) -> list[dict]: + names = { + FLAT_TASK: "平地速度控制", + ROUGH_TASK: "地形自适应(仅后端)", + OBSTACLE_TASK: "前视射线避障导航", + } + return [ + { + "id": task, + "name": names.get(task, task), + "browserCompatible": task in (FLAT_TASK, OBSTACLE_TASK), + "terrainPresets": list(TERRAIN_PRESETS) if task in names else [], + "terrainParameters": TERRAIN_PARAMETERS if task in names else {}, + "sensorTypes": ["raycast"] if task == OBSTACLE_TASK else [], + "sensorModes": ["single_ring_raycast", "multi_ring_raycast"] + if task == OBSTACLE_TASK + else [], + "sensorParameters": SENSOR_PARAMETERS if task == OBSTACLE_TASK else {}, + "mapSyncScope": ( + "已应用静态碰撞场景→custom_boxes(boxes-v1);AABB与标准底板近似;不支持mesh/hfield" + ), + } + for task in tasks + ] diff --git a/training_server/tests/fixtures/custom-boxes.json b/training_server/tests/fixtures/custom-boxes.json new file mode 100644 index 00000000..ff797614 --- /dev/null +++ b/training_server/tests/fixtures/custom-boxes.json @@ -0,0 +1,14 @@ +{ + "representation": "boxes-v1", + "approximation": true, + "size": 12, + "friction": 0.8, + "boxes": [ + { "pos": [0, 0, -0.1], "size": [6, 6, 0.1], "yaw": 0 }, + { "pos": [1, 2, 0.5], "size": [0.4, 0.3, 0.5], "yaw": 0 } + ], + "spawn": [-2, -1, 0.32], + "spawnQuaternion": [0.7071067811865476, 0, 0, 0.7071067811865476], + "target": [2, -1], + "actualObstacleCount": 1 +} diff --git a/training_server/tests/generate_multi_ring_golden.py b/training_server/tests/generate_multi_ring_golden.py new file mode 100644 index 00000000..6c705bac --- /dev/null +++ b/training_server/tests/generate_multi_ring_golden.py @@ -0,0 +1,90 @@ +"""Regenerate CPU mj_ray references: run from repository root in .venv.""" + +import json +import math +import sys +from pathlib import Path + +import mujoco +import numpy as np + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from task_config import OBSTACLE_TASK, deployment_metadata, validate_task_config # noqa: E402 + + +def generate(): + config = validate_task_config( + OBSTACLE_TASK, {"sensorCfg": {"sensorMode": "multi_ring_raycast"}}, 42 + ) + deployment = deployment_metadata(OBSTACLE_TASK, config, 42) + boxes = deployment["terrain"]["boxes"] + # Low obstacle directly ahead; platform edge is a bounded-floor drop (no invented height map). + layouts = { + "default": boxes, + "low": [boxes[0], {"pos": [0.9, 0, 0.025], "size": [0.2, 0.5, 0.025], "yaw": 0}], + "edge": [boxes[0]], + } + cases = [] + for layout, position, euler in [ + ("default", [-5, 0, 0.32], [0, 0, 0]), + ("default", [-3, 1, 0.42], [0.23, -0.31, 0.51]), + ("default", [-2, -2, 0.5], [-0.35, 0.27, -0.63]), + ("low", [0, 0, 0.32], [0, 0, 0]), + ("low", [0, 0, 0.32], [0.2, 0.1, -0.1]), + ("edge", [5.45, 0, 0.32], [0, 0, 0]), + ]: + xml = ( + "" + + "".join( + ''.format( + " ".join(map(str, b["pos"])), " ".join(map(str, b["size"])) + ) + for b in layouts[layout] + ) + + "" + ) + model = mujoco.MjModel.from_xml_string(xml) + data = mujoco.MjData(model) + mujoco.mj_forward(model, data) + q = np.zeros(4) + mujoco.mju_euler2Quat(q, np.array(euler), "xyz") + matrix = np.zeros(9) + mujoco.mju_quat2Mat(matrix, q) + matrix = matrix.reshape(3, 3) + origin = np.array(position) + matrix @ np.array([0.3, 0, 0.05]) + distances, ids = [], [] + for pitch in [0, -20, -45]: + for i in range(16): + yaw = math.radians(-45 + i * 90 / 15) + p = math.radians(pitch) + direction = matrix @ np.array( + [math.cos(p) * math.cos(yaw), math.cos(p) * math.sin(yaw), math.sin(p)] + ) + geom = np.array([-1], dtype=np.int32) + distance = mujoco.mj_ray(model, data, origin, direction, None, 1, -1, geom) + distances.append(distance) + ids.append(int(geom[0])) + cases.append( + dict( + layout=layout, + position=position, + quaternion=q.tolist(), + distances=distances, + hitIds=ids, + depth=[1 if d < 0 else min(1, d / 4) for d in distances], + ) + ) + root = Path("web_platform/src/rl/fixtures") + (root / "multiRingGolden.json").write_text( + json.dumps( + dict(source=f"CPU MuJoCo {mujoco.__version__} mj_ray", layouts=layouts, cases=cases), + indent=2, + ) + + "\n" + ) + (root / "multiRingDeployment.json").write_text(json.dumps(deployment, indent=2) + "\n") + Path("/tmp/go2-multi-ring-stage5/task.json").write_text(json.dumps(config)) + + +if __name__ == "__main__": + generate() diff --git a/training_server/tests/test_custom_boxes.py b/training_server/tests/test_custom_boxes.py new file mode 100644 index 00000000..2a850c6a --- /dev/null +++ b/training_server/tests/test_custom_boxes.py @@ -0,0 +1,254 @@ +"""Authoritative custom boxes validation and CPU compilation, no training/RNG.""" + +import copy +import json +import math +import os +import sys +import tempfile +import unittest +from pathlib import Path +from unittest.mock import patch + +ROOT = Path(__file__).resolve().parents[1] +for path in (ROOT, ROOT / "rl"): + sys.path.insert(0, str(path)) +from task_config import ( # noqa: E402 + OBSTACLE_TASK, + TaskConfigError, + build_terrain_layout, + deployment_metadata, + validate_custom_terrain, + validate_task_config, +) + + +def layout(): + return json.loads((Path(__file__).parent / "fixtures/custom-boxes.json").read_text()) + + +class CustomBoxesTest(unittest.TestCase): + def test_payload_roundtrip_without_rng_and_aliasing(self): + source = layout() + with patch("task_config.random.Random", side_effect=AssertionError("must not run RNG")): + config = validate_task_config( + OBSTACLE_TASK, + { + "terrainPreset": "custom_boxes", + "customTerrainBoxes": source, + }, + 42, + ) + self.assertEqual(build_terrain_layout(config), source) + self.assertEqual(deployment_metadata(OBSTACLE_TASK, config, 42)["terrain"], source) + source["boxes"][1]["pos"][0] = 9 + self.assertEqual(config["customTerrainBoxes"], layout()) + + def test_strict_fields_numbers_floor_safety_and_count(self): + mutations = [ + lambda t: t.update(path="../../etc/passwd"), + lambda t: t.pop("target"), + lambda t: t.update(approximation=False), + lambda t: t.update(actualObstacleCount=True), + lambda t: t.update(actualObstacleCount=9), + lambda t: t.update(size=25), + lambda t: t.update(friction=True), + lambda t: t.update(spawnQuaternion=[2, 0, 0, 0]), + lambda t: t.update(spawn=[-2, -1, 0.4]), + lambda t: t.update(target=[6, 0]), + lambda t: t["boxes"][0]["pos"].__setitem__(2, 0), + lambda t: t["boxes"][1].update(mesh="../../model.stl"), + lambda t: t["boxes"][1].update(yaw=True), + lambda t: t["boxes"][1].update(pos=[6, 2, 0.5]), + lambda t: t["boxes"][1].update(pos=[1, 2, -0.3]), + lambda t: t["boxes"][1].update(pos=[1, 2, 12]), + lambda t: t.update( + boxes=t["boxes"] + [copy.deepcopy(t["boxes"][1])] * 256, actualObstacleCount=257 + ), + ] + for bad in (float("nan"), float("inf"), -float("inf"), 10**400, 0, -0.0, -1, True): + mutations.append(lambda t, bad=bad: t["boxes"][1]["size"].__setitem__(0, bad)) + for mutate in mutations: + value = layout() + mutate(value) + with self.subTest(value=str(value)[:200]), self.assertRaises(TaskConfigError): + validate_custom_terrain(value) + for key in ("spawn", "target"): + value = layout() + # Circle tangent to the right face: reject; floor alone is exempt. + value[key][:2] = [1.9, 2] + with self.assertRaisesRegex(TaskConfigError, "安全区"): + validate_custom_terrain(value) + value[key][:2] = [1.91, 2] + validate_custom_terrain(value) + # Precise circular corner test, not an expanded square approximation. + value[key][:2] = [1.8, 2.7] + validate_custom_terrain(value) + + def test_conflict_missing_path_and_training_entry(self): + from scripts.train import _load_task_config + + for payload in ( + {"terrainPreset": "custom_boxes"}, + {"terrainPreset": "custom_boxes", "customTerrainBoxes": "../../file.json"}, + {"terrainPreset": "plane", "customTerrainBoxes": layout()}, + { + "terrainPreset": "custom_boxes", + "customTerrainBoxes": layout(), + "terrainParams": {"size": 8, "friction": 0.8}, + }, + ): + with self.assertRaises(TaskConfigError): + validate_task_config(OBSTACLE_TASK, payload, 42) + config = validate_task_config( + OBSTACLE_TASK, {"terrainPreset": "custom_boxes", "customTerrainBoxes": layout()}, 42 + ) + with tempfile.TemporaryDirectory() as directory: + path = Path(directory) / "task.json" + path.write_text(json.dumps(config)) + self.assertEqual(_load_task_config(OBSTACLE_TASK, str(path), 42), config) + config["customTerrainBoxes"]["boxes"][1]["size"][0] = -0.0 + path.write_text(json.dumps(config)) + with self.assertRaises(TaskConfigError): + _load_task_config(OBSTACLE_TASK, str(path), 42) + config["customTerrainBoxesPath"] = "/etc/passwd" + path.write_text(json.dumps(config)) + with self.assertRaises(ValueError): + _load_task_config(OBSTACLE_TASK, str(path), 42) + + def test_server_payload_and_limit(self): + from server import DEFAULT_TASKS, MAX_REQUEST_BYTES, ApiError, TrainingManager + + with tempfile.TemporaryDirectory() as directory: + manager = TrainingManager( + Path(directory), sys.executable, DEFAULT_TASKS, check_environment=False + ) + payload = dict( + taskId=OBSTACLE_TASK, + terrainPreset="custom_boxes", + customTerrainBoxes=layout(), + numEnvs=2, + maxIterations=1, + seed=42, + device="cpu", + gpuIds=[], + ) + config = manager.parse_config(payload) + self.assertEqual(config.deployment["terrain"], layout()) + self.assertEqual(config.task_config["customTerrainBoxes"], layout()) + self.assertIn("custom_boxes", manager.health()["taskMetadata"][2]["terrainPresets"]) + self.assertIn( + "--task-config", manager.command_for(config, Path(directory) / "server-owned.json") + ) + with self.assertRaises(ApiError): + manager.parse_config({**payload, "customTerrainBoxesPath": "/etc/passwd"}) + full = layout() + full["boxes"] += [copy.deepcopy(full["boxes"][1])] * 255 + full["actualObstacleCount"] = 256 + # Worst typical double precision expansion still fits the bounded request envelope. + full["boxes"][1:] = [ + dict( + pos=[1.123456789012345, 2.123456789012345, 0.5123456789012345], + size=[0.4123456789012345, 0.3123456789012345, 0.5123456789012345], + yaw=0, + ) + for _ in range(256) + ] + self.assertLess( + len(json.dumps({**payload, "customTerrainBoxes": full}).encode()), MAX_REQUEST_BYTES + ) + manager.parse_config({**payload, "customTerrainBoxes": full}) + + def test_http_body_length_is_bounded_and_allows_full_layout(self): + import io + + from server import MAX_REQUEST_BYTES, ApiError, TrainingRequestHandler + + value = layout() + value["boxes"] += [copy.deepcopy(value["boxes"][1])] * 255 + value["actualObstacleCount"] = 256 + body = json.dumps({"terrainPreset": "custom_boxes", "customTerrainBoxes": value}).encode() + handler = object.__new__(TrainingRequestHandler) + handler.headers = {"Content-Length": str(len(body))} + handler.rfile = io.BytesIO(body) + self.assertEqual(handler._payload()["customTerrainBoxes"], value) + for size in (0, MAX_REQUEST_BYTES + 1): + handler.headers = {"Content-Length": str(size)} + with self.assertRaises(ApiError): + handler._payload() + + def test_real_cpu_compilation_and_custom_origin_quaternion_goal(self): + import mujoco + import numpy as np + from mjlab.terrains import TerrainGenerator + from src.tasks.obstacle_avoidance.env_cfg import ( + apply_obstacle_configuration, + unitree_go2_obstacle_env_cfg, + ) + + value = layout() + config = validate_task_config( + OBSTACLE_TASK, {"terrainPreset": "custom_boxes", "customTerrainBoxes": value}, 42 + ) + cfg = unitree_go2_obstacle_env_cfg() + apply_obstacle_configuration(cfg, config) + generator = TerrainGenerator(cfg.scene.terrain.terrain_generator) + spec = mujoco.MjSpec() + generator.compile(spec) + model = spec.compile() + np.testing.assert_allclose(model.geom_pos, [b["pos"] for b in value["boxes"]], atol=1e-12) + np.testing.assert_allclose(model.geom_size, [b["size"] for b in value["boxes"]], atol=1e-12) + np.testing.assert_allclose(generator.terrain_origins[0, 0], [-2, -1, 0]) + self.assertEqual( + cfg.scene.entities["robot"].init_state.rot, tuple(value["spawnQuaternion"]) + ) + self.assertEqual(cfg.commands["twist"].goal_offset, (4, 0)) + self.assertEqual(cfg.terminations["outside_map"].params, {"size": 12}) + self.assertTrue(math.isfinite(model.geom_size.sum())) + + @unittest.skipUnless(os.environ.get("GO2_CUSTOM_BOXES_SMOKE") == "1", "opt-in 2env/1step smoke") + def test_two_env_one_step_custom_layout(self): + import numpy as np + import torch + import warp as wp + from mjlab.envs import ManagerBasedRlEnv + from src.tasks.obstacle_avoidance.env_cfg import ( + apply_obstacle_configuration, + unitree_go2_obstacle_env_cfg, + ) + from warp._src import context + + if not torch.cuda.is_available(): + self.skipTest("CUDA unavailable") + if not hasattr(wp, "context"): + wp.context = context + cfg = unitree_go2_obstacle_env_cfg() + custom = validate_task_config( + OBSTACLE_TASK, {"terrainPreset": "custom_boxes", "customTerrainBoxes": layout()}, 42 + ) + apply_obstacle_configuration(cfg, custom, randomize_navigation=False) + cfg.scene.num_envs = 2 + env = ManagerBasedRlEnv(cfg, device="cuda:0") + try: + obs, _ = env.reset() + self.assertEqual(tuple(obs["actor"].shape), (2, 81)) + np.testing.assert_allclose(env.scene.env_origins.cpu(), [[-2, -1, 0]] * 2, atol=1e-6) + np.testing.assert_allclose( + env.scene["robot"].data.root_link_pos_w.cpu(), [layout()["spawn"]] * 2, atol=1e-6 + ) + np.testing.assert_allclose( + env.scene["robot"].data.root_link_quat_w.cpu(), + [layout()["spawnQuaternion"]] * 2, + atol=1e-6, + ) + np.testing.assert_allclose( + env.command_manager.get_term("twist").errors()[1].cpu(), [4, 4], atol=1e-6 + ) + obs, reward, terminated, truncated, _ = env.step( + torch.zeros((2, 12), device=env.device) + ) + self.assertTrue(torch.isfinite(obs["actor"]).all()) + self.assertTrue(torch.isfinite(reward).all()) + self.assertFalse(terminated.any() or truncated.any()) + finally: + env.close() diff --git a/training_server/tests/test_multi_ring.py b/training_server/tests/test_multi_ring.py new file mode 100644 index 00000000..0c23ba27 --- /dev/null +++ b/training_server/tests/test_multi_ring.py @@ -0,0 +1,214 @@ +"""Multi-ring contract, CPU reference generation and opt-in real 97-D environment.""" + +import os +import sys +import unittest +from pathlib import Path +from types import SimpleNamespace + +ROOT = Path(__file__).resolve().parents[1] +for p in (ROOT, ROOT / "rl"): + sys.path.insert(0, str(p)) + +from task_config import ( # noqa: E402 + OBSTACLE_TASK, + TaskConfigError, + deployment_metadata, + validate_task_config, +) + + +class MultiRingTest(unittest.TestCase): + def test_fresh_cli_registers_flat_rough_obstacle(self): + import subprocess + + for task in ("Unitree-Go2-Flat", "Unitree-Go2-Rough", OBSTACLE_TASK): + result = subprocess.run( + [sys.executable, "-u", "scripts/train.py", task, "--help"], + cwd=ROOT / "rl", + capture_output=True, + text=True, + timeout=60, + ) + self.assertEqual(result.returncode, 0, result.stdout + result.stderr) + self.assertIn("--env.scene.num-envs", result.stdout) + + def test_whitelist_and_legacy_defaults(self): + for mode, count in (("single_ring_raycast", 32), ("multi_ring_raycast", 48)): + c = validate_task_config(OBSTACLE_TASK, {"sensorCfg": {"sensorMode": mode}}, 42) + s = c["sensorCfg"] + self.assertEqual(s["rayCount"], count) + self.assertEqual( + deployment_metadata(OBSTACLE_TASK, c, 42)["observationSize"], 49 + count + ) + self.assertEqual(validate_task_config(OBSTACLE_TASK, c, 42), c) + for patch in ( + {"rayCount": 64}, + {"pitchAngles": [0, -45, -20]}, + {"yawCount": True}, + {"angleUnit": "rad"}, + {"rayOrder": "yaw-major"}, + {"sensorMode": "camera_depth"}, + {"yawAngles": [0] * s["yawCount"]}, + {"garbage": 1}, + ): + with self.subTest(patch=patch), self.assertRaises(TaskConfigError): + validate_task_config(OBSTACLE_TASK, {"sensorCfg": {**s, **patch}}, 42) + self.assertEqual(validate_task_config(OBSTACLE_TASK, {}, 42)["sensorCfg"]["rayCount"], 32) + + def test_floor_identity_fail_closed_and_fixed_body_world_transform(self): + import mujoco + from src.tasks.obstacle_avoidance.mdp import standard_floor_id + + floor = '' + model = mujoco.MjModel.from_xml_string( + '' + floor + "" + ) + self.assertEqual(standard_floor_id(model, 12), 0) + for geoms in ( + floor.replace("terrain_0", "other"), + floor.replace("6 6 .1", "5 5 .1"), + floor + floor.replace("terrain_0", "duplicate"), + ): + m = mujoco.MjModel.from_xml_string( + '' + geoms + "" + ) + with self.assertRaises(ValueError): + standard_floor_id(m, 12) + + def test_pattern_floor_classification_and_reward(self): + import torch + from src.tasks.obstacle_avoidance.mdp import ( + ForwardFanPatternCfg, + forward_depth, + obstacle_proximity, + ) + + offsets, rays = ForwardFanPatternCfg(sensor_mode="multi_ring_raycast").generate_rays( + None, "cpu" + ) + self.assertEqual(tuple(rays.shape), (48, 3)) + torch.testing.assert_close(offsets, torch.tensor([0.3, 0, 0.05]).repeat(48, 1)) + torch.testing.assert_close(rays.norm(dim=1), torch.ones(48)) + for i, pitch in enumerate([0, -20, -45]): + self.assertAlmostEqual( + rays[i * 16, 2].item(), + __import__("math").sin(pitch * __import__("math").pi / 180), + places=6, + ) + self.assertLess(rays[i * 16, 1], 0) + self.assertGreater(rays[i * 16 + 15, 1], 0) + # floor, 5cm obstacle, side, outside floor, miss; obs is never overwritten. + data = SimpleNamespace( + distances=torch.tensor([[0.2], [0.2], [0.2], [0.2], [-1.0]]), + hit_pos_w=torch.tensor( + [[[0.0, 0, 0]], [[0, 0, 0.05]], [[0, 0, 0]], [[7, 0, 0]], [[0, 0, 0]]] + ), + normals_w=torch.tensor( + [[[0.0, 0, 1]], [[0, 0, 1]], [[1, 0, 0]], [[0, 0, 1]], [[0, 0, 0]]] + ), + ) + env = SimpleNamespace( + scene={"forward_scan": SimpleNamespace(data=data)}, _multi_ring_floor_id=0 + ) + torch.testing.assert_close( + forward_depth(env)[:, 0], torch.tensor([0.05, 0.05, 0.05, 0.05, 1]) + ) + torch.testing.assert_close( + obstacle_proximity(env, floor_size=12), torch.tensor([0.0, 0.36, 0.36, 0.36, 0.0]) + ) + self.assertGreater(obstacle_proximity(env)[0], 0) # legacy unchanged + + @unittest.skipUnless(os.environ.get("GO2_RUN_MULTI_SMOKE") == "1", "opt-in GPU smoke") + def test_real_2env_1step_97_shape_floor_reward_and_cpu_parity(self): + import mujoco + import numpy as np + import torch + import warp as wp + from mjlab.envs import ManagerBasedRlEnv + from src.tasks.obstacle_avoidance.env_cfg import ( + apply_obstacle_configuration, + unitree_go2_obstacle_env_cfg, + ) + from src.tasks.obstacle_avoidance.mdp import ( + ForwardFanPatternCfg, + floor_top_hits, + obstacle_proximity, + ) + from warp._src import context + + if not hasattr(wp, "context"): + wp.context = context + custom = validate_task_config( + OBSTACLE_TASK, + { + "terrainPreset": "plane", + "sensorCfg": {"sensorMode": "multi_ring_raycast", "safetyDistance": 1}, + }, + 42, + ) + cfg = unitree_go2_obstacle_env_cfg() + apply_obstacle_configuration(cfg, custom) + cfg.scene.num_envs = 2 + env = ManagerBasedRlEnv(cfg, device="cuda:0") + try: + obs, _ = env.reset() + m = env.sim.mj_model + print( + "FLOOR_DIAGNOSTIC", + [ + ( + m.geom(i).name, + int(m.geom_bodyid[i]), + m.geom_pos[i].tolist(), + m.geom_size[i].tolist(), + m.geom_quat[i].tolist(), + ) + for i in range(m.ngeom) + if m.geom_group[i] == 0 + ], + ) + obs, reward, _, _, _ = env.step(torch.zeros((2, 12), device=env.device)) + self.assertEqual(tuple(obs["actor"].shape), (2, 97)) + self.assertTrue(torch.isfinite(obs["actor"]).all() and torch.isfinite(reward).all()) + scan = env.scene["forward_scan"].data + self.assertTrue(floor_top_hits(scan, 12).any()) + self.assertTrue((obs["actor"][:, 63:95] < 1).any()) + torch.testing.assert_close( + obstacle_proximity(env, safety_distance=1, floor_size=12), + torch.zeros(2, device=env.device), + ) + model = env.sim.mj_model + data = mujoco.MjData(model) + offsets, rays = ForwardFanPatternCfg(sensor_mode="multi_ring_raycast").generate_rays( + None, "cpu" + ) + for e in range(2): + data.qpos[:] = env.sim.data.qpos[e].cpu().numpy() + mujoco.mj_forward(model, data) + body = model.body("robot/base_link").id + rotation = data.xmat[body].reshape(3, 3) + expected = [] + for o, d in zip(offsets.numpy(), rays.numpy(), strict=True): + t = mujoco.mj_ray( + model, + data, + data.xpos[body] + rotation @ o, + rotation @ d, + np.array([1, 0, 0, 0, 0, 0], dtype=np.uint8), + 1, + -1, + np.array([-1], dtype=np.int32), + ) + expected.append(1 if t < 0 else min(1, t / 4)) + np.testing.assert_allclose(obs["actor"][e, 47:95].cpu(), expected, atol=2e-5) + print( + "MULTI_SMOKE: 2env x 1step actor=(2,97); CPU mj_ray all48 parity; " + "floor obs retained, proximity=0; reward finite" + ) + finally: + env.close() + + +if __name__ == "__main__": + unittest.main() diff --git a/training_server/tests/test_obstacle_env.py b/training_server/tests/test_obstacle_env.py new file mode 100644 index 00000000..1a91cb50 --- /dev/null +++ b/training_server/tests/test_obstacle_env.py @@ -0,0 +1,223 @@ +"""Optional installed-mjlab checks; GPU smoke is opt-in, never a full training run.""" + +import importlib.util +import os +import sys +import unittest +from pathlib import Path +from types import SimpleNamespace + +SERVICE_ROOT = Path(__file__).resolve().parents[1] +for root in (SERVICE_ROOT, SERVICE_ROOT / "rl"): + sys.path.insert(0, str(root)) + +HAS_MJLAB = importlib.util.find_spec("mjlab") is not None + + +@unittest.skipUnless(HAS_MJLAB, "mjlab is not installed in this Python") +class ObstacleContractTest(unittest.TestCase): + def test_training_config_file_is_revalidated(self): + import json + import tempfile + + from scripts.train import _load_task_config + from task_config import OBSTACLE_TASK, TaskConfigError, validate_task_config + + with tempfile.TemporaryDirectory() as directory: + source = Path(directory) / "training_config.json" + custom = validate_task_config(OBSTACLE_TASK, {"sensorCfg": {"fov": 60}}, 42) + source.write_text(json.dumps(custom)) + self.assertEqual(_load_task_config(OBSTACLE_TASK, str(source), 42), custom) + with self.assertRaises(ValueError): + _load_task_config(OBSTACLE_TASK, str(source), 43) + custom["sensorCfg"]["maxDistance"] = -1 + source.write_text(json.dumps(custom)) + with self.assertRaises(TaskConfigError): + _load_task_config(OBSTACLE_TASK, str(source), 42) + source.write_text("x" * (128 * 1024 + 1)) + with self.assertRaises(ValueError): + _load_task_config(OBSTACLE_TASK, str(source), 42) + + def test_pattern_order_normalization_and_body_offsets(self): + import torch + from src.tasks.obstacle_avoidance.mdp import ForwardFanPatternCfg, forward_depth + + offsets, directions = ForwardFanPatternCfg(fov=90).generate_rays(None, "cpu") + self.assertEqual(tuple(directions.shape), (32, 3)) + torch.testing.assert_close(offsets, torch.tensor([0.3, 0, 0.05]).repeat(32, 1)) + torch.testing.assert_close(directions.norm(dim=1), torch.ones(32)) + self.assertAlmostEqual(directions[0, 1].item(), -(0.5**0.5), places=6) + self.assertAlmostEqual(directions[-1, 1].item(), 0.5**0.5, places=6) + scene = { + "forward_scan": SimpleNamespace( + data=SimpleNamespace( + distances=torch.tensor([[-1, 0, 1, 4, 8]], dtype=torch.float32) + ) + ) + } + result = forward_depth(SimpleNamespace(scene=scene), max_distance=4) + torch.testing.assert_close(result, torch.tensor([[1, 0, 0.25, 1, 1]])) + + def test_navigation_matches_body_heading_and_arrival_stop(self): + import torch + from src.tasks.obstacle_avoidance.mdp import NavigationCommandCfg + + data = SimpleNamespace( + root_link_pos_w=torch.tensor([[-5.0, 0, 0.32]]), + root_link_quat_w=torch.tensor([[1.0, 0, 0, 0]]), + ) + + class Scene(dict): + env_origins = torch.tensor([[-5.0, 0, 0]]) + + env = SimpleNamespace( + num_envs=1, device="cpu", scene=Scene(robot=SimpleNamespace(data=data)) + ) + command = NavigationCommandCfg(resampling_time_range=(1e9, 1e9)).build(env) + torch.testing.assert_close(command.command, torch.tensor([[0.6, 0, 0]])) + self.assertAlmostEqual(command.errors()[1].item(), 10) + data.root_link_quat_w[:] = torch.tensor([[0.5**0.5, 0, 0, 0.5**0.5]]) + self.assertAlmostEqual(command.command[0, 2].item(), -1) + data.root_link_pos_w[0, 0] = 4.8 + torch.testing.assert_close(command.command, torch.zeros((1, 3))) + + def test_navigation_reset_randomizes_safe_distant_pairs_and_heading(self): + import torch + from src.tasks.obstacle_avoidance.mdp import NavigationCommandCfg + + class Robot: + def __init__(self, count): + self.data = SimpleNamespace( + default_root_state=torch.tensor([[0, 0, 0.32] + [0] * 10] * count), + root_link_pos_w=torch.zeros((count, 3)), + root_link_quat_w=torch.zeros((count, 4)), + ) + + def write_root_link_pose_to_sim(self, pose, env_ids): + self.data.root_link_pos_w[env_ids] = pose[:, :3] + self.data.root_link_quat_w[env_ids] = pose[:, 3:] + + def write_root_link_velocity_to_sim(self, velocity, env_ids): + self.velocity = velocity + + class Scene(dict): + env_origins = torch.zeros((64, 3)) + + robot = Robot(64) + env = SimpleNamespace(num_envs=64, device="cpu", scene=Scene(robot=robot)) + command = NavigationCommandCfg( + resampling_time_range=(1e9, 1e9), + navigation_points=((-2, 0), (-1, 0), (0, 0), (1, 0), (2, 0)), + component_starts=(0,), + component_counts=(5,), + fallback_pairs=(((-2, 0), (2, 0)),), + min_goal_distance=2, + ).build(env) + torch.manual_seed(7) + command.sample_episode(torch.arange(64)) + self.assertTrue( + ((command.goals_w - robot.data.root_link_pos_w[:, :2]).norm(dim=1) >= 2).all() + ) + self.assertGreater(torch.unique(robot.data.root_link_pos_w[:, :2], dim=0).shape[0], 1) + self.assertGreater(torch.unique(command.goals_w, dim=0).shape[0], 1) + torch.testing.assert_close(robot.data.root_link_quat_w.norm(dim=1), torch.ones(64)) + self.assertTrue((robot.velocity == 0).all()) + + def test_all_terrain_presets_compile_at_exported_world_coordinates(self): + import mujoco + import numpy as np + from mjlab.terrains import TerrainGenerator + from src.tasks.obstacle_avoidance.env_cfg import unitree_go2_obstacle_env_cfg + from src.tasks.obstacle_avoidance.terrain import apply_terrain_configuration + from task_config import ( + OBSTACLE_TASK, + TERRAIN_PRESETS, + build_terrain_layout, + validate_task_config, + ) + + for preset in (p for p in TERRAIN_PRESETS if p != "custom_boxes"): + with self.subTest(preset=preset): + custom = validate_task_config(OBSTACLE_TASK, {"terrainPreset": preset}, 7) + cfg = unitree_go2_obstacle_env_cfg() + apply_terrain_configuration(cfg, custom) + generator = TerrainGenerator(cfg.scene.terrain.terrain_generator) + spec = mujoco.MjSpec() + generator.compile(spec) + model = spec.compile() + layout = build_terrain_layout(custom) + np.testing.assert_allclose(model.geom_pos, [box["pos"] for box in layout["boxes"]]) + np.testing.assert_allclose( + model.geom_size, [box["size"] for box in layout["boxes"]] + ) + np.testing.assert_allclose(generator.terrain_origins[0, 0], [-5, 0, 0]) + np.testing.assert_allclose(model.geom_friction[:, 0], layout["friction"]) + self.assertTrue((model.geom_group == 0).all()) + + @unittest.skipUnless(os.environ.get("GO2_RUN_MJLAB_SMOKE") == "1", "opt-in GPU smoke") + def test_real_environment_81_observation_and_mujoco_raycast_parity(self): + import mujoco + import numpy as np + import torch + import warp as wp + from mjlab.envs import ManagerBasedRlEnv + from src.tasks.obstacle_avoidance.env_cfg import ( + apply_obstacle_configuration, + unitree_go2_obstacle_env_cfg, + ) + from src.tasks.obstacle_avoidance.mdp import ForwardFanPatternCfg + from task_config import JOINT_NAMES, OBSTACLE_TASK, validate_task_config + from warp._src import context + + if not torch.cuda.is_available(): + self.skipTest("CUDA is unavailable") + if not hasattr(wp, "context"): + wp.context = context + cfg = unitree_go2_obstacle_env_cfg() + apply_obstacle_configuration( + cfg, validate_task_config(OBSTACLE_TASK, {}, 42), randomize_navigation=False + ) + cfg.scene.num_envs = 2 + env = ManagerBasedRlEnv(cfg, device="cuda:0") + try: + obs, _ = env.reset() + self.assertEqual(tuple(obs["actor"].shape), (2, 81)) + self.assertEqual(list(env.scene["robot"].joint_names), JOINT_NAMES) + np.testing.assert_allclose(env.scene.env_origins.cpu(), [[-5, 0, 0]] * 2) + np.testing.assert_allclose( + env.scene["robot"].data.root_link_pos_w.cpu(), [[-5, 0, 0.32]] * 2, atol=1e-6 + ) + for _ in range(3): + obs, reward, terminated, truncated, _ = env.step( + torch.zeros((2, 12), device=env.device) + ) + self.assertTrue(torch.isfinite(obs["actor"]).all()) + self.assertTrue(torch.isfinite(reward).all()) + self.assertFalse(terminated.any() or truncated.any()) + model = env.sim.mj_model + data = mujoco.MjData(model) + data.qpos[:] = env.sim.data.qpos[0].cpu().numpy() + mujoco.mj_forward(model, data) + body = model.body("robot/base_link").id + rotation = data.xmat[body].reshape(3, 3) + offsets, directions = ForwardFanPatternCfg().generate_rays(None, "cpu") + expected = [] + for offset, direction in zip(offsets.numpy(), directions.numpy(), strict=True): + distance = mujoco.mj_ray( + model, + data, + data.xpos[body] + rotation @ offset, + rotation @ direction, + np.array([1, 0, 0, 0, 0, 0], dtype=np.uint8), + 1, + -1, + np.array([-1], dtype=np.int32), + ) + expected.append(1 if distance < 0 else min(1, distance / 4)) + np.testing.assert_allclose(obs["actor"][0, 47:79].cpu(), expected, atol=2e-5) + finally: + env.close() + + +if __name__ == "__main__": + unittest.main() diff --git a/training_server/tests/test_obstacle_tuning.py b/training_server/tests/test_obstacle_tuning.py new file mode 100644 index 00000000..078fc38d --- /dev/null +++ b/training_server/tests/test_obstacle_tuning.py @@ -0,0 +1,512 @@ +"""Obstacle-specific schema, orchestration, objective math and opt-in real rollout.""" + +import json +import os +import sys +import tempfile +import threading +import unittest +from copy import deepcopy +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import patch + +SERVICE = Path(__file__).resolve().parents[1] +for root in (SERVICE, SERVICE / "rl"): + sys.path.insert(0, str(root)) + +from task_config import build_terrain_layout, validate_task_config # noqa: E402 +from tuning import obstacle_scoring as scoring # noqa: E402 +from tuning.advisor import DeepSeekAdvisor # noqa: E402 +from tuning.manager import TuningError, TuningManager # noqa: E402 +from tuning.process import GpuLease, ResourceBusyError # noqa: E402 +from tuning.schema import ( # noqa: E402 + OBSTACLE_TASK, + RewardConfigError, + apply_reward_configuration, + base_configuration, + merge_proposal, + validate_configuration, + validate_constraints, + validate_proposal, +) +from tuning.scoring import EvaluationError # noqa: E402 + + +def sample(**kw): + return ( + dict( + distance=1.0, + clearance=0.5, + action_delta=0.0, + ray_hit=0.0, + collision=0.0, + fall=0.0, + terminal=0.0, + **{}, + ) + | kw + ) + + +def evaluation(custom, n=2): + metrics = scoring.score_trajectory([sample()] * scoring.STEPS) + return { + "protocol": scoring.protocol(custom, n), + "metrics": metrics, + "seedMetrics": [ + {"seed": seed, "episodes": n, "rolloutSteps": scoring.STEPS, "metrics": metrics} + for seed in scoring.SEEDS + ], + } + + +class ObstacleSchemaTest(unittest.TestCase): + def test_ranges_task_isolation_nonfinite_and_patch_no_mutation(self): + base = base_configuration(OBSTACLE_TASK) + original = deepcopy(base) + for section, key, low, high in ( + ("weights", "avoidance_weight", 0.5, 5), + ("weights", "collision_penalty", -10, -0.5), + ("weights", "action_smoothness", -0.05, -0.001), + ("params", "target_velocity", 0.3, 1.2), + ): + for value in (low, high): + changed = deepcopy(base) + changed[section][key] = value + self.assertEqual(validate_configuration(changed, OBSTACLE_TASK), changed) + for value in (low - 0.00001, high + 0.00001, float("nan"), float("inf"), True): + with self.subTest(key=key, value=value), self.assertRaises(RewardConfigError): + validate_proposal({section: {key: value}}, base, task_id=OBSTACLE_TASK) + for proposal in ( + {"weights": {"pose": 1.1}}, + {"weights": {"target_velocity": 0.8}}, + {"params": {"seed": 101}}, + {"unknown": {}}, + {"weights": {"avoidance_weight": 4.01}}, + ): + with self.assertRaises(RewardConfigError): + validate_proposal(proposal, base, task_id=OBSTACLE_TASK) + with self.assertRaises(RewardConfigError): + validate_configuration(base) + with self.assertRaises(RewardConfigError): + validate_configuration(base_configuration(), OBSTACLE_TASK) + with self.assertRaises(RewardConfigError): + validate_constraints({"weights.pose": {"kind": "fixed", "value": 1}}, OBSTACLE_TASK) + candidate = merge_proposal( + base, {"params": {"target_velocity": 1.2}}, task_id=OBSTACLE_TASK + ) + self.assertEqual(candidate["params"]["target_velocity"], 1.2) + self.assertEqual(base, original) + + def test_real_training_config_mapping_and_command(self): + import torch + from scripts.train import ( + _configure_task_and_rewards, + _load_reward_config, + _load_task_config, + ) + from src.tasks.obstacle_avoidance.env_cfg import ( + apply_obstacle_configuration, + unitree_go2_obstacle_env_cfg, + ) + + base = base_configuration(OBSTACLE_TASK) + base["weights"].update(avoidance_weight=3, collision_penalty=-7, action_smoothness=-0.02) + base["params"]["target_velocity"] = 0.9 + custom = validate_task_config(OBSTACLE_TASK, {"sensorCfg": {"fov": 60}}, 42) + with tempfile.TemporaryDirectory() as directory: + reward, task = Path(directory) / "reward.json", Path(directory) / "task.json" + reward.write_text(json.dumps(base)) + task.write_text(json.dumps(custom)) + cfg = unitree_go2_obstacle_env_cfg() + apply_obstacle_configuration(cfg, _load_task_config(OBSTACLE_TASK, str(task), 42)) + apply_reward_configuration( + cfg, _load_reward_config(str(reward), None, OBSTACLE_TASK), OBSTACLE_TASK + ) + deployment, checked = _configure_task_and_rewards( + OBSTACLE_TASK, + SimpleNamespace( + env=unitree_go2_obstacle_env_cfg(), + task_config=str(task), + reward_config=str(reward), + reward_config_json=None, + agent=SimpleNamespace(seed=42), + ), + ) + self.assertEqual(checked, base) + self.assertEqual(deployment["navigation"]["speed"], 0.9) + self.assertEqual(deployment["sensorCfg"]["avoidanceWeight"], 3) + self.assertEqual(cfg.rewards["obstacle_proximity"].weight, -3) + self.assertEqual(cfg.rewards["obstacle_collision"].weight, -7) + self.assertIs( + cfg.rewards["obstacle_collision"].func, cfg.terminations["illegal_contact"].func + ) + self.assertEqual(cfg.rewards["action_rate_l2"].weight, -0.02) + self.assertNotIn("target_velocity", cfg.rewards) + self.assertEqual(cfg.commands["twist"].speed, 0.9) + + class Scene(dict): + env_origins = torch.tensor([[-5.0, 0, 0]]) + + scene = Scene( + robot=SimpleNamespace( + data=SimpleNamespace( + root_link_pos_w=torch.tensor([[-5.0, 0, 0.32]]), + root_link_quat_w=torch.tensor([[1.0, 0, 0, 0]]), + ) + ) + ) + command = cfg.commands["twist"].build( + SimpleNamespace(num_envs=1, device="cpu", scene=scene) + ) + self.assertAlmostEqual(command.command[0, 0].item(), 0.9, places=6) + + def test_mock_deepseek_task_context_no_network(self): + advisor = DeepSeekAdvisor() + output = SimpleNamespace( + weights={"avoidance_weight": 2.2}, + params={"target_velocity": 0.7}, + rationale="绕行与目标导航", + expected_impact={}, + confidence=0.8, + ) + fake = SimpleNamespace(run_sync=lambda prompt: SimpleNamespace(output=output)) + with patch.object(advisor, "_agent", return_value=fake): + result = advisor.propose({"task": OBSTACLE_TASK}, base_configuration(OBSTACLE_TASK)) + self.assertEqual(result["patch"]["params"], {"target_velocity": 0.7}) + + +class ObstacleScoringTest(unittest.TestCase): + def test_hand_calculated_scores_and_first_terminal(self): + samples = [ + sample(clearance=0.25, action_delta=0.5), + sample(distance=0.4, clearance=0.25, action_delta=0.5), + ] + metrics = scoring.score_trajectory(samples, 2) + self.assertEqual(metrics["success"], 1) + self.assertEqual(metrics["time"], 0) + self.assertEqual(metrics["clearance"], 0.5) + self.assertEqual(metrics["smooth"], 0.5) + self.assertAlmostEqual(scoring.score_evaluation(metrics)["score"], 0.65) + failed = scoring.score_trajectory([sample(distance=0.1, terminal=1, fall=1)], 1000) + self.assertEqual((failed["success"], failed["time"], failed["clearance"]), (0, 0, 0)) + contact = scoring.score_trajectory([sample(distance=0.1, terminal=1, collision=1)], 1000) + self.assertEqual((contact["arrival_rate"], contact["success"], contact["time"]), (1, 0, 0)) + with self.assertRaises(EvaluationError): + scoring.score_trajectory([sample()], 1000) + with self.assertRaises(EvaluationError): + scoring.score_trajectory([sample(terminal=1), sample()], 2) + with self.assertRaises(EvaluationError): + scoring.score_trajectory([sample(clearance=float("nan"))], 1) + + def test_recorder_excludes_floor_and_keeps_terminal_before_reset(self): + import torch + from scripts.evaluate_obstacle import FirstEpisodeRecorder + + class Scene(dict): + env_origins = torch.zeros((2, 3)) + + robot = SimpleNamespace( + root_link_pos_w=torch.tensor([[0.0, 0, 0.32], [0.0, 0, 0.1]]), + projected_gravity_b=torch.tensor([[0.0, 0, -1.0], [0.0, 1.0, 0.0]]), + ) + scene = Scene( + robot=SimpleNamespace(data=robot), + nonfoot_ground_touch=SimpleNamespace( + data=SimpleNamespace(force_history=torch.zeros((2, 1, 4, 3))) + ), + forward_scan=SimpleNamespace( + data=SimpleNamespace(distances=torch.ones((2, 32))), + cfg=SimpleNamespace(max_distance=4), + ), + ) + command = SimpleNamespace(errors=lambda: (None, torch.tensor([0.2, 0.2]), None)) + env = SimpleNamespace( + scene=scene, + device="cpu", + num_envs=2, + action_manager=SimpleNamespace(action=torch.ones((2, 12)) * 0.5), + command_manager=SimpleNamespace(get_term=lambda name: command), + termination_manager=SimpleNamespace(compute=lambda: torch.ones(2, dtype=torch.bool)), + ) + floor = {"pos": [0, 0, -0.1], "size": [6, 6, 0.1]} + recorder = FirstEpisodeRecorder(env, {"spawn": [0, 0, 0.32], "boxes": [floor]}) + env.termination_manager.compute() + robot.root_link_pos_w[:] = torch.tensor([0.0, 0, 0.32]) # Simulated auto-reset overwrite. + robot.projected_gravity_b[:] = torch.tensor([0.0, 0, -1.0]) + env.termination_manager.compute() + self.assertEqual([len(s) for s in recorder.samples], [1, 1]) + upright, lying = [scoring.score_trajectory(s, 2) for s in recorder.samples] + self.assertEqual(upright["clearance"], 0.5) # 1 valid safe step / 2, not floor distance 0. + self.assertEqual(upright["success"], 1) + self.assertEqual(lying["clearance"], 0) + self.assertEqual(lying["success"], 0) + self.assertEqual(lying["time"], 0) + self.assertEqual(lying["fall_rate"], 1) + + def test_threshold_boundary_and_fail_closed(self): + base = scoring.score_trajectory([sample(distance=0.1)], 1) + base.update(success=0.5, fall_rate=0.1, no_fall=0.9) + current = dict(base, success=0.48, fall_rate=0.12, no_fall=0.88) + self.assertTrue(scoring.score_evaluation(current, base)["eligible"]) + for changed in ( + dict(current, success=0.479999), + dict(current, fall_rate=0.120001, no_fall=0.879999), + ): + self.assertFalse(scoring.score_evaluation(changed, base)["eligible"]) + self.assertEqual(scoring.score_evaluation(changed, base)["score"], -1) + custom = validate_task_config(OBSTACLE_TASK, {}, 42) + expected = scoring.protocol(custom, 2) + valid = evaluation(custom) + scoring.validate_evaluation(valid, expected) + for mutate in ( + lambda v: v["seedMetrics"].pop(), + lambda v: v["metrics"].pop("success"), + lambda v: v["protocol"].update(stepsPerSeed=999), + lambda v: v["seedMetrics"][0].update(episodes=1), + lambda v: v["metrics"].update(smooth=float("inf")), + ): + bad = deepcopy(valid) + mutate(bad) + with self.assertRaises(EvaluationError): + scoring.validate_evaluation(bad, expected) + + def test_fixed_three_seed_and_custom_authoritative_map(self): + custom = validate_task_config(OBSTACLE_TASK, {}, 42) + scenes = scoring.evaluation_scenarios(custom) + self.assertEqual([s["seed"] for s in scenes], [101, 202, 303]) + self.assertNotEqual(scenes[0]["terrain"], scenes[1]["terrain"]) + layout = build_terrain_layout(custom) + layout["approximation"] = True + custom = validate_task_config( + OBSTACLE_TASK, {"terrainPreset": "custom_boxes", "customTerrainBoxes": layout}, 42 + ) + fixed = scoring.protocol(custom, 2) + self.assertEqual(fixed["sceneMode"], "fixed-custom-map") + self.assertEqual([s["terrain"] for s in fixed["scenarios"]], [layout] * 3) + self.assertEqual(fixed, scoring.protocol(deepcopy(custom), 2)) + + +class ObstacleManagerTest(unittest.TestCase): + def setUp(self): + self.temp = tempfile.TemporaryDirectory() + root = Path(self.temp.name) + self.manager = TuningManager(SERVICE / "rl", sys.executable, root, GpuLease()) + payload = { + "taskId": OBSTACLE_TASK, + "mode": "approval", + "evalNumEnvs": 2, + "numEnvs": 2, + "initialIterations": 1, + "middleIterations": 1, + "finalIterations": 1, + "sensorCfg": {"fov": 60, "avoidanceWeight": 3}, + } + mode, config, objective, fallback = self.manager.parse_create(payload) + self.session = self.manager.storage.create_session(mode, config, objective, fallback) + self.trial = self.manager.storage.create_trial( + self.session["id"], + 0, + 0, + 1, + self.manager._base_configuration(self.session), + None, + "trial-000-rung-0", + ) + + def tearDown(self): + self.manager.shutdown() + self.temp.cleanup() + + def test_argv_json_real_train_and_eval_ingress(self): + commands = [] + custom = self.session["config"]["taskConfig"] + + def run(session_id, command, cwd, environment, log_path): + commands.append(command) + task_path = Path(command[command.index("--task-config") + 1]) + reward_path = Path(command[command.index("--reward-config") + 1]) + from scripts.train import _load_reward_config, _load_task_config + + self.assertEqual(_load_task_config(OBSTACLE_TASK, str(task_path), 42), custom) + self.assertEqual( + _load_reward_config(str(reward_path), None, OBSTACLE_TASK), + self.trial["rewardConfig"], + ) + if "scripts/train.py" in command: + (log_path.parent / "model_0.pt").write_bytes(b"mock-checkpoint") + (log_path.parent / "policy.onnx").write_bytes(b"mock-onnx") + else: + Path(command[command.index("--output") + 1]).write_text( + json.dumps(evaluation(custom)) + ) + return 0 + + self.manager.cancel_events[self.session["id"]] = threading.Event() + with ( + patch.object(self.manager, "_run_command", side_effect=run), + patch.object(self.manager.studies, "record", return_value=0), + ): + result = self.manager._execute_trial(self.session, self.trial) + self.assertEqual(result["state"], "completed") + self.assertIn("--steps-per-seed=1000", commands[1]) + self.assertIn(OBSTACLE_TASK, commands[0]) + context = self.manager._proposal_context(self.session) + self.assertIn("32", context["taskContext"]) + self.assertEqual(set(context["allowlist"]["params"]), {"target_velocity"}) + + def test_cas_stale_guardrails_and_mode_roundtrip(self): + sid = self.session["id"] + constraint = {"params.target_velocity": {"kind": "range", "min": 0.4, "max": 0.8}} + self.manager.set_constraints(sid, {"revision": 0, "constraints": constraint}) + before = self.manager.detail(sid) + with self.assertRaises(ResourceBusyError): + self.manager.set_constraints(sid, {"revision": 0, "constraints": {}}) + self.assertEqual(self.manager.detail(sid), before) + for bad in ( + {"weights.pose": {"kind": "fixed", "value": 1}}, + {"params.target_velocity": {"kind": "fixed", "value": float("nan")}}, + ): + with self.assertRaises(RewardConfigError): + self.manager.set_constraints(sid, {"revision": 1, "constraints": bad}) + self.assertEqual(self.manager.detail(sid), before) + self.manager.storage.update_session(sid, state="awaiting_approval") + proposal = self.manager.storage.create_proposal( + sid, self.trial["id"], {"params": {"target_velocity": 0.7}}, "test", {}, 0.8, "agent" + ) + self.assertEqual(self.manager.set_mode(sid, {"mode": "automatic"})["mode"], "automatic") + self.assertEqual(self.manager.storage.get_proposal(proposal["id"])["state"], "approved") + self.assertEqual(self.manager.set_mode(sid, {"mode": "approval"})["mode"], "approval") + with self.assertRaises(RewardConfigError): + validate_proposal( + {"params": {"target_velocity": 0.9}}, + self.trial["rewardConfig"], + constraint, + OBSTACLE_TASK, + ) + + def test_protocol_fields_cannot_be_changed(self): + for values in ( + {"evalSteps": 999}, + {"seeds": [1, 2, 3]}, + {"objectiveWeights": {}}, + {"taskConfig": {"unknown": 1}}, + {"taskConfig": {"seed": True}}, + {"seed": 1, "taskConfig": {"seed": True}}, + ): + with self.assertRaises(TuningError): + self.manager.parse_create({"taskId": OBSTACLE_TASK, **values}) + + +class IsolatedEvaluationTest(unittest.TestCase): + def test_spawn_results_and_failure_cancel_missing_seed_fail_closed(self): + import subprocess + + from scripts.evaluate import EvaluateConfig + from scripts.evaluate_obstacle import evaluate_isolated_seeds + + custom = validate_task_config(OBSTACLE_TASK, {}, 42) + fixed = scoring.protocol(custom, 2) + reward = base_configuration(OBSTACLE_TASK) + with tempfile.TemporaryDirectory() as directory: + checkpoint = Path(directory) / "model.pt" + checkpoint.write_bytes(b"test checkpoint snapshot") + cfg = EvaluateConfig(checkpoint=str(checkpoint), output="unused", num_envs=2) + + def run(command, **kwargs): + self.assertFalse(kwargs["shell"]) + self.assertEqual(command[0], sys.executable) + request = json.loads(Path(command[2]).read_text()) + result = evaluation(custom)["seedMetrics"][0] + result["seed"] = request["scenario"]["seed"] + result["checkpointSha256"] = request["checkpointSha256"] + Path(command[3]).write_text(json.dumps(result)) + + with patch("scripts.evaluate_obstacle.subprocess.run", side_effect=run) as mocked: + results = evaluate_isolated_seeds(OBSTACLE_TASK, cfg, fixed, reward) + self.assertEqual([r["seed"] for r in results], [101, 202, 303]) + self.assertEqual(mocked.call_count, 3) + for failure in ( + subprocess.CalledProcessError(1, "worker"), + subprocess.TimeoutExpired("worker", 1200), + KeyboardInterrupt(), + ): + with patch( + "scripts.evaluate_obstacle.subprocess.run", side_effect=failure + ) as mocked: + with self.assertRaises(type(failure)): + evaluate_isolated_seeds(OBSTACLE_TASK, cfg, fixed, reward) + self.assertEqual(mocked.call_count, 1) + with ( + patch("scripts.evaluate_obstacle.subprocess.run"), + self.assertRaises(FileNotFoundError), + ): + evaluate_isolated_seeds(OBSTACLE_TASK, cfg, fixed, reward) + + def wrong_seed(command, **kwargs): + run(command, **kwargs) + path = Path(command[3]) + result = json.loads(path.read_text()) + result["seed"] = 0 + path.write_text(json.dumps(result)) + + with ( + patch("scripts.evaluate_obstacle.subprocess.run", side_effect=wrong_seed), + self.assertRaises(EvaluationError), + ): + evaluate_isolated_seeds(OBSTACLE_TASK, cfg, fixed, reward) + + +@unittest.skipUnless( + os.environ.get("GO2_RUN_TUNING_SMOKE") == "1", "opt-in real GPU checkpoint rollout" +) +class ObstacleRolloutSmoke(unittest.TestCase): + def test_actual_checkpoint_policy_statistics_and_first_terminal(self): + from dataclasses import asdict + + import src.tasks # noqa: F401 + import torch + from mjlab.envs import ManagerBasedRlEnv + from mjlab.rl import RslRlVecEnvWrapper + from mjlab.tasks.registry import load_rl_cfg, load_runner_cls + from scripts.evaluate import EvaluateConfig + from scripts.evaluate_obstacle import FirstEpisodeRecorder, configure_seed + + cfg = EvaluateConfig( + checkpoint=os.environ["GO2_TUNING_CHECKPOINT"], + output="/tmp/unused.json", + num_envs=2, + device="cuda:0", + ) + custom = validate_task_config(OBSTACLE_TASK, {}, 42) + scenario = scoring.evaluation_scenarios(custom)[0] + env_cfg = configure_seed(OBSTACLE_TASK, cfg, scenario, base_configuration(OBSTACLE_TASK)) + env_cfg.episode_length_s = 0.02 # Test-only one-step timeout exercises auto-reset capture. + env = ManagerBasedRlEnv(env_cfg, device="cuda:0") + agent_cfg = load_rl_cfg(OBSTACLE_TASK) + wrapped = RslRlVecEnvWrapper(env, clip_actions=agent_cfg.clip_actions) + try: + runner = load_runner_cls(OBSTACLE_TASK)( + wrapped, asdict(agent_cfg), log_dir=None, device="cuda:0" + ) + runner.load( + cfg.checkpoint, load_cfg={"actor": True}, strict=True, map_location="cuda:0" + ) + saved = torch.load(cfg.checkpoint, weights_only=False)["actor_state_dict"] + actual = runner.alg.actor.state_dict() + normalizers = [k for k in saved if "normaliz" in k] + self.assertTrue(normalizers) + for key in normalizers: + torch.testing.assert_close(actual[key], saved[key].to(actual[key].device)) + obs, _ = env.reset(seed=101) + recorder = FirstEpisodeRecorder(env, scenario["terrain"]) + with torch.inference_mode(): + policy = runner.get_inference_policy(device="cuda:0") + for _ in range(2): + obs, _, _, _ = wrapped.step(policy(obs)) + self.assertEqual([len(s) for s in recorder.samples], [1, 1]) + self.assertTrue(all(s[0]["terminal"] for s in recorder.samples)) + scoring.validate_metrics(recorder.metrics(horizon=2)) + finally: + wrapped.close() diff --git a/training_server/tests/test_pretrained.py b/training_server/tests/test_pretrained.py new file mode 100644 index 00000000..bd10fc71 --- /dev/null +++ b/training_server/tests/test_pretrained.py @@ -0,0 +1,475 @@ +"""CPU transfer regression; real source and <=4env/1iteration GPU checks are opt-in.""" + +import copy +import importlib.util +import io +import json +import os +import sys +import tempfile +import unittest +from dataclasses import asdict, replace +from pathlib import Path +from types import SimpleNamespace + +ROOT = Path(__file__).resolve().parents[1] +for root in (ROOT, ROOT / "rl"): + sys.path.insert(0, str(root)) +HAS_STACK = all(importlib.util.find_spec(m) is not None for m in ("mjlab", "torch", "onnxruntime")) + + +@unittest.skipUnless(HAS_STACK, "installed RL + CPU ORT stack required") +class PretrainedTest(unittest.TestCase): + @classmethod + def setUpClass(cls): + import torch + + torch.set_num_threads(1) + from pretrained import ValidatedSource, make_reference_actor + 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 + + torch.manual_seed(42) + actor = make_reference_actor() + with torch.no_grad(): + actor.obs_normalizer._mean.uniform_(-0.2, 0.2) + actor.obs_normalizer._var.uniform_(0.1, 1.2) + actor.obs_normalizer._std.copy_(actor.obs_normalizer._var.sqrt()) + actor.obs_normalizer.count.fill_(983138304) + cls.source = ValidatedSource(actor.state_dict(), {"source_id": "synthetic"}) + cls.env_cfg = asdict(unitree_go2_flat_env_cfg()) + cls.agent_cfg = asdict(unitree_go2_ppo_runner_cfg()) + + def test_extension_preserves_actions_and_learns_new_columns(self): + import torch + from pretrained import comparison_observations, make_reference_actor, warm_start_actor + from tensordict import TensorDict + + source_actor = make_reference_actor().eval() + source_actor.load_state_dict(self.source.actor_state) + base = comparison_observations() + expected = source_actor.mlp(source_actor.obs_normalizer(base)).detach() + for dim in (47, 81, 97): + with self.subTest(dim=dim): + actor = make_reference_actor(dim) + warm_start_actor(actor, self.source) + for key in ("_mean", "_var", "_std"): + value = getattr(actor.obs_normalizer, key) + self.assertTrue( + torch.equal(value[:, :47], self.source.actor_state[f"obs_normalizer.{key}"]) + ) + self.assertTrue( + torch.equal( + value[:, 47:], + torch.full_like(value[:, 47:], 0 if key == "_mean" else 1), + ) + ) + self.assertTrue( + torch.equal( + actor.obs_normalizer.count, self.source.actor_state["obs_normalizer.count"] + ) + ) + self.assertEqual(actor.mlp[0].weight[:, 47:].count_nonzero().item(), 0) + self.assertTrue(actor.mlp[0].weight.requires_grad) + extra = torch.rand(len(base), dim - 47) * 2 - 1 + x = torch.cat((base, extra), dim=1) + torch.testing.assert_close( + actor.mlp(actor.obs_normalizer(x)), expected, atol=3e-6, rtol=3e-6 + ) + initial_error = (actor.mlp(actor.obs_normalizer(x)) - expected).abs().max().item() + x[:, 47:] *= -1 + torch.testing.assert_close( + actor.mlp(actor.obs_normalizer(x)), expected, atol=3e-6, rtol=3e-6 + ) + before = actor.obs_normalizer._mean.clone() + obs = TensorDict({"actor": x}, batch_size=[len(base)]) + actor.update_normalization(obs) + self.assertEqual(actor.obs_normalizer.count.item(), 983138304 + len(base)) + self.assertLess( + (actor.obs_normalizer._mean[:, :47] - before[:, :47]).abs().max().item(), 1e-6 + ) + # Updating statistics is not eval identity. Random gravity probes include + # out-of-distribution values in nearly constant source channels. + updated_source = copy.deepcopy(source_actor).train() + updated_source.update_normalization( + TensorDict({"actor": base}, batch_size=[len(base)]) + ) + updated_expected = updated_source.mlp(updated_source.obs_normalizer(base)).detach() + torch.testing.assert_close(actor(obs), updated_expected, atol=3e-6, rtol=3e-6) + self.assertLess((actor(obs) - expected).abs().max().item(), 2e-4) + if hasattr(self, "evidence"): + self.evidence.setdefault("cpu_transfer", {})[str(dim)] = { + "initial_max_abs_error": initial_error, + "updated_source_max_abs_error": (actor(obs) - updated_expected) + .abs() + .max() + .item(), + "random_batch_action_max_abs_change": (actor(obs) - expected) + .abs() + .max() + .item(), + } + loss = actor(obs).square().mean() + loss.backward() + grad = actor.mlp[0].weight.grad + self.assertTrue(torch.isfinite(grad).all()) + if dim > 47: + self.assertGreater(grad[:, 47:].abs().max().item(), 0) + torch.optim.Adam(actor.parameters(), lr=1e-4).step() + if dim > 47: + self.assertGreater(actor.mlp[0].weight[:, 47:].abs().max().item(), 0) + stream = io.BytesIO() + torch.save(actor.state_dict(), stream) + stream.seek(0) + restored = make_reference_actor(dim) + restored.load_state_dict(torch.load(stream, weights_only=True), strict=True) + for k, v in actor.state_dict().items(): + self.assertTrue(torch.equal(v, restored.state_dict()[k]), k) + + def test_rejects_shapes_nonfinite_and_runtime_activation(self): + import torch + from pretrained import ( + PretrainedError, + ValidatedSource, + make_reference_actor, + warm_start_actor, + ) + + for key, value in ( + ("mlp.0.weight", torch.zeros(512, 46)), + ("obs_normalizer._std", torch.full((1, 47), float("nan"))), + ("distribution.std_param", torch.zeros(12)), + ): + state = dict(self.source.actor_state) + state[key] = value + with self.assertRaises(PretrainedError): + warm_start_actor(make_reference_actor(81), ValidatedSource(state, {})) + actor = make_reference_actor(81) + actor.mlp[1] = torch.nn.ReLU() + with self.assertRaisesRegex(PretrainedError, "architecture"): + warm_start_actor(actor, self.source) + with self.assertRaises(PretrainedError): + warm_start_actor(make_reference_actor(82), self.source) + + def test_semantics_fail_closed(self): + from pretrained import PretrainedError, _plain, validate_semantics + + env, agent = _plain(self.env_cfg), _plain(self.agent_cfg) + validate_semantics(env, agent, self.env_cfg, self.agent_cfg) + for edit in ( + lambda e: e["observations"]["actor"]["terms"]["phase"]["params"].update(period=0.7), + lambda e: e["observations"]["actor"]["terms"]["joint_vel"].update(scale=0.1), + lambda e: e["actions"]["joint_pos"].update(scale=0.5), + lambda e: e["scene"]["entities"]["robot"]["articulation"]["actuators"][0].update( + armature=0.1 + ), + ): + bad = copy.deepcopy(env) + edit(bad) + with self.assertRaises(PretrainedError): + validate_semantics(bad, agent, self.env_cfg, self.agent_cfg) + with self.assertRaises(PretrainedError): + validate_semantics(env, agent, bad, self.agent_cfg) + bad_agent = copy.deepcopy(agent) + bad_agent["actor"]["activation"] = "relu" + with self.assertRaises(PretrainedError): + validate_semantics(env, bad_agent, self.env_cfg, self.agent_cfg) + + def test_compiled_runtime_contract(self): + from unittest.mock import patch + + from pretrained import BASE_TERMS, JOINTS, PretrainedError, validate_runtime_contract + + metadata = { + "joint_names": JOINTS, + "observation_names": BASE_TERMS + ["forward_depth", "target_error"], + "command_names": ["twist"], + "action_scale": 0.25, + "joint_stiffness": [20, 20, 40] * 4, + "joint_damping": [1, 1, 2] * 4, + "default_joint_pos": [-0.1, 0.9, -1.8, 0.1, 0.9, -1.8] * 2, + } + with patch("mjlab.rl.exporter_utils.get_base_metadata", return_value=metadata): + validate_runtime_contract(None) + metadata["joint_names"] = list(reversed(JOINTS)) + with self.assertRaisesRegex(PretrainedError, "joint order"): + validate_runtime_contract(None) + + def test_safe_yaml_and_allowed_roots(self): + from pretrained import PretrainedError, _read_allowed, _yaml_data + + with self.assertRaises(PretrainedError): + _yaml_data(b"x: !!python/object/apply:os.system ['touch /tmp/never-pretrained']") + with self.assertRaises(PretrainedError): + _yaml_data(b"x: 1\nx: 2") + self.assertEqual(_yaml_data(b"axis: {0: 1, 1: 2}"), {"axis": {0: 1, 1: 2}}) + self.assertEqual( + _yaml_data(b"x: !!python/name:os.system ''"), {"x": {"symbol": "os.system"}} + ) + with tempfile.TemporaryDirectory() as directory, tempfile.TemporaryDirectory() as outside: + root = Path(directory) + secret = Path(outside) / "secret.pt" + secret.write_bytes(b"x") + (root / "link.pt").symlink_to(secret) + for path in (secret, root / "link.pt", root / ".." / Path(outside).name / "secret.pt"): + with self.assertRaises(PretrainedError): + _read_allowed(path, [root], 128) + with self.assertRaises(PretrainedError): + _read_allowed(secret, [Path(outside)], 0) + + def test_fresh_runner_only_and_trial_consistency(self): + import torch + from pretrained import PretrainedError, initialize_runner, make_reference_actor + + actors = [] + for _ in range(2): + actor = make_reference_actor(81) + critic = torch.nn.Linear(108, 1) + original = copy.deepcopy(critic.state_dict()) + optimizer = torch.optim.Adam(list(actor.parameters()) + list(critic.parameters())) + runner = SimpleNamespace( + current_learning_iteration=0, + alg=SimpleNamespace(actor=actor, critic=critic, optimizer=optimizer), + ) + initialize_runner(runner, self.source) + self.assertFalse(optimizer.state) + self.assertEqual(runner.current_learning_iteration, 0) + for k in original: + self.assertTrue(torch.equal(original[k], critic.state_dict()[k])) + actors.append(actor.state_dict()) + runner.current_learning_iteration = 1 + with self.assertRaises(PretrainedError): + initialize_runner(runner, self.source) + runner.current_learning_iteration = 0 + optimizer.state[actor.mlp[0].weight] = {"step": torch.tensor(1)} + with self.assertRaises(PretrainedError): + initialize_runner(runner, self.source) + for k in actors[0]: + self.assertTrue(torch.equal(actors[0][k], actors[1][k])) + + def test_cli_modes_do_not_reinitialize_resume(self): + 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 + + cfg = TrainConfig(env=unitree_go2_flat_env_cfg(), agent=unitree_go2_ppo_runner_cfg()) + self.assertIsNone(_load_pretrained(cfg)) + self.assertIsNone(_load_pretrained(replace(cfg, resume_checkpoint="trial.pt"))) + with self.assertRaisesRegex(ValueError, "mutually exclusive"): + _load_pretrained( + replace(cfg, resume_checkpoint="trial.pt", pretrained_checkpoint="source.pt") + ) + cfg.agent.resume = True + with self.assertRaisesRegex(ValueError, "mutually exclusive"): + _load_pretrained(replace(cfg, pretrained_checkpoint="source.pt")) + with self.assertRaisesRegex(ValueError, "ONNX alone"): + _load_pretrained(replace(cfg, pretrained_onnx="policy.onnx")) + + +@unittest.skipUnless( + HAS_STACK and os.environ.get("GO2_PRETRAINED_SOURCE"), + "set GO2_PRETRAINED_SOURCE to explicit allowed source directory", +) +class RealPretrainedTest(PretrainedTest): + @classmethod + def setUpClass(cls): + super().setUpClass() + from pretrained import read_pretrained_source + + directory = Path(os.environ["GO2_PRETRAINED_SOURCE"]) + cls.source = read_pretrained_source( + directory / "model_10000.pt", + allowed_roots=[directory], + target_env=cls.env_cfg, + target_agent=cls.agent_cfg, + ) + cls.evidence = {"source": cls.source.manifest} + + @classmethod + def tearDownClass(cls): + if os.environ.get("GO2_PRETRAINED_EVIDENCE_DIR"): + directory = Path(os.environ["GO2_PRETRAINED_EVIDENCE_DIR"]) + directory.mkdir(parents=True, exist_ok=True) + (directory / "real-source-evidence.json").write_text( + json.dumps(cls.evidence, indent=2) + "\n" + ) + + def test_wrong_checkpoint_actor_rejected_by_onnx(self): + from pretrained import PretrainedError, make_reference_actor, verify_onnx + + # No bulk loading: a fresh random actor is sufficient to prove fail-closed identity. + onnx = (Path(os.environ["GO2_PRETRAINED_SOURCE"]) / "policy.onnx").read_bytes() + with self.assertRaisesRegex(PretrainedError, "does not match"): + verify_onnx(onnx, make_reference_actor()) + + def test_extended_export_matches_cpu_actor(self): + import numpy as np + import onnxruntime as ort + import torch + from pretrained import comparison_observations, make_reference_actor, warm_start_actor + + errors = {} + with tempfile.TemporaryDirectory() as directory: + for dim in (81, 97): + actor = make_reference_actor(dim).eval() + warm_start_actor(actor, self.source) + export = actor.as_onnx(verbose=False) + path = str(Path(directory) / f"policy-{dim}.onnx") + torch.onnx.export( + export, + export.get_dummy_inputs(), + path, + input_names=export.input_names, + output_names=export.output_names, + opset_version=18, + dynamo=False, + ) + session = ort.InferenceSession(path, providers=["CPUExecutionProvider"]) + x = torch.cat((comparison_observations(), torch.rand(48, dim - 47) * 2 - 1), dim=1) + with torch.no_grad(): + expected = export(x).numpy() + actual = np.concatenate( + [session.run(None, {"obs": row[None].numpy()})[0] for row in x] + ) + np.testing.assert_allclose(actual, expected, atol=2e-5, rtol=2e-5) + errors[str(dim)] = float(np.abs(actual - expected).max()) + self.evidence["extended_onnx_max_abs_error"] = errors + + @unittest.skipUnless( + os.environ.get("GO2_PRETRAINED_GPU_SMOKE") == "1", "opt-in <=4env x 1iteration GPU smoke" + ) + def test_real_observations_and_one_ppo_iteration(self): + import torch + import warp as wp + + if not hasattr(wp, "context"): + from warp._src import context + + wp.context = context + from mjlab.envs import ManagerBasedRlEnv + from mjlab.rl import RslRlVecEnvWrapper + from pretrained import initialize_runner, make_reference_actor, validate_runtime_contract + from src.tasks.obstacle_avoidance.env_cfg import unitree_go2_obstacle_env_cfg + from src.tasks.velocity.config.go2.rl_cfg import unitree_go2_ppo_runner_cfg + from src.tasks.velocity.rl.runner import VelocityOnPolicyRunner + from task_config import OBSTACLE_TASK, deployment_metadata, validate_task_config + from tensordict import TensorDict + + env_cfg = unitree_go2_obstacle_env_cfg() + env_cfg.scene.num_envs = 4 + env_cfg.seed = 42 + agent_cfg = unitree_go2_ppo_runner_cfg() + agent_cfg.logger = "tensorboard" + agent_cfg.max_iterations = 1 + with tempfile.TemporaryDirectory() as directory: + raw = ManagerBasedRlEnv(env_cfg, device="cuda:0") + self.evidence["compiled_runtime"] = validate_runtime_contract(raw) + raw.platform_deployment = deployment_metadata( + OBSTACLE_TASK, validate_task_config(OBSTACLE_TASK, {}, 42), 42 + ) + env = RslRlVecEnvWrapper(raw) + try: + runner = VelocityOnPolicyRunner(env, asdict(agent_cfg), directory, "cuda:0") + manifest = initialize_runner(runner, self.source) + actor = runner.alg.actor + cpu_actor = make_reference_actor().eval() + cpu_actor.load_state_dict(self.source.actor_state) + obs = env.get_observations() + base = obs["actor"][:, :47].cpu() + with torch.no_grad(): + expected = cpu_actor(TensorDict({"actor": base}, batch_size=[4])) + before = actor(obs).cpu() + torch.testing.assert_close(before, expected, atol=2e-5, rtol=2e-5) + old_mean = actor.obs_normalizer._mean.clone() + actor.update_normalization(obs) + mean_change = ( + (actor.obs_normalizer._mean[:, :47] - old_mean[:, :47]).abs().max().item() + ) + after = actor(obs).detach().cpu() + self.assertTrue(torch.isfinite(after).all()) + torch.testing.assert_close(before, after, atol=2e-5, rtol=2e-5) + loss = actor(obs).square().mean() + loss.backward() + grad = actor.mlp[0].weight.grad[:, 47:] + self.assertTrue(torch.isfinite(grad).all()) + self.assertGreater(grad.abs().max().item(), 0) + grad_max = grad.abs().max().item() + runner.alg.optimizer.zero_grad() + critic_before = copy.deepcopy(runner.alg.critic.state_dict()) + runner.learn(num_learning_iterations=1, init_at_random_ep_len=True) + self.assertGreater(actor.mlp[0].weight[:, 47:].abs().max().item(), 0) + self.assertTrue(runner.alg.optimizer.state) + self.assertTrue( + any( + not torch.equal(v, runner.alg.critic.state_dict()[k]) + for k, v in critic_before.items() + ) + ) + saved = torch.load( + Path(directory) / "model_0.pt", map_location="cpu", weights_only=True + ) + # Same-trial resume uses the installed PPO loader: restore all states. + actor_state = copy.deepcopy(actor.state_dict()) + # Production promotion starts a fresh process/runner; do not reload + # into tensors created by the previous rollout's inference_mode. + resumed = VelocityOnPolicyRunner(env, asdict(agent_cfg), directory, "cuda:0") + self.assertTrue(resumed.alg.load(saved, None, strict=True)) + resumed.current_learning_iteration = saved["iter"] + 1 + self.assertEqual(resumed.current_learning_iteration, 1) + self.assertTrue(resumed.alg.optimizer.state) + actor = resumed.alg.actor + for k, v in actor_state.items(): + self.assertTrue(torch.equal(v.cpu(), actor.state_dict()[k].cpu()), k) + import numpy as np + import onnxruntime as ort + + session = ort.InferenceSession( + str(Path(directory) / "policy.onnx"), providers=["CPUExecutionProvider"] + ) + final_obs = env.get_observations() + with torch.no_grad(): + expected_final = actor(final_obs).cpu().numpy() + actual = np.concatenate( + [ + session.run(None, {"obs": row[None].cpu().numpy()})[0] + for row in final_obs["actor"] + ] + ) + np.testing.assert_allclose(actual, expected_final, atol=2e-5, rtol=2e-5) + # Independently run source ORT on observations actually sampled from MuJoCo. + source_session = ort.InferenceSession( + str(Path(os.environ["GO2_PRETRAINED_SOURCE"]) / "policy.onnx"), + providers=["CPUExecutionProvider"], + ) + source_actions = np.concatenate( + [source_session.run(None, {"obs": row[None].numpy()})[0] for row in base] + ) + np.testing.assert_allclose(source_actions, before.numpy(), atol=2e-5, rtol=2e-5) + self.evidence["gpu_smoke"] = { + "num_envs": 4, + "iterations": 1, + "initialization": manifest, + "real_observation_source_onnx_max_abs_error": float( + np.abs(source_actions - before.numpy()).max() + ), + "normalization_batch_action_max_abs_change": (after - before) + .abs() + .max() + .item(), + "normalization_batch_base_mean_max_abs_change": mean_change, + "new_columns_gradient_max": grad_max, + "new_columns_weight_max_after_ppo": actor.mlp[0] + .weight[:, 47:] + .abs() + .max() + .item(), + "export_max_abs_error": float(np.abs(actual - expected_final).max()), + "count_restored": actor.obs_normalizer.count.item(), + "all_actor_tensors_restored_exactly": True, + } + finally: + env.close() + + +if __name__ == "__main__": + unittest.main() diff --git a/training_server/tests/test_pretrained_sources.py b/training_server/tests/test_pretrained_sources.py new file mode 100644 index 00000000..718f90d1 --- /dev/null +++ b/training_server/tests/test_pretrained_sources.py @@ -0,0 +1,307 @@ +"""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() diff --git a/training_server/tests/test_pretrained_upload.py b/training_server/tests/test_pretrained_upload.py new file mode 100644 index 00000000..dd5ad48c --- /dev/null +++ b/training_server/tests/test_pretrained_upload.py @@ -0,0 +1,735 @@ +"""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() diff --git a/training_server/tests/test_pretrained_upload_runner.py b/training_server/tests/test_pretrained_upload_runner.py new file mode 100644 index 00000000..d7e147e5 --- /dev/null +++ b/training_server/tests/test_pretrained_upload_runner.py @@ -0,0 +1,163 @@ +"""Opt-in real CPU runner integration: four environments, one ONNX-derived PPO iteration.""" + +import copy +import io +import json +import os +import sys +import tempfile +import unittest +from dataclasses import asdict +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[1] +for root in (ROOT, ROOT / "rl"): + sys.path.insert(0, str(root)) + + +@unittest.skipUnless(os.environ.get("GO2_UPLOAD_RUNNER_OUTPUT"), "opt-in four-env CPU runner smoke") +class UploadedRunnerTest(unittest.TestCase): + def test_real_upload_runner_ppo_export_and_same_trial_restore(self): + import numpy as np + import onnxruntime as ort + import torch + import warp as wp + + if not hasattr(wp, "context"): + from warp._src import context + + wp.context = context + from mjlab.envs import ManagerBasedRlEnv + from mjlab.rl import RslRlVecEnvWrapper + from pretrained import initialize_runner, make_reference_actor, validate_runtime_contract + from pretrained_sources import PretrainedSources + from pretrained_upload import read_uploaded_source + from src.tasks.obstacle_avoidance.env_cfg import unitree_go2_obstacle_env_cfg + from src.tasks.velocity.config.go2.rl_cfg import unitree_go2_ppo_runner_cfg + from src.tasks.velocity.rl.runner import VelocityOnPolicyRunner + from task_config import OBSTACLE_TASK, deployment_metadata, validate_task_config + + torch.set_num_threads(1) + source_root = Path(os.environ["GO2_UPLOAD_REAL_DIR"]) + output = Path(os.environ["GO2_UPLOAD_RUNNER_OUTPUT"]) + output.mkdir(parents=True, exist_ok=True) + # Only these explicit files; each decoder sees its independent empty root. + original_onnx = (source_root / "policy.onnx").read_bytes() + options = ort.SessionOptions() + options.intra_op_num_threads = options.inter_op_num_threads = 1 + original = ort.InferenceSession(original_onnx, options, providers=["CPUExecutionProvider"]) + evidence = {} + for fmt, data in ( + ("pt", (source_root / "model_10000.pt").read_bytes()), + ("onnx", original_onnx), + ): + with tempfile.TemporaryDirectory() as store: + registry = PretrainedSources(None, Path(store), sys.executable, ROOT / "rl") + record = registry.receive_upload( + io.BytesIO(data), len(data), fmt, "go2-legacy47-v1", f"single.{fmt}" + ) + bound = record["initialization"] + directory = registry.verify(bound) + cfg = unitree_go2_obstacle_env_cfg() + cfg.scene.num_envs = 4 + cfg.seed = 42 + agent = unitree_go2_ppo_runner_cfg() + agent.logger = "tensorboard" + agent.max_iterations = 1 + source = read_uploaded_source( + directory / "actor.pt", + allowed_roots=[directory], + manifest_path=directory / "upload.json", + target_env=asdict(cfg), + target_agent=asdict(agent), + ) + raw = ManagerBasedRlEnv(cfg, device="cpu") + env = RslRlVecEnvWrapper(raw) + try: + validate_runtime_contract(raw) + raw.platform_deployment = deployment_metadata( + OBSTACLE_TASK, validate_task_config(OBSTACLE_TASK, {}, 42), 42 + ) + log = output / fmt + log.mkdir(exist_ok=True) + runner = VelocityOnPolicyRunner(env, asdict(agent), str(log), "cpu") + raw.platform_initialization = initialize_runner(runner, source) + (log / "initialization.json").write_text( + json.dumps(raw.platform_initialization) + ) + actor = runner.alg.actor + obs = env.get_observations() + reference = make_reference_actor().eval() + reference.load_state_dict(source.actor_state) + with torch.no_grad(): + before = actor(obs) + expected = reference.mlp(reference.obs_normalizer(obs["actor"][:, :47])) + torch.testing.assert_close(before, expected, atol=2e-5, rtol=2e-5) + onnx_actions = np.concatenate( + [ + original.run(None, {"obs": row[None].numpy()})[0] + for row in obs["actor"][:, :47] + ] + ) + np.testing.assert_allclose(before.numpy(), onnx_actions, atol=2e-5, rtol=2e-5) + self.assertEqual(actor.mlp[0].weight[:, 47:].count_nonzero().item(), 0) + self.assertFalse(runner.alg.optimizer.state) + actor(obs).square().mean().backward() + gradient = actor.mlp[0].weight.grad[:, 47:] + self.assertTrue(torch.isfinite(gradient).all()) + self.assertGreater(gradient.abs().max().item(), 0) + evidence[fmt] = { + "source_id": record["id"], + "real_observation_ort_error": float( + np.abs(before.numpy() - onnx_actions).max() + ), + "new_column_gradient_max": gradient.abs().max().item(), + } + runner.alg.optimizer.zero_grad() + if fmt == "pt": + continue # Only ONNX-derived branch runs the one approved PPO iteration. + runner.learn(num_learning_iterations=1, init_at_random_ep_len=True) + self.assertGreater(actor.mlp[0].weight[:, 47:].abs().max().item(), 0) + saved = torch.load(log / "model_0.pt", map_location="cpu", weights_only=True) + updated = copy.deepcopy(actor.state_dict()) + resumed = VelocityOnPolicyRunner(env, asdict(agent), str(log), "cpu") + self.assertTrue(resumed.alg.load(saved, None, strict=True)) + self.assertTrue(resumed.alg.optimizer.state) + for key, tensor in updated.items(): + self.assertTrue( + torch.equal(tensor, resumed.alg.actor.state_dict()[key]), key + ) + self.assertGreater(resumed.alg.actor.obs_normalizer.count.item(), 1_000_000) + exported = ort.InferenceSession( + str(log / "policy.onnx"), options, providers=["CPUExecutionProvider"] + ) + metadata = json.loads( + exported.get_modelmeta().custom_metadata_map["pretrained_initialization"] + ) + self.assertEqual(metadata["sourceFormat"], "onnx") + self.assertEqual( + metadata["uploadSha256"], source.manifest["artifacts"]["upload"]["sha256"] + ) + self.assertNotIn(str(source_root), json.dumps(metadata)) + final_obs = env.get_observations() + with torch.no_grad(): + expected = resumed.alg.actor(final_obs).numpy() + actual = np.concatenate( + [ + exported.run(None, {"obs": row[None].numpy()})[0] + for row in final_obs["actor"] + ] + ) + np.testing.assert_allclose(actual, expected, atol=2e-5, rtol=2e-5) + evidence[fmt].update( + ppo_iterations=1, + num_envs=4, + count_restored=resumed.alg.actor.obs_normalizer.count.item(), + export_max_error=float(np.abs(actual - expected).max()), + new_columns_max_after_ppo=actor.mlp[0].weight[:, 47:].abs().max().item(), + source_metadata=metadata, + ) + finally: + env.close() + (output / "evidence.json").write_text(json.dumps(evidence, indent=2)) + print("REAL_RUNNER_UPLOAD_EVIDENCE=" + json.dumps(evidence)) diff --git a/training_server/tests/test_reward_preset_tasks.py b/training_server/tests/test_reward_preset_tasks.py new file mode 100644 index 00000000..56bf5595 --- /dev/null +++ b/training_server/tests/test_reward_preset_tasks.py @@ -0,0 +1,137 @@ +"""Persisted preset identity -> resolver -> training validation; no training/API calls.""" + +import json +import sys +import tempfile +import unittest +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from server import ApiError, TrainingManager # noqa: E402 +from tuning.manager import TuningManager # noqa: E402 +from tuning.schema import ( # noqa: E402 + FLAT_TASK, + OBSTACLE_TASK, + RewardConfigError, + base_configuration, +) +from tuning.storage import TuningStorage # noqa: E402 + + +class RewardPresetTaskTest(unittest.TestCase): + def setUp(self): + self.temp = tempfile.TemporaryDirectory() + root = Path(self.temp.name) + (root / "scripts").mkdir() + (root / "scripts/train.py").write_text("raise AssertionError('must not train')") + self.storage = TuningStorage(root / "tuning.sqlite3") + # Real resolver with real storage; no advisor, processes or training workers are needed. + self.tuning = object.__new__(TuningManager) + self.tuning.storage = self.storage + self.training = TrainingManager( + root, sys.executable, (FLAT_TASK, OBSTACLE_TASK), check_environment=False + ) + self.training.preset_resolver = self.tuning.preset_config + + def tearDown(self): + self.storage.connection().close() + self.temp.cleanup() + + def save(self, config, reward=None): + session = self.storage.create_session("approval", config, {}, False) + reward = reward or base_configuration(config.get("taskId", FLAT_TASK)) + trial = self.storage.create_trial(session["id"], 0, 0, 1, reward, None, "trial") + return self.storage.save_preset(session["id"], session["id"], trial["id"], reward) + + def payload(self, preset): + return dict( + taskId=FLAT_TASK, + rewardPresetId=preset["id"], + numEnvs=2, + maxIterations=1, + seed=42, + device="cpu", + ) + + def assert_rejected(self, preset): + for entry in (self.training.parse_config, self.training.start): + with self.assertRaises(ApiError) as raised: + entry(self.payload(preset)) + self.assertEqual(raised.exception.status, 400) + self.assertEqual(self.training.jobs, {}) + self.assertIsNone(self.training.lease.public()) + + def test_obstacle_list_identity_and_cross_task_rejected_before_job_creation(self): + preset = self.save({"taskId": OBSTACLE_TASK}) + self.assertEqual(preset["taskId"], OBSTACLE_TASK) + self.assertEqual(self.storage.list_presets(), [preset]) + self.assertEqual(self.storage.get_preset(preset["id"]), preset) + self.assertEqual( + self.tuning.preset_config(preset["id"], OBSTACLE_TASK), preset["rewardConfig"] + ) + with self.assertRaisesRegex(RewardConfigError, "任务"): + self.tuning.preset_config(preset["id"]) + self.assert_rejected(preset) + + def test_historical_flat_without_task_id_and_explicit_flat_remain_valid(self): + for config in ({}, {"taskId": FLAT_TASK}): + with self.subTest(config=config): + preset = self.save(config) + self.assertEqual(preset["taskId"], FLAT_TASK) + self.assertEqual(self.tuning.preset_config(preset["id"]), preset["rewardConfig"]) + parsed = self.training.parse_config(self.payload(preset)) + self.assertEqual(parsed.reward_config, base_configuration()) + self.assertIn("--reward-config-json", self.training.command_for(parsed)) + self.assertEqual({p["taskId"] for p in self.storage.list_presets()}, {FLAT_TASK}) + + def test_malformed_preset_and_corrupt_or_missing_source_fail_closed(self): + preset = self.save({}) + connection = self.storage.connection() + valid_reward = json.dumps(preset["rewardConfig"]) + for bad_reward in ( + '{"weights":{},"params":{}}', + "null", + "{broken", + json.dumps(base_configuration(OBSTACLE_TASK)), + ): + with self.subTest(reward=bad_reward): + connection.execute( + "UPDATE presets SET reward_config_json=? WHERE id=?", (bad_reward, preset["id"]) + ) + with self.assertRaises(RewardConfigError): + self.storage.list_presets() + self.assert_rejected(preset) + connection.execute( + "UPDATE presets SET reward_config_json=? WHERE id=?", (valid_reward, preset["id"]) + ) + for source in ( + "null", + "[]", + "{broken", + '{"taskId":null}', + '{"taskId":"unknown"}', + json.dumps({"taskId": OBSTACLE_TASK}), + ): + with self.subTest(source=source): + connection.execute( + "UPDATE sessions SET config_json=? WHERE id=?", (source, preset["sessionId"]) + ) + with self.assertRaises(RewardConfigError): + self.storage.get_preset(preset["id"]) + self.assert_rejected(preset) + connection.execute("DELETE FROM sessions WHERE id=?", (preset["sessionId"],)) + with self.assertRaises(RewardConfigError): + self.storage.list_presets() + self.assert_rejected(preset) + + def test_save_and_service_both_validate_complete_task_schema(self): + with self.assertRaises(RewardConfigError): + self.save({}, base_configuration(OBSTACLE_TASK)) + self.assertEqual(self.storage.list_presets(), []) + # The service itself must reject partial/malformed configs even from a broken resolver. + self.training.preset_resolver = lambda _id, _task: {"weights": {"pose": 1}, "params": {}} + self.assert_rejected({"id": "f" * 32}) + + +if __name__ == "__main__": + unittest.main() diff --git a/training_server/tests/test_server.py b/training_server/tests/test_server.py index b0d7f8cf..796d588b 100644 --- a/training_server/tests/test_server.py +++ b/training_server/tests/test_server.py @@ -10,6 +10,7 @@ from unittest.mock import patch sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from server import ( # noqa: E402 + DEFAULT_TASKS, MAX_JOBS, ApiError, TrainingJob, @@ -79,6 +80,70 @@ out.write_bytes(b'onnx') with self.assertRaises(ApiError): self.manager.parse_config(self.payload(wandbMode="login")) + def test_custom_task_metadata_and_validated_deployment(self): + self.manager.tasks = DEFAULT_TASKS + health = self.manager.health() + task = next( + item for item in health["taskMetadata"] if item["id"] == "Unitree-Go2-ObstacleAvoidance" + ) + self.assertEqual(task["sensorTypes"], ["raycast"]) + config = self.manager.parse_config( + self.payload( + taskId=task["id"], + terrainPreset="discrete_obstacles", + terrainParams={"obstacle_count": 8, "friction": 0.9}, + sensorCfg={"type": "raycast", "fov": 100, "maxDistance": 5}, + ) + ) + self.assertEqual(config.deployment["observationSize"], 81) + self.assertEqual(config.deployment["terrain"]["actualObstacleCount"], 8) + self.assertEqual(config.deployment["sensorCfg"]["rayCount"], 32) + self.assertEqual( + TrainingJob(id="a" * 32, config=config).public()["deployment"], config.deployment + ) + rough = self.manager.parse_config(self.payload(taskId="Unitree-Go2-Rough")) + self.assertFalse(rough.deployment["browserCompatible"]) + self.assertEqual(rough.deployment["observationSize"], 234) + + def test_custom_config_rejects_unknown_nonfinite_and_incompatible_fields(self): + self.manager.tasks = DEFAULT_TASKS + for fields in ( + {"terrainPreset": "plane;touch /tmp/injected"}, + {"terrainPreset": ["plane"]}, + {"terrainParams": {"friction": float("nan")}}, + {"terrainParams": {"size": float("inf")}}, + {"terrainParams": {"size": 10**400}}, + {"terrainParams": {"obstacle_count": True}}, + {"terrainParams": {"obstacle_count": 1.5}}, + {"terrainParams": {"mjcf": ""}}, + {"terrainParams": {"obstacle_height_min": 1, "obstacle_height_max": 0.1}}, + {"sensorCfg": {"rayCount": 1_000_000}}, + {"sensorCfg": {"maxDistance": 1, "safetyDistance": 1}}, + {"sensorCfg": {"fov": float("nan")}}, + {"sensorType": "camera_depth"}, + {"rewardPresetId": "a" * 32}, + ): + with self.subTest(fields=fields), self.assertRaises(ApiError): + self.manager.parse_config( + self.payload(taskId="Unitree-Go2-ObstacleAvoidance", **fields) + ) + with self.assertRaises(ApiError): + self.manager.parse_config(self.payload(sensorCfg={"fov": 90})) + + def test_custom_job_passes_server_owned_json_file_without_shell(self): + self.manager.tasks = DEFAULT_TASKS + config = self.manager.parse_config(self.payload(taskId="Unitree-Go2-ObstacleAvoidance")) + job = TrainingJob(id="b" * 32, config=config) + self.manager.jobs[job.id] = job + with patch("server.subprocess.Popen", wraps=subprocess.Popen) as popen: + self.manager._run(job) + args = popen.call_args.args[0] + self.assertNotIn("shell", popen.call_args.kwargs) + path = Path(args[args.index("--task-config") + 1]) + self.assertTrue(path.is_relative_to(self.root)) + self.assertEqual(json.loads(path.read_text()), config.task_config) + self.assertEqual(job.state, "succeeded") + def test_builds_argument_array_without_shell(self): config = self.manager.parse_config(self.payload(device="gpu", gpuIds=[0, 2])) command = self.manager.command_for(config) @@ -89,8 +154,13 @@ out.write_bytes(b'onnx') def test_resolves_reward_preset_to_inline_validated_trainer_argument(self): preset_id = "f" * 32 - reward_config = {"weights": {"pose": 1.2}, "params": {}} - self.manager.preset_resolver = lambda value: reward_config if value == preset_id else None + from tuning.schema import base_configuration + + reward_config = base_configuration() + reward_config["weights"]["pose"] = 1.2 + self.manager.preset_resolver = lambda value, task: ( + reward_config if value == preset_id and task == "Unitree-Go2-Flat" else None + ) config = self.manager.parse_config(self.payload(rewardPresetId=preset_id)) command = self.manager.command_for(config) index = command.index("--reward-config-json") diff --git a/training_server/tests/test_task_config.py b/training_server/tests/test_task_config.py new file mode 100644 index 00000000..2686c919 --- /dev/null +++ b/training_server/tests/test_task_config.py @@ -0,0 +1,83 @@ +import json +import math +import sys +import unittest +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from task_config import ( # noqa: E402 + FLAT_TASK, + OBSTACLE_TASK, + TERRAIN_PRESETS, + build_terrain_layout, + deployment_metadata, + navigation_candidates, + validate_task_config, +) + + +class TaskConfigTest(unittest.TestCase): + def test_flat_legacy_request_does_not_override_environment(self): + self.assertIsNone(validate_task_config(FLAT_TASK, {}, 42)) + self.assertNotIn("terrain", deployment_metadata(FLAT_TASK, None, 42)) + + def test_layouts_are_bounded_finite_and_leave_flat_spawn_goal(self): + for preset in (p for p in TERRAIN_PRESETS if p != "custom_boxes"): + with self.subTest(preset=preset): + custom = validate_task_config(OBSTACLE_TASK, {"terrainPreset": preset}, 42) + layout = build_terrain_layout(custom) + self.assertLessEqual(len(layout["boxes"]), 257) + self.assertLess(len(json.dumps(layout)), 64 * 1024) + self.assertEqual(layout["spawn"], [-5, 0, 0.32]) + self.assertEqual(layout["target"], [5, 0]) + for box in layout["boxes"]: + self.assertTrue(all(math.isfinite(v) for v in box["pos"] + box["size"])) + self.assertTrue(all(v > 0 for v in box["size"])) + self.assertEqual(box["yaw"], 0) + for box in layout["boxes"][1:]: + self.assertGreaterEqual(box["pos"][0] - box["size"][0], -4) + self.assertLessEqual(box["pos"][0] + box["size"][0], 4) + + def test_layout_is_deterministic_and_export_is_authoritative(self): + custom = validate_task_config(OBSTACLE_TASK, {}, 42) + original = build_terrain_layout(custom) + self.assertEqual(original, build_terrain_layout(custom)) + self.assertNotEqual(original, build_terrain_layout({**custom, "seed": 43})) + metadata = deployment_metadata(OBSTACLE_TASK, custom, 42) + self.assertEqual(metadata["terrain"], original) + self.assertEqual(metadata["observationSize"], 47 + 32 + 2) + self.assertEqual(metadata["jointNames"][0], "FL_hip_joint") + self.assertEqual(metadata["navigation"]["distanceScale"], 12) + self.assertTrue(metadata["sensorCfg"]["includeGround"]) + self.assertEqual(metadata["navigation"]["trainingReset"], "random-connected-free-pair") + + def test_navigation_candidates_are_safe_connected_and_have_distant_fallbacks(self): + layout = build_terrain_layout(validate_task_config(OBSTACLE_TASK, {}, 42)) + candidates = navigation_candidates(layout) + self.assertGreater(len(candidates["points"]), 2) + self.assertEqual(len(candidates["componentStarts"]), len(candidates["fallbackPairs"])) + for point in candidates["points"]: + self.assertTrue(all(abs(value) <= layout["size"] / 2 - 0.55 for value in point)) + for box in layout["boxes"][1:]: + distance = math.hypot( + max(abs(point[0] - box["pos"][0]) - box["size"][0], 0), + max(abs(point[1] - box["pos"][1]) - box["size"][1], 0), + ) + self.assertGreater(distance, candidates["clearance"]) + for first, second in candidates["fallbackPairs"]: + self.assertGreaterEqual(math.dist(first, second), candidates["minDistance"]) + + def test_obstacle_count_is_capped_by_spacing_capacity(self): + custom = validate_task_config( + OBSTACLE_TASK, + { + "terrainParams": {"size": 8, "spacing": 3, "obstacle_count": 100}, + }, + 0, + ) + layout = build_terrain_layout(custom) + self.assertEqual(layout["actualObstacleCount"], 2) + + +if __name__ == "__main__": + unittest.main() diff --git a/training_server/tuning/advisor.py b/training_server/tuning/advisor.py index afeffed5..d4d0beb2 100644 --- a/training_server/tuning/advisor.py +++ b/training_server/tuning/advisor.py @@ -12,7 +12,8 @@ from .schema import validate_proposal SYSTEM_PROMPT = """你是 Unitree Go2 强化学习奖励调参专家。 只根据提供的数值配置、训练曲线摘要和固定评估结果提出下一轮稀疏修改。 -必须优先保持速度跟踪与跌倒安全门槛;每轮最多修改四个白名单标量,不得改变符号、函数、传感器或结构。 +必须优先保持当前任务的客观安全门槛:Flat速度跟踪/跌倒,Obstacle无碰撞到达/跌倒。 +每轮最多修改四个白名单标量,不得改变符号、函数、传感器、评估协议或结构。 不要建议 Python 代码、命令、文件路径或白名单外参数。输出必须符合 RewardProposal schema。 """ @@ -101,7 +102,9 @@ class DeepSeekAdvisor: result = self._agent().run_sync(prompt) output = result.output patch = validate_proposal( - {"weights": dict(output.weights), "params": dict(output.params)}, previous + {"weights": dict(output.weights), "params": dict(output.params)}, + previous, + task_id=context.get("task", "Unitree-Go2-Flat"), ) try: usage = result.usage() diff --git a/training_server/tuning/manager.py b/training_server/tuning/manager.py index d2306012..2454829c 100644 --- a/training_server/tuning/manager.py +++ b/training_server/tuning/manager.py @@ -13,11 +13,20 @@ from copy import deepcopy from pathlib import Path from typing import Any +from pretrained_sources import PretrainedSources, SourceError +from task_config import TaskConfigError, validate_task_config + +from . import obstacle_scoring from .advisor import AdvisorUnavailable, DeepSeekAdvisor from .process import GpuLease, ResourceBusyError, terminate_process from .schema import ( - BASE_REWARD_CONFIGURATION, + FLAT_TASK, + OBSTACLE_TASK, + RewardConfigError, + base_configuration, merge_proposal, + task_specs, + validate_configuration, validate_configuration_constraints, validate_constraints, validate_proposal, @@ -51,6 +60,7 @@ class TuningManager: data_root: Path, lease: GpuLease, advisor: Any | None = None, + sources: PretrainedSources | None = None, ): self.trainer_root = trainer_root.expanduser().resolve() self.python = python @@ -66,10 +76,22 @@ class TuningManager: self.workers: dict[str, threading.Thread] = {} self.processes: dict[str, subprocess.Popen[str]] = {} self.cancel_events: dict[str, threading.Event] = {} + self.sources = sources def capability(self) -> dict[str, Any]: capability = self.advisor.capability() - capability.update({"ready": (self.trainer_root / "scripts" / "evaluate.py").is_file()}) + capability.update( + { + "ready": (self.trainer_root / "scripts" / "evaluate.py").is_file(), + "pretrainedSources": self.sources.catalog() if self.sources else [], + "pretrainedUpload": { + "enabled": self.sources is not None, + "templateId": "go2-legacy47-v1", + "formats": {"pt": 256 * 1024**2, "onnx": 64 * 1024**2}, + "endpoint": "/api/training/pretrained-sources/upload", + }, + } + ) return capability @staticmethod @@ -85,8 +107,34 @@ class TuningManager: mode = payload.get("mode", "automatic") if mode not in ("automatic", "approval"): raise TuningError("mode 必须是 automatic 或 approval") - if payload.get("taskId", "Unitree-Go2-Flat") != "Unitree-Go2-Flat": - raise TuningError("第一版只支持 Unitree-Go2-Flat") + task_id = payload.get("taskId", "Unitree-Go2-Flat") + task_specs(task_id) + allowed = { + "taskId", + "mode", + "runName", + "gpuIds", + "trialCount", + "initialIterations", + "middleIterations", + "finalIterations", + "numEnvs", + "seed", + "evalNumEnvs", + "evalSteps", + "earlyStopPatience", + "objectiveWeights", + "fallbackEnabled", + "terrainPreset", + "terrainParams", + "sensorType", + "sensorCfg", + "customTerrainBoxes", + "taskConfig", + "pretrainedSourceId", + } + if payload.keys() - allowed: + raise TuningError("未知session字段") run_name = payload.get("runName", "auto-tune") if not isinstance(run_name, str) or not RUN_NAME.fullmatch(run_name): raise TuningError("runName 格式无效") @@ -105,7 +153,7 @@ class TuningManager: rung1 = self._integer(payload, "middleIterations", 900, rung0, 1000000) rung2 = self._integer(payload, "finalIterations", 2000, rung1, 1000000) config = { - "taskId": "Unitree-Go2-Flat", + "taskId": task_id, "numEnvs": self._integer(payload, "numEnvs", 4096, 1, 16384), "seed": self._integer(payload, "seed", 42, 0, 2147483647), "runName": run_name, @@ -117,12 +165,78 @@ class TuningManager: "evalSteps": self._integer(payload, "evalSteps", 1000, 10, 100000), "earlyStopPatience": self._integer(payload, "earlyStopPatience", 4, 1, 20), } - objective = validate_objective_weights( - payload.get("objectiveWeights", DEFAULT_OBJECTIVE_WEIGHTS) - ) + if task_id == OBSTACLE_TASK: + if config["evalSteps"] != obstacle_scoring.STEPS: + raise TuningError("避障评估固定1000步,不能调整") + objective = deepcopy(obstacle_scoring.WEIGHTS) + if "objectiveWeights" in payload and payload["objectiveWeights"] != objective: + raise TuningError("避障评估权重固定") + task_payload = payload.get("taskConfig", payload) + if "taskConfig" in payload: + if not isinstance(task_payload, dict) or task_payload.keys() - { + "terrainPreset", + "terrainParams", + "sensorCfg", + "sensorType", + "seed", + "customTerrainBoxes", + }: + raise TuningError("taskConfig字段无效") + if any( + k in payload + for k in ( + "terrainPreset", + "terrainParams", + "sensorCfg", + "sensorType", + "customTerrainBoxes", + ) + ): + raise TuningError("taskConfig不能与顶层场景配置混用") + task_seed = task_payload.get("seed", config["seed"]) + if ( + isinstance(task_seed, bool) + or not isinstance(task_seed, int) + or task_seed != config["seed"] + ): + raise TuningError("taskConfig seed必须与session一致") + try: + config["taskConfig"] = validate_task_config(task_id, task_payload, config["seed"]) + except TaskConfigError as error: + raise TuningError(str(error)) from error + base = base_configuration(task_id) + base["weights"]["avoidance_weight"] = config["taskConfig"]["sensorCfg"][ + "avoidanceWeight" + ] + validate_configuration_constraints(base, {}, task_id) + else: + if any( + k in payload + for k in ( + "terrainPreset", + "terrainParams", + "sensorCfg", + "sensorType", + "customTerrainBoxes", + "taskConfig", + ) + ): + raise TuningError("Flat调参不接受避障场景字段") + objective = validate_objective_weights( + payload.get("objectiveWeights", DEFAULT_OBJECTIVE_WEIGHTS) + ) fallback = payload.get("fallbackEnabled", False) if not isinstance(fallback, bool): raise TuningError("fallbackEnabled 必须是布尔值") + if "pretrainedSourceId" in payload: + if self.sources is None: + raise TuningError("服务尚未注册基础策略,请配置--pretrained-sources") + try: + config["pretrained"] = self.sources.bind( + payload["pretrainedSourceId"], task_id, config.get("taskConfig") + ) + except SourceError as error: + raise TuningError(str(error)) from error return mode, config, objective, fallback def create(self, payload: Any) -> dict: @@ -262,8 +376,36 @@ class TuningManager: "--reward-config", str(reward_path), ] + task_path = None + if config["taskId"] == OBSTACLE_TASK: + task_path = run_dir / "task_config.json" + task_path.write_text( + json.dumps(config["taskConfig"], allow_nan=False), encoding="utf-8" + ) + command.extend(("--task-config", str(task_path))) + pretrained = config.get("pretrained") + if pretrained is not None: + if self.sources is None: + raise TuningError("基础策略快照服务未配置,拒绝随机初始化") + self.sources.verify(pretrained) if resume_checkpoint is not None: + if pretrained is not None: + origin_path = resume_checkpoint.parent / "initialization.json" + if ( + not origin_path.is_file() + or origin_path.is_symlink() + or origin_path.stat().st_size > 64 * 1024 + ): + raise TuningError("续训checkpoint缺少有效基础策略来源记录,拒绝丢失来源") + origin = json.loads(origin_path.read_text(encoding="utf-8")) + if ( + origin.get("source_id") != pretrained["sourceId"] + or origin.get("artifacts") != pretrained["manifest"]["artifacts"] + ): + raise TuningError("续训checkpoint基础来源身份不匹配") command.extend(("--resume-checkpoint", str(resume_checkpoint))) + elif pretrained is not None: + command.extend(self.sources.arguments(pretrained)) environment = os.environ.copy() environment["WANDB_MODE"] = "disabled" environment["WANDB_SILENT"] = "true" @@ -304,6 +446,8 @@ class TuningManager: "--gpu-ids", json.dumps(config["gpuIds"], separators=(",", ":")), ] + if task_path is not None: + eval_command.extend(("--task-config", str(task_path))) self.storage.update_trial(trial_id, state="evaluating", message="正在固定协议评估") self.storage.update_session( session_id, state="evaluating", message=f"正在评估 trial {trial['number']}" @@ -316,7 +460,15 @@ class TuningManager: raise TuningError(f"评估失败(返回码 {return_code})") evaluation = json.loads(eval_output.read_text(encoding="utf-8")) baseline_trial = self.storage.list_trials(session_id)[0] - if baseline_trial["evaluation"] is None: + if config["taskId"] == OBSTACLE_TASK: + expected = obstacle_scoring.protocol(config["taskConfig"], config["evalNumEnvs"]) + metrics = obstacle_scoring.validate_evaluation(evaluation, expected) + baseline = baseline_trial["evaluation"] + baseline_metrics = ( + obstacle_scoring.validate_evaluation(baseline, expected) if baseline else None + ) + scored = obstacle_scoring.score_evaluation(metrics, baseline_metrics) + elif baseline_trial["evaluation"] is None: scored = { "eligible": True, "score": 0.0, @@ -394,7 +546,28 @@ class TuningManager: return { "task": session["config"]["taskId"], "objectiveWeights": session["objectiveWeights"], - "allowlist": "服务端将验证固定 schema;最多四项修改", + "allowlist": { + section: { + key: {"min": spec.minimum, "max": spec.maximum} for key, spec in specs.items() + } + for section, specs in zip( + ("weights", "params"), task_specs(session["config"]["taskId"]), strict=True + ) + }, + "taskContext": ( + ( + "97维目标导航:48条前视ray,3x16层pitch=[0,-20,-45]deg;仍有侧后/层间/坑盲区。" + if session["config"]["taskConfig"]["sensorCfg"].get("sensorMode") + == "multi_ring_raycast" + else "81维目标导航:32条前视ray仅单高度切片,存在侧后方/矮障碍/跌落盲区。" + ) + + "避障权重对应近障平方惩罚,collision_penalty为非足端>10N接触惩罚;" + "target_velocity是真实导航command速度而非reward权重。关注擦碰、绕行、目标到达和动作抖动。" + "评估固定seed/1000步/客观权重,不允许修改地图、传感器、起终点、协议或以训练奖励代替指标。" + ) + if session["config"]["taskId"] == OBSTACLE_TASK + else "Flat速度跟踪与步态稳定;保留原六目标评估。", + "taskConfig": session["config"].get("taskConfig"), "parameterConstraints": self.storage.get_control(session["id"])["constraints"], "rejectedFeedback": rejected_feedback, "trials": [ @@ -410,12 +583,14 @@ class TuningManager: ], } - def _fallback_patch(self, previous: dict, index: int) -> dict: + def _fallback_patch(self, previous: dict, index: int, task_id="Unitree-Go2-Flat") -> dict: names = ("track_linear_velocity", "action_rate_l2", "body_orientation_l2", "foot_slip") + if task_id == OBSTACLE_TASK: + names = ("avoidance_weight", "collision_penalty") name = names[index % len(names)] old = previous["weights"][name] factor = 1.1 if index % 2 == 0 else 0.9 - return validate_proposal({"weights": {name: old * factor}}, previous) + return validate_proposal({"weights": {name: old * factor}}, previous, task_id=task_id) def _request_proposal( self, session: dict, previous: dict, base_trial_id: str, index: int @@ -427,7 +602,7 @@ class TuningManager: if not session["fallbackEnabled"]: raise AdvisorUnavailable(str(error)) from error result = { - "patch": self._fallback_patch(previous, index), + "patch": self._fallback_patch(previous, index, session["config"]["taskId"]), "rationale": f"Agent 不可用,显式 fallback:{error}", "expectedImpact": {}, "confidence": 0.2, @@ -437,7 +612,9 @@ class TuningManager: } source = "fallback" constraints = self.storage.get_control(session["id"])["constraints"] - result["patch"] = validate_proposal(result["patch"], previous, constraints) + result["patch"] = validate_proposal( + result["patch"], previous, constraints, session["config"]["taskId"] + ) proposal = self.storage.create_proposal( session["id"], base_trial_id, @@ -471,11 +648,34 @@ class TuningManager: self.condition.wait(timeout=1.0) raise TuningError("session 已取消") + @staticmethod + def _base_configuration(session): + base = base_configuration(session["config"]["taskId"]) + if session["config"]["taskId"] == OBSTACLE_TASK: + base["weights"]["avoidance_weight"] = session["config"]["taskConfig"]["sensorCfg"][ + "avoidanceWeight" + ] + return base + def _run_session(self, session_id: str, resume: bool, cancel: threading.Event) -> None: try: session = self.storage.get_session(session_id) trials = self.storage.list_trials(session_id) if resume: + for proposal in self.storage.list_proposals(session_id): + if proposal["state"] == "pending" and self.storage.decide_proposal( + proposal["id"], + "rejected", + "service_restart/recovery_invalidated:服务重启,需重新提案并审批(非用户拒绝)", + ): + self.storage.audit( + session_id, + "recovery_proposal_invalidated", + { + "proposalId": proposal["id"], + "reason": "service_restart/recovery_invalidated", + }, + ) root = self._session_root(session_id) for interrupted in [trial for trial in trials if trial["state"] == "interrupted"]: run_dir = (root / interrupted["runDir"]).resolve() @@ -499,7 +699,7 @@ class TuningManager: 0, 0, session["config"]["rungs"][0], - deepcopy(BASE_REWARD_CONFIGURATION), + self._base_configuration(session), None, baseline_dir, ) @@ -551,7 +751,12 @@ class TuningManager: proposal = self.storage.get_proposal(proposal["id"]) self._claim_dispatch(session_id, cancel) constraints = self.storage.get_control(session_id)["constraints"] - reward_config = merge_proposal(base["rewardConfig"], proposal["patch"], constraints) + reward_config = merge_proposal( + base["rewardConfig"], + proposal["patch"], + constraints, + session["config"]["taskId"], + ) trial = self.storage.create_trial( session_id, next_number, @@ -698,7 +903,12 @@ class TuningManager: patch = payload["patch"] if feedback is not None and (not isinstance(feedback, str) or len(feedback) > 2000): raise TuningError("feedback 无效") - patch = validate_proposal(patch, base["rewardConfig"], constraints) + patch = validate_proposal( + patch, + base["rewardConfig"], + constraints, + self.storage.get_session(session_id)["config"]["taskId"], + ) if not self.storage.decide_proposal(proposal_id, "approved", feedback, patch): raise TuningError("proposal 已处理") self.storage.audit( @@ -743,7 +953,12 @@ class TuningManager: continue base = self.storage.get_trial(proposal["baseTrialId"]) try: - validate_proposal(proposal["patch"], base["rewardConfig"], constraints) + validate_proposal( + proposal["patch"], + base["rewardConfig"], + constraints, + session["config"]["taskId"], + ) except Exception as error: self.storage.decide_proposal( proposal["id"], "rejected", f"参数护栏已变化:{error}" @@ -771,8 +986,8 @@ class TuningManager: revision = payload["revision"] if isinstance(revision, bool) or not isinstance(revision, int) or revision < 0: raise TuningError("constraints revision 必须是非负整数") - constraints = validate_constraints(payload["constraints"]) session = self.storage.get_session(session_id) + constraints = validate_constraints(payload["constraints"], session["config"]["taskId"]) if session["state"] not in ACTIVE_SESSION_STATES: raise TuningError("终态 session 不能修改参数护栏") control = self.storage.get_control(session_id) @@ -784,7 +999,9 @@ class TuningManager: else: base = self._best_highest_rung(session_id) if base is not None: - validate_configuration_constraints(base["rewardConfig"], constraints) + validate_configuration_constraints( + base["rewardConfig"], constraints, session["config"]["taskId"] + ) try: updated = self.storage.replace_constraints(session_id, revision, constraints) except StorageConflict as error: @@ -796,7 +1013,12 @@ class TuningManager: continue proposal_base = self.storage.get_trial(proposal["baseTrialId"]) try: - validate_proposal(proposal["patch"], proposal_base["rewardConfig"], constraints) + validate_proposal( + proposal["patch"], + proposal_base["rewardConfig"], + constraints, + session["config"]["taskId"], + ) except Exception as error: message = f"参数护栏 revision {updated['constraintsRevision']}:{error}" if self.storage.decide_proposal(proposal["id"], "rejected", message): @@ -914,7 +1136,16 @@ class TuningManager: return self.detail(session_id) def resume(self, session_id: str) -> dict: + # Serialize state transition and worker registration against duplicate requests. + with self.condition: + return self._resume_locked(session_id) + + def _resume_locked(self, session_id: str) -> dict: session = self.storage.get_session(session_id) + if session["config"].get("pretrained") is not None: + if self.sources is None: + raise TuningError("基础策略快照服务未配置,无法恢复") + self.sources.verify(session["config"]["pretrained"]) if session["state"] == "paused": self.storage.set_run_policy(session_id, "continuous") has_pending = any( @@ -982,8 +1213,11 @@ class TuningManager: raise TuningError("最佳策略文件不存在") return path - def preset_config(self, preset_id: str) -> dict: - return self.storage.get_preset(preset_id)["rewardConfig"] + def preset_config(self, preset_id: str, task_id: str = FLAT_TASK) -> dict: + preset = self.storage.get_preset(preset_id) + if preset["taskId"] != task_id: + raise RewardConfigError("奖励 preset 来源任务与训练任务不匹配") + return validate_configuration(preset["rewardConfig"], task_id) def test_agent(self) -> dict: try: diff --git a/training_server/tuning/obstacle_scoring.py b/training_server/tuning/obstacle_scoring.py new file mode 100644 index 00000000..011aac01 --- /dev/null +++ b/training_server/tuning/obstacle_scoring.py @@ -0,0 +1,130 @@ +"""Fixed obstacle protocol; objective measurements never consume training rewards.""" + +import math +from statistics import fmean + +from task_config import build_terrain_layout, validate_task_config + +from .schema import OBSTACLE_TASK +from .scoring import EvaluationError + +SEEDS = (101, 202, 303) +STEPS = 1000 +WEIGHTS = {"success": 0.4, "time": 0.2, "clearance": 0.2, "smooth": 0.1, "no_fall": 0.1} +METRICS = (*WEIGHTS, "collision_rate", "fall_rate", "arrival_rate", "ray_hit_rate") + + +def evaluation_scenarios(custom): + scenarios = [] + for seed in SEEDS: + value = validate_task_config(OBSTACLE_TASK, custom, seed) + scenarios.append( + {"seed": seed, "taskConfig": value, "terrain": build_terrain_layout(value)} + ) + return scenarios + + +def protocol(custom, num_envs): + return { + "protocolVersion": "obstacle-v1", + "seeds": list(SEEDS), + "stepsPerSeed": STEPS, + "numEnvs": num_envs, + "objectiveWeights": WEIGHTS, + "sceneMode": "fixed-custom-map" + if custom["terrainPreset"] == "custom_boxes" + else "three-seed-layouts", + "scenarios": evaluation_scenarios(custom), + "episodePolicy": "first-episode-only; terminal snapshot before auto-reset; fixed horizon", + } + + +def score_trajectory(samples, horizon=STEPS): + """One first-episode trajectory, one sample per policy step including first terminal.""" + if not samples or len(samples) > horizon: + raise EvaluationError("轨迹样本数不完整") + keys = {"distance", "clearance", "action_delta", "ray_hit", "collision", "fall", "terminal"} + for sample in samples: + if set(sample) != keys: + raise EvaluationError("轨迹指标缺失/未知") + for key in keys: + value = sample[key] + if not isinstance(value, (int, float)) or not math.isfinite(value) or value < 0: + raise EvaluationError(f"无效轨迹指标 {key}") + if key in {"ray_hit", "collision", "fall", "terminal"} and value > 1: + raise EvaluationError(f"无效标志 {key}") + if any(s["terminal"] for s in samples[:-1]) or ( + len(samples) != horizon and not samples[-1]["terminal"] + ): + raise EvaluationError("首episode轨迹不完整") + collision = any(s["collision"] for s in samples) + fall = any(s["fall"] for s in samples) + arrival = next((i + 1 for i, s in enumerate(samples) if s["distance"] < 0.5), None) + success = arrival is not None and not collision and not fall + # Missing steps after an early terminal earn zero clearance/smoothness, not a bonus. + return { + "success": float(success), + "time": 1 - arrival / horizon if success else 0.0, + "clearance": 0.0 if fall else sum(min(s["clearance"] / 0.5, 1) for s in samples) / horizon, + "smooth": sum(1 - min(s["action_delta"], 1) for s in samples) / horizon, + "no_fall": float(not fall), + "collision_rate": float(collision), + "fall_rate": float(fall), + "arrival_rate": float(arrival is not None), + "ray_hit_rate": fmean(s["ray_hit"] for s in samples), + } + + +def validate_metrics(value): + if not isinstance(value, dict) or set(value) != set(METRICS): + raise EvaluationError("避障评估指标缺失/未知") + for key, number in value.items(): + if ( + isinstance(number, bool) + or not isinstance(number, (int, float)) + or not math.isfinite(number) + or not 0 <= number <= 1 + ): + raise EvaluationError(f"避障指标 {key} 必须在0–1且有限") + if not math.isclose(value["no_fall"] + value["fall_rate"], 1, abs_tol=1e-8): + raise EvaluationError("跌倒指标不一致") + return value + + +def validate_evaluation(value, expected_protocol): + if value.get("protocol") != expected_protocol: + raise EvaluationError("避障评估协议/场景不一致") + seeds = value.get("seedMetrics", []) + if len(seeds) != 3: + raise EvaluationError("避障评估必须完整包含3seed") + for item, seed in zip(seeds, SEEDS, strict=True): + if ( + item.get("seed") != seed + or item.get("episodes") != expected_protocol["numEnvs"] + or item.get("rolloutSteps") != STEPS + ): + raise EvaluationError("避障评估样本数不完整") + validate_metrics(item.get("metrics")) + metrics = validate_metrics(value.get("metrics")) + if any( + not math.isclose(metrics[k], fmean(s["metrics"][k] for s in seeds), abs_tol=1e-8) + for k in METRICS + ): + raise EvaluationError("避障评估聚合不一致") + return metrics + + +def score_evaluation(current, baseline=None): + current = validate_metrics(current) + baseline = validate_metrics(baseline) if baseline is not None else current + eligible = ( + current["fall_rate"] <= baseline["fall_rate"] + 0.02 + 1e-12 + and current["success"] >= baseline["success"] - 0.02 - 1e-12 + ) + total = sum(WEIGHTS[k] * current[k] for k in WEIGHTS) + return { + "eligible": eligible, + "score": total if eligible else -1.0, + "components": {k: current[k] for k in WEIGHTS}, + "metrics": current, + } diff --git a/training_server/tuning/schema.py b/training_server/tuning/schema.py index e7a2c819..c68ee715 100644 --- a/training_server/tuning/schema.py +++ b/training_server/tuning/schema.py @@ -68,6 +68,32 @@ BASE_REWARD_CONFIGURATION: dict[str, dict[str, float]] = { } +FLAT_TASK = "Unitree-Go2-Flat" +OBSTACLE_TASK = "Unitree-Go2-ObstacleAvoidance" +OBSTACLE_WEIGHT_SPECS = { + "avoidance_weight": NumericSpec(0.5, 5.0, 2.0, False), + "collision_penalty": NumericSpec(-10.0, -0.5, -5.0, False), + "action_smoothness": NumericSpec(-0.05, -0.001, -0.05, False), +} +OBSTACLE_PARAMETER_SPECS = {"target_velocity": NumericSpec(0.3, 1.2, 0.6, False)} + + +def task_specs(task_id=FLAT_TASK): + if task_id == FLAT_TASK: + return WEIGHT_SPECS, PARAMETER_SPECS + if task_id == OBSTACLE_TASK: + return OBSTACLE_WEIGHT_SPECS, OBSTACLE_PARAMETER_SPECS + raise RewardConfigError(f"只支持Flat或Obstacle调参任务:{task_id}") + + +def base_configuration(task_id=FLAT_TASK): + weights, params = task_specs(task_id) + return { + "weights": {k: v.default for k, v in weights.items()}, + "params": {k: v.default for k, v in params.items()}, + } + + def _number(name: str, value: Any, spec: NumericSpec) -> float: if isinstance(value, bool) or not isinstance(value, (int, float)): raise RewardConfigError(f"{name} 必须是数值") @@ -87,19 +113,22 @@ def _mapping(value: Any, name: str) -> Mapping[str, Any]: return value -def _cross_validate(config: Mapping[str, Mapping[str, float]]) -> None: +def _cross_validate(config: Mapping[str, Mapping[str, float]], task_id=FLAT_TASK) -> None: + if task_id == OBSTACLE_TASK: + return params = config["params"] if params["pose.walking_threshold"] >= params["pose.running_threshold"]: raise RewardConfigError("pose.walking_threshold 必须小于 pose.running_threshold") -def _path_spec(path: str) -> tuple[str, str, NumericSpec]: +def _path_spec(path: str, task_id=FLAT_TASK) -> tuple[str, str, NumericSpec]: + weights, params = task_specs(task_id) if path.startswith("weights."): section, name = "weights", path.removeprefix("weights.") - spec = WEIGHT_SPECS.get(name) + spec = weights.get(name) elif path.startswith("params."): section, name = "params", path.removeprefix("params.") - spec = PARAMETER_SPECS.get(name) + spec = params.get(name) else: section, name, spec = "", "", None if spec is None: @@ -107,16 +136,16 @@ def _path_spec(path: str) -> tuple[str, str, NumericSpec]: return section, name, spec -def validate_constraints(value: Any) -> dict[str, dict[str, float | str]]: +def validate_constraints(value: Any, task_id=FLAT_TASK) -> dict[str, dict[str, float | str]]: """Validate sparse per-session range/fixed safety constraints.""" root = _mapping(value, "constraints") - if len(root) > len(WEIGHT_SPECS) + len(PARAMETER_SPECS): + if len(root) > sum(len(specs) for specs in task_specs(task_id)): raise RewardConfigError("constraints 数量超过白名单参数总数") result: dict[str, dict[str, float | str]] = {} for raw_path, raw_constraint in root.items(): if not isinstance(raw_path, str): raise RewardConfigError("constraint path 必须是字符串") - _, _, spec = _path_spec(raw_path) + _, _, spec = _path_spec(raw_path, task_id) constraint = _mapping(raw_constraint, raw_path) kind = constraint.get("kind") if kind == "fixed": @@ -137,12 +166,12 @@ def validate_constraints(value: Any) -> dict[str, dict[str, float | str]]: return result -def validate_configuration_constraints(value: Any, constraints: Any) -> None: +def validate_configuration_constraints(value: Any, constraints: Any, task_id=FLAT_TASK) -> None: """Ensure a complete reward configuration satisfies every session constraint.""" - config = validate_configuration(value) - checked = validate_constraints(constraints) + config = validate_configuration(value, task_id) + checked = validate_constraints(constraints, task_id) for path, constraint in checked.items(): - section, name, _ = _path_spec(path) + section, name, _ = _path_spec(path, task_id) current = config[section][name] if constraint["kind"] == "fixed": if current != constraint["value"]: @@ -155,36 +184,37 @@ def validate_configuration_constraints(value: Any, constraints: Any) -> None: ) -def validate_configuration(value: Any) -> dict[str, dict[str, float]]: +def validate_configuration(value: Any, task_id=FLAT_TASK) -> dict[str, dict[str, float]]: """Validate a complete configuration and reject missing/unknown fields.""" root = _mapping(value, "rewardConfig") if set(root) != {"weights", "params"}: raise RewardConfigError("rewardConfig 只能包含 weights 和 params") + weights, params = task_specs(task_id) raw_weights = _mapping(root["weights"], "weights") raw_params = _mapping(root["params"], "params") - if set(raw_weights) != set(WEIGHT_SPECS): + if set(raw_weights) != set(weights): raise RewardConfigError("weights 必须完整且不能包含未知奖励项") - if set(raw_params) != set(PARAMETER_SPECS): + if set(raw_params) != set(params): raise RewardConfigError("params 必须完整且不能包含未知参数") config = { "weights": { name: _number(f"weights.{name}", raw_weights[name], spec) - for name, spec in WEIGHT_SPECS.items() + for name, spec in weights.items() }, "params": { - name: _number(f"params.{name}", raw_params[name], spec) - for name, spec in PARAMETER_SPECS.items() + name: _number(f"params.{name}", raw_params[name], spec) for name, spec in params.items() }, } - _cross_validate(config) + _cross_validate(config, task_id) return config def validate_proposal( - value: Any, previous: Any, constraints: Any | None = None + value: Any, previous: Any, constraints: Any | None = None, task_id=FLAT_TASK ) -> dict[str, dict[str, float]]: """Validate a sparse Agent patch relative to a complete previous config and guardrails.""" - current = validate_configuration(previous) + weights, params = task_specs(task_id) + current = validate_configuration(previous, task_id) root = _mapping(value, "proposal") if not set(root).issubset({"weights", "params"}): raise RewardConfigError("proposal 只能包含 weights 和 params") @@ -194,8 +224,8 @@ def validate_proposal( raise RewardConfigError("proposal 至少需要一项修改") if len(raw_weights) + len(raw_params) > MAX_PROPOSAL_CHANGES: raise RewardConfigError(f"proposal 每轮最多修改 {MAX_PROPOSAL_CHANGES} 项") - unknown_weights = set(raw_weights) - set(WEIGHT_SPECS) - unknown_params = set(raw_params) - set(PARAMETER_SPECS) + unknown_weights = set(raw_weights) - set(weights) + unknown_params = set(raw_params) - set(params) if unknown_weights: raise RewardConfigError(f"未知奖励项:{', '.join(sorted(unknown_weights))}") if unknown_params: @@ -203,7 +233,7 @@ def validate_proposal( patch: dict[str, dict[str, float]] = {"weights": {}, "params": {}} for name, raw in raw_weights.items(): - value_number = _number(f"weights.{name}", raw, WEIGHT_SPECS[name]) + value_number = _number(f"weights.{name}", raw, weights[name]) old = current["weights"][name] if old != 0.0 and value_number != 0.0: ratio = abs(value_number / old) @@ -216,7 +246,7 @@ def validate_proposal( raise RewardConfigError(f"weights.{name} 没有发生变化") patch["weights"][name] = value_number for name, raw in raw_params.items(): - value_number = _number(f"params.{name}", raw, PARAMETER_SPECS[name]) + value_number = _number(f"params.{name}", raw, params[name]) old = current["params"][name] ratio = abs(value_number / old) if ratio < MIN_CHANGE_RATIO or ratio > MAX_CHANGE_RATIO: @@ -230,26 +260,39 @@ def validate_proposal( candidate = deepcopy(current) candidate["weights"].update(patch["weights"]) candidate["params"].update(patch["params"]) - _cross_validate(candidate) + _cross_validate(candidate, task_id) if constraints is not None: - validate_configuration_constraints(candidate, constraints) + validate_configuration_constraints(candidate, constraints, task_id) return patch def merge_proposal( - previous: Any, proposal: Any, constraints: Any | None = None + previous: Any, proposal: Any, constraints: Any | None = None, task_id=FLAT_TASK ) -> dict[str, dict[str, float]]: - current = validate_configuration(previous) - patch = validate_proposal(proposal, current, constraints) + current = validate_configuration(previous, task_id) + patch = validate_proposal(proposal, current, constraints, task_id) merged = deepcopy(current) merged["weights"].update(patch["weights"]) merged["params"].update(patch["params"]) - return validate_configuration(merged) + return validate_configuration(merged, task_id) -def apply_reward_configuration(env_cfg: Any, value: Any) -> None: +def apply_reward_configuration(env_cfg: Any, value: Any, task_id=FLAT_TASK) -> None: """Apply a validated full config to a fresh mjlab environment config.""" - config = validate_configuration(value) + config = validate_configuration(value, task_id) + if task_id == OBSTACLE_TASK: + from mjlab.managers import RewardTermCfg + + contact = env_cfg.terminations["illegal_contact"] + env_cfg.rewards["obstacle_collision"] = RewardTermCfg( + func=contact.func, + params=deepcopy(contact.params), + weight=config["weights"]["collision_penalty"], + ) + env_cfg.rewards["obstacle_proximity"].weight = -config["weights"]["avoidance_weight"] + env_cfg.rewards["action_rate_l2"].weight = config["weights"]["action_smoothness"] + env_cfg.commands["twist"].speed = config["params"]["target_velocity"] + return for name, weight in config["weights"].items(): if name not in env_cfg.rewards: raise RewardConfigError(f"环境缺少奖励项:{name}") diff --git a/training_server/tuning/storage.py b/training_server/tuning/storage.py index cfba8abd..a5dc000a 100644 --- a/training_server/tuning/storage.py +++ b/training_server/tuning/storage.py @@ -12,6 +12,8 @@ from datetime import UTC, datetime from pathlib import Path from typing import Any +from .schema import FLAT_TASK, RewardConfigError, validate_configuration + SCHEMA_VERSION = 2 @@ -169,16 +171,37 @@ class TuningStorage: def recover_interrupted(self) -> None: at = now_iso() with self.transaction() as connection: + previous = connection.execute( + "SELECT id,state FROM sessions WHERE state IN " + "('queued','running','evaluating','paused','awaiting_approval')" + ).fetchall() + for row in previous: + connection.execute( + "INSERT INTO audit_events(session_id,event_type,payload_json,created_at) " + "VALUES (?,?,?,?)", + ( + row["id"], + "service_restart_interrupted", + _json( + { + "previousState": row["state"], + "reason": "service_restart", + "requiresExplicitResume": True, + } + ), + at, + ), + ) connection.execute( "UPDATE trials SET state='interrupted', ended_at=?, " "message='服务重启中断,等待显式恢复' " - "WHERE state IN ('training','evaluating')", + "WHERE state IN ('queued','training','evaluating')", (at,), ) connection.execute( "UPDATE sessions SET state='interrupted', updated_at=?, " "message='服务重启中断,可从完整 checkpoint 恢复' " - "WHERE state IN ('running','evaluating')", + "WHERE state IN ('queued','running','evaluating','paused','awaiting_approval')", (at,), ) @@ -643,47 +666,51 @@ class TuningStorage: for row in rows ] + def _preset_task(self, session_id: str) -> str: + # Identity comes only from the persisted source session, never client/preset labels. + try: + config = self.get_session(session_id)["config"] + except (KeyError, ValueError, TypeError) as error: + raise RewardConfigError("奖励 preset 来源 session 丢失或损坏") from error + if not isinstance(config, dict): + raise RewardConfigError("奖励 preset 来源 config 必须是对象") + # Original Flat-only sessions did not require taskId. Validate their preset below. + return config.get("taskId", FLAT_TASK) + + def _preset(self, row) -> dict: + task_id = self._preset_task(row["session_id"]) + try: + reward_config = validate_configuration(_decode(row["reward_config_json"]), task_id) + except (ValueError, TypeError) as error: + raise RewardConfigError("奖励 preset 配置无效:" + str(error)) from error + return { + "id": row["id"], + "name": row["name"], + "sessionId": row["session_id"], + "trialId": row["trial_id"], + "taskId": task_id, + "rewardConfig": reward_config, + "createdAt": row["created_at"], + } + def save_preset(self, name: str, session_id: str, trial_id: str, reward_config: dict) -> dict: + reward_config = validate_configuration(reward_config, self._preset_task(session_id)) preset_id, at = uuid.uuid4().hex, now_iso() self.connection().execute( "INSERT INTO presets(id,name,session_id,trial_id,reward_config_json,created_at) " "VALUES (?,?,?,?,?,?)", (preset_id, name, session_id, trial_id, _json(reward_config), at), ) - return { - "id": preset_id, - "name": name, - "sessionId": session_id, - "trialId": trial_id, - "rewardConfig": reward_config, - "createdAt": at, - } + return self.get_preset(preset_id) def get_preset(self, preset_id: str) -> dict: row = self.connection().execute("SELECT * FROM presets WHERE id=?", (preset_id,)).fetchone() if row is None: raise KeyError(preset_id) - return { - "id": row["id"], - "name": row["name"], - "sessionId": row["session_id"], - "trialId": row["trial_id"], - "rewardConfig": _decode(row["reward_config_json"]), - "createdAt": row["created_at"], - } + return self._preset(row) def list_presets(self) -> list[dict]: rows = ( self.connection().execute("SELECT * FROM presets ORDER BY created_at DESC").fetchall() ) - return [ - { - "id": row["id"], - "name": row["name"], - "sessionId": row["session_id"], - "trialId": row["trial_id"], - "rewardConfig": _decode(row["reward_config_json"]), - "createdAt": row["created_at"], - } - for row in rows - ] + return [self._preset(row) for row in rows] diff --git a/web_platform/TRAINING.md b/web_platform/TRAINING.md new file mode 100644 index 00000000..052a0cc6 --- /dev/null +++ b/web_platform/TRAINING.md @@ -0,0 +1,80 @@ +# 自定义训练与前视射线避障 + +1. 启动仓库本地训练服务,在「控制台 → 强化学习任务」输入令牌并连接。 +2. 选择「前视射线避障导航」,设置地形、种子、障碍物参数、FOV、探测/安全距离及避障奖励权重,发起训练。界面显示进度、最近价值/策略/熵损失和原始日志。 +3. 「同步当前场景地图」从全部**已应用**实例的已编译MuJoCo静态碰撞几何导出权威 `custom_boxes/boxes-v1`,支持多实例、平移/旋转、工程地图路径及嵌套body。成功显示“已将视口中 N 个自定义障碍物编译为训练地图布局”。不重新运行随机预设,不同步草稿/过时场景。初始姿态取场景加载时机器人位姿,出生高度标准化为.32m;目标建议为初始朝向前方3m,不自动寻找安全点。出生/目标XY可修改,必须重新同步,并保留与所有障碍XY AABB不相交的.5m圆形安全区;绝不删除或移动障碍来避开安全区。 +4. 训练完成后,先导入并加载有对应12个关节的Go2模型(Go2-W仅为非同构实验)。点击「导入策略」会校验作业配置与ONNX内嵌部署契约,在候选场景中完成真实ORT初始化、graph与机器人绑定检查后才事务替换当前**物理**地形、复位到出生点,启动策略与仿真。当前编辑器地图保留,重新加载模型恢复它。失败保留原场景、策略、物理状态及原暂停/播放状态。旧Flat无配套地形时仍按原方式加载后手动启用;新服务的默认Flat作业允许旧外部训练器无metadata导出,但真实graph必须固定47→12;显式自定义地图/避障不允许此降级。 +5. 避障任务自动导航到配套目标,目标半径0.5m内停止速度指令。摄像头PiP自动打开,可单独隐藏;「显示避障射线」可切换红色命中/绿色未命中线段。 + +## 交互式导航目标与训练趋势 + +- 避障策略加载后,主视口显示青色呼吸目标信标,按当前训练碰撞地形的顶部高度落位。ONNX 面板显示当前目标坐标及水平剩余距离。 +- 点击「设定目标」,再在主视口地形单击:目标限制在地图边缘内缩0.5m的区域,并自动退出设定模式。Esc/取消可退出;空白、拖动及摄像头PiP不会更新目标。模式使用捕获阶段拦截输入,不选择刚体或操纵地图Gizmo。 +- 暂停或停用策略时也可更改目标,但不会自动启动仿真或策略。「复位目标点」只恢复训练地图默认目标;「重置仿真」同时恢复默认目标与出生状态。换目标不修改部署JSON、不重建ONNX会话,下一次观测仍保持该策略的81或97维。信标仅表示目标位置,不保证该点可到达或路径无障碍。 +- 训练卡片中的「训练指标趋势」可折叠,提供价值/策略损失及综合页;综合页的熵、平均奖励、平均回合长度各用独立纵轴。复用自调参工作台的轻量Canvas图表,EMA系数0.4仅用于曲线,悬停图例和摘要保留原始值,支持缩放。 +- 按日志中的 `Learning iteration` 合并指标,重复轮询不重复入点,同迭代分批日志可补全;内存最多保留最近500个有指标的迭代。换作业/清除作业会清空。服务器仅保留日志尾部,断线重连不能恢复已丢失历史;没有迭代标题或已知重叠上下文的孤立指标会跳过,不猜测step。缺失指标不显示,原始日志仍可展开。 + +## 自定义场景导出边界 + +- 使用实际碰撞几何而非Three装饰包围盒;box的世界AABB半尺寸为abs(R)×halfsize,sphere/capsule/ellipsoid/cylinder有精确解析bounds。mesh/hfield暂明确拒绝,其他未声明静态碰撞也报错,不能漏掉。既有工程地图include导入限制不变。 +- **AABB会膨胀旋转/非box形状,底板标准化可能填补支撑面之间空白**;custom总是approximation=true,仅保证训练和浏览器部署相同boxes,不保证原OBB/原场景几何同构。仅明确命名且顶面z≈0的水平支撑归并为floor z=[-.2,0];地下/坑底、倾斜或非零高度plane拒绝,不静默填坑。 +- 固定floor加最多256障碍,超量报错不截断。保持世界坐标、覆盖范围8–24m,超出时明确拒绝,不clamp/平移障碍。所有半尺寸须严格正值;统一friction=.2–2且其余摩擦分量为.005/.0001,混合摩擦需先在地图中显式统一。 +- 服务和训练入口都验证完整JSON白名单、尺寸/位置、标准floor、count/approximation一致性及起终点安全区;`terrainPreset=custom_boxes`配套`customTerrainBoxes`布局全程传入任务JSON/deployment/ONNX,不降级预设。旧服务无custom_boxes能力时同步按钮报升级提示。 + +## 明确边界 + +- 默认32条水平前向物理射线;显式multi模式为3×16=48条(pitch 0/-20/-45度),**不是真实camera_depth或RGB视觉策略**。PiP是观察窗口,不作为网络输入。低障碍、跌落与切片外障碍存在感知盲区。 +- 配套`boxes-v1`布局直接来自训练器,使用相同世界坐标、半尺寸、摩擦、出生点和目标;`rough/wave/pyramid_stairs`明确为训练专用box近似,不等于编辑器同名高度场。 +- 浏览器对**编译后的**静态训练box `geom_xpos/geom_xmat/geom_size`做解析slab求交,包括地板。它与MuJoCo box几何求交等价;完整机身姿态旋转射线,按命名与group2仅选训练地形,排除整个机器人。不接受额外无限plane、未声明静态几何或非box训练地形。地图组合先用`mj_saveLastXML`展开include,避免原地面残留。 +- 使用原47维本体观测+32归一化距离+2目标误差,共81维;multi改为47+48+2=97维。保持单in-flight、50Hz异步held-action,不阻塞WASM子步。不另加动态观测归一化。导出器Actor内已有归一化。 +- 原Go2自定义地图部署按`go2_constants.py`补齐hip/thigh/calf armature=0.01/0.01/0.02、零关节阻尼/摩擦损失、非足端condim1/priority0、足端condim3/priority1/solimp宽度0.023,碰撞contype1/conaffinity0关闭自碰撞;足端摩擦使用配套地形。旧无自定义地图的Flat路径不覆盖这些参数。 +- 初始关节姿态写入`MjData.qpos`,不改变模型`qpos0`的关节参考角;重置会恢复出生位姿、初始关节、零速度/历史动作及步态时间。 +- 浏览器评测在20秒、机身高度低于0.12m或越界时停止并报错,需重置后重新启用;**不是训练端的接触/姿态终止自动reset**。到达目标不会强制结束episode。 +- 旧`Unitree-Go2-Rough`的234维观测未在浏览器实现,可训练但禁止一键导入。原Flat、Go2-W速度策略、Python控制器、Flat自调参工作台保持兼容。 +- 模型拓扑及PD配置校验不等于动力学完全一致。短PPO和零动作fixture只能验证链路,不能证明避障收敛或Sim2Sim导航成功。 + +## 验证入口 + +```bash +npm run typecheck +npm run check +npm run test:e2e -- web_platform/e2e/obstacle.spec.ts +# 可选:加载后端1iteration实际导出的模型(无需复制模型进仓库) +GO2_SMOKE_POLICY=/tmp/go2-obstacle-train-smoke/policy.onnx npm run test:e2e -- web_platform/e2e/obstacle.spec.ts +``` + +`src/rl/fixtures/obstacleRayGolden.json`保存CPU MuJoCo `mj_ray`参考距离;单测覆盖yaw/pitch/roll、地板命中、box内部/表面/平行/边缘、range边界、miss;真实WASM单测编译Go2碰撞/惯量模型并逐元素验证32条射线。E2E默认模型为`fixtures/obstacle/zero-action.onnx`(测试专用Constant零动作,非训练策略),检查真实浏览器WASM+ORT、地图联动、PiP、射线开关、metadata不一致拒绝。 + +完整后端字段定义见[部署契约](../training_server/OBSTACLE_AVOIDANCE.md)。 + +## Obstacle 自调参 + +自调参工作台新建Session可选Flat或Obstacle,选择Obstacle自动绑定四个专属标量/范围和五项固定客观权重(success .4/time .2/clearance .2/smooth .1/no-fall .1)。目标速度位于`params.target_velocity`,是导航command而非奖励。Approval/Automatic及运行时切换、护栏revision保持兼容,Loss看板不变。 + +从训练面板打开工作台时传递当前任务、seed、terrain/sensor(包括已验证custom_boxes快照);未同步/过时场景不发送凭据或配置。独立打开时可直接选任务,在“避障场景配置JSON”输入与训练API同构的terrain/sensor配置;服务再次严格验证。自定义地图明确为同一权威地图/起终点上的三个固定seed重复评估,不冒充三地形。Agent上下文按32/48ray实际模式描述侧后、层间与高度盲区、目标导航、擦碰/绕行/动作抖动约束。 + +Obstacle最佳策略通过ONNX内嵌metadata事务导入,导航速度随部署契约应用(.3–1.2),旧81维.6m/s模型与Flat导入保持原行为。客观评分、成功/碰撞/跌倒定义及3seed串行隔离评估详见[部署契约](../training_server/OBSTACLE_AVOIDANCE.md#obstacle-deepseek-自调参obstacle-v1)。小GPU smoke只验证链路和指标,不代表学会导航。 + +## 多层射线与性能(阶段5) + +避障高级设置的“传感器模式”可选默认水平32ray/81obs或显式三层48ray/97obs。旧服务未公布multi能力时选项禁用;加载策略以后按metadata动态分配PiP射线,不截断到32条。前47维及heading/π符号不变,Click-to-Navigate目标/面板、Loss曲线、custom_boxes、Obstacle调参速度与Flat保留。旧81 graph不能通过97部署shape校验。 + +**不是局部高程图**:尚无网格/用途定义,本阶段未实现。下倾仍可能漏掉层间矮物、遮挡及坑;标准地图底板会填内部空洞。floor真实命中进入观测,只有multi近障奖励使用已验证标准底板顶面的10微米几何分类过滤,不把floor测距伪装成miss。与floor顶面共面的其他几何无法区分。完整字段、公式和分类限制见[多层契约](../training_server/OBSTACLE_AVOIDANCE.md#显式多层射线阶段5不是局部高程图)。 + +`multiRingGolden.json`来自CPU mj_ray,覆盖全部48条完整姿态、5cm低障碍、有界地板跌落边缘和层序;真实WASM绑定逐元素比对。`multi-zero-action.onnx`为测试Constant零动作97维fixture,非训练策略。真实短训模型可用: + +```bash +GO2_MULTI_SMOKE_POLICY=/tmp/go2-multi-ring-stage5/train/policy.onnx npm run test:e2e -- web_platform/e2e/multiRing.spec.ts +``` + +`src/rl/tasks/raycastBenchmark.ts`导出`benchmarkForwardRays()`:每帧真实调用生产sampleForwardRays与slab,预计算方向/静态box,包含完整48ray全部box求交、姿态旋转和depth/线段分配;不测单ray、不用平均值代替尾延迟。默认25box和257box的16×16最坏数量布局分别warmup3000帧、采样10000帧,逐帧改变姿态/位置,不缓存命中结果。仅报告数量最坏布局,不声称穷尽所有空间排布或整浏览器帧耗时。 + +本机Intel i7-14700F/Ubuntu24.04,Node24.19:25box p50/p95约.019/.020ms,257box约.169/.187ms;HeadlessChromium151(同机,cross-origin-isolated高精度计时)约.015/.020ms与.160/.165ms,测得均低于完整48ray .2ms目标。初版通用旋转slab的257box p95≈.219ms未达标;优化为严格单位旋转快速slab后达标,保留原日志。不隔离浏览器计时分辨率约.1ms,p95量化为.2ms,不能用该读数证明严格小于目标;高精度基准不改变生产隔离/时序配置。 + +结果随硬件、GC、浏览器调度变化,不是实时保证;只测解析感知,不含渲染/ORT/物理步。复现打包脚本、Node/浏览器原始分位数与硬件说明位于`/tmp/go2-multi-ring-stage5/run_benchmark.mjs`及`perf-*.json`;日志与独立阶段incremental.diff同目录,后续review可直接读取。 + +## 复审修复:preset与导航拖动 + +Flat奖励菜单只展示服务根据来源session确认`taskId=Unitree-Go2-Flat`的preset;Obstacle及未声明身份项不展示。历史Flat的任务身份由服务恢复并完整验证schema,跨任务/损坏preset在创建作业前返回400,不推迟到训练器失败。旧服务没有返回任务身份时需升级服务才能使用这些preset;不增加Obstacle普通训练preset入口。 + +设定导航目标时超过5px即永久标记本次手势为拖动,即使返回起点也不会设置目标;画布外的同pointer移动也记录。pointercancel、Esc、失焦、退出模式、clear/dispose会清空手势。取消后旧pointerup不生效,重新按下的正常点击仍可设定。 diff --git a/web_platform/e2e/multiRing.spec.ts b/web_platform/e2e/multiRing.spec.ts new file mode 100644 index 00000000..82a4dcc6 --- /dev/null +++ b/web_platform/e2e/multiRing.spec.ts @@ -0,0 +1,153 @@ +import { expect, test } from '@playwright/test'; +import { readFileSync } from 'node:fs'; +import { resolve } from 'node:path'; +const fixture = JSON.parse( + readFileSync(resolve('web_platform/src/rl/fixtures/multiRingDeployment.json'), 'utf8'), +); + +// Original Go2 collision/inertia model, visual meshes omitted to keep the smoke fixture lightweight. +const go2 = readFileSync( + resolve('training_server/rl/src/assets/robots/unitree_go2/xmls/go2.xml'), + 'utf8', +) + .replace(/]*\/>/g, '') + .replace(/]*\bmesh="[^"]*"[^>]*\/>/g, ''); +let model = readFileSync( + process.env.GO2_MULTI_SMOKE_POLICY ?? + resolve('web_platform/fixtures/obstacle/multi-zero-action.onnx'), +); + +test('训练作业一键导入:真实WASM地图+97维ORT+PiP+射线开关', async ({ page }) => { + await page.addInitScript(() => { + localStorage.setItem('mujoco-local-training-job-id', 'a'.repeat(32)); + sessionStorage.setItem('mujoco-local-training-token', 'test'); + }); + const job = { + id: 'a'.repeat(32), + taskId: 'Unitree-Go2-ObstacleAvoidance', + state: 'succeeded', + artifactReady: true, + progress: 1, + iteration: 1, + maxIterations: 1, + logs: [ + 'Learning iteration 0 / 1', + 'Mean value loss: 0.9', + 'Mean surrogate loss: -0.1', + 'Mean entropy loss: -1', + 'Mean reward: 2', + 'Mean episode length: 20', + 'Learning iteration 1 / 1', + 'Mean value loss: 0.5', + 'Mean surrogate loss: -0.2', + 'Mean entropy loss: -0.8', + 'Mean reward: 3', + 'Mean episode length: 30', + ], + message: '测试专用零动作策略', + deployment: fixture, + }; + await page.route('http://127.0.0.1:8765/**', (route) => { + const url = route.request().url(); + if (url.endsWith('/policy.onnx')) + return route.fulfill({ contentType: 'application/octet-stream', body: model }); + if (url.endsWith('/health')) + return route.fulfill({ + json: { ready: true, trainerRoot: '/test', tasks: [job.taskId], activeJobId: job.id }, + }); + if (url.includes('/presets')) return route.fulfill({ json: { presets: [] } }); + return route.fulfill({ json: job }); + }); + await page.goto('/'); + await page + .locator('input[type="file"]') + .first() + .setInputFiles({ name: 'go2.xml', mimeType: 'text/xml', buffer: Buffer.from(go2) }); + await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 }); + await page.getByRole('tab', { name: '控制台' }).click(); + const tools = page.getByRole('tabpanel', { name: '控制台' }); + await tools.getByRole('button', { name: /强化学习任务/ }).click(); + await tools.getByRole('button', { name: '连接', exact: true }).click(); + await expect(tools.getByRole('button', { name: '导入策略' })).toBeEnabled(); + await tools.getByRole('button', { name: '导入策略' }).click(); + await expect(page.getByText('训练配套物理地图', { exact: false })).toBeVisible({ + timeout: 30_000, + }); + await expect(page.getByLabel('摄像头画面', { exact: true })).toBeVisible(); + const rays = page.getByRole('checkbox', { name: '显示避障射线' }); + await expect(rays).toBeChecked(); + await rays.uncheck(); + await rays.check(); + const section = tools.getByRole('button', { name: /ONNX 策略运行/ }); + if ((await section.getAttribute('aria-expanded')) === 'false') await section.click(); + await expect(tools.getByText('97 / 12')).toBeVisible(); + await expect(tools.getByText('Go2 前视射线避障导航', { exact: true })).toBeVisible(); + const count = tools.getByText('推理次数', { exact: true }).locator('..'); + await expect + .poll(async () => Number((await count.textContent())?.replace(/\D/g, ''))) + .toBeGreaterThan(0); + const pause = page.getByRole('button', { name: '⏸ 暂停' }); + if (await pause.isVisible()) await pause.click(); + const stop = tools.getByRole('button', { name: '停止', exact: true }); + if (await stop.isVisible()) await stop.click(); + const targetRow = tools.getByText('当前目标 (X, Y)', { exact: true }).locator('..'); + const initialTarget = await targetRow.textContent(); + const selection = () => page.locator('[aria-selected="true"]').allTextContents(); + const beforeSelection = await selection(); + await tools.getByRole('button', { name: '设定目标', exact: true }).click(); + await page.keyboard.press('Escape'); + await expect(tools.getByRole('button', { name: '设定目标', exact: true })).toHaveAttribute( + 'aria-pressed', + 'false', + ); + await expect(targetRow).toHaveText(initialTarget!); + await tools.getByRole('button', { name: '设定目标', exact: true }).click(); + const canvas = page.locator('.viewport-shell canvas').first(); + const bounds = (await canvas.boundingBox())!; + await canvas.click({ position: { x: bounds.width * 0.5, y: bounds.height * 0.55 } }); + await expect(targetRow).not.toHaveText(initialTarget!); + await expect(tools.getByRole('button', { name: '设定目标', exact: true })).toHaveAttribute( + 'aria-pressed', + 'false', + ); + expect(await selection()).toEqual(beforeSelection); + await expect(tools.getByRole('button', { name: '启用', exact: true })).toBeVisible(); + await tools.getByRole('button', { name: '复位目标点' }).click(); + await expect(targetRow).toHaveText(initialTarget!); + await tools.getByRole('button', { name: '导入策略' }).click(); + await expect(tools.getByRole('alert')).toHaveCount(0); + // Switch real valid graph/metadata pairs in both directions, without relaxing transactions. + const multiModel = model; + job.deployment = JSON.parse( + readFileSync(resolve('web_platform/src/rl/fixtures/obstacleDeployment.json'), 'utf8'), + ); + model = readFileSync(resolve('web_platform/fixtures/obstacle/zero-action.onnx')); + await tools.getByRole('button', { name: '连接', exact: true }).click(); + await tools.getByRole('button', { name: '导入策略' }).click(); + await expect(tools.getByText('81 / 12')).toBeVisible(); + job.deployment = fixture; + model = multiModel; + await tools.getByRole('button', { name: '连接', exact: true }).click(); + await tools.getByRole('button', { name: '导入策略' }).click(); + await expect(tools.getByText('97 / 12')).toBeVisible(); + const pauseAgain = page.getByRole('button', { name: '⏸ 暂停' }); + if (await pauseAgain.isVisible()) await pauseAgain.click(); + const stopAgain = tools.getByRole('button', { name: '停止', exact: true }); + if (await stopAgain.isVisible()) await stopAgain.click(); + // A real 81-feature ONNX graph with valid 97 metadata must fail transactionally. + model = readFileSync(resolve('web_platform/fixtures/obstacle/multi-wrong-graph.onnx')); + await tools.getByRole('button', { name: '导入策略' }).click(); + await expect(tools.getByRole('alert')).toContainText('维度'); + await expect(tools.getByText('97 / 12')).toBeVisible(); + await expect(targetRow).toHaveText(initialTarget!); + await tools.getByRole('button', { name: '设定目标', exact: true }).click(); + await tools.getByRole('button', { name: '卸载', exact: true }).click(); + await expect(targetRow).toHaveCount(0); + await expect(tools.locator('.uplot')).toHaveCount(0); + await tools.getByRole('button', { name: /训练指标趋势/ }).click(); + await expect(tools.locator('.uplot canvas')).toHaveCount(1); + await tools.getByRole('tab', { name: '综合', exact: true }).click(); + await expect(tools.locator('.uplot canvas')).toHaveCount(5); + await tools.getByRole('button', { name: /训练指标趋势/ }).click(); + await expect(tools.locator('.uplot')).toHaveCount(0); +}); diff --git a/web_platform/e2e/obstacle.spec.ts b/web_platform/e2e/obstacle.spec.ts new file mode 100644 index 00000000..ad68dafe --- /dev/null +++ b/web_platform/e2e/obstacle.spec.ts @@ -0,0 +1,291 @@ +import { expect, test } from '@playwright/test'; +import { readFileSync } from 'node:fs'; +import { resolve } from 'node:path'; +const fixture = JSON.parse( + readFileSync(resolve('web_platform/src/rl/fixtures/obstacleDeployment.json'), 'utf8'), +); + +// Original Go2 collision/inertia model, visual meshes omitted to keep the smoke fixture lightweight. +const go2 = readFileSync( + resolve('training_server/rl/src/assets/robots/unitree_go2/xmls/go2.xml'), + 'utf8', +) + .replace(/]*\/>/g, '') + .replace(/]*\bmesh="[^"]*"[^>]*\/>/g, ''); +const model = readFileSync( + process.env.GO2_SMOKE_POLICY ?? resolve('web_platform/fixtures/obstacle/zero-action.onnx'), +); + +test('训练作业一键导入:真实WASM地图+81维ORT+PiP+射线开关', async ({ page }) => { + await page.addInitScript(() => { + localStorage.setItem('mujoco-local-training-job-id', 'a'.repeat(32)); + sessionStorage.setItem('mujoco-local-training-token', 'test'); + }); + const job = { + id: 'a'.repeat(32), + taskId: 'Unitree-Go2-ObstacleAvoidance', + state: 'succeeded', + artifactReady: true, + progress: 1, + iteration: 1, + maxIterations: 1, + logs: [ + 'Learning iteration 0 / 1', + 'Mean value loss: 0.9', + 'Mean surrogate loss: -0.1', + 'Mean entropy loss: -1', + 'Mean reward: 2', + 'Mean episode length: 20', + 'Learning iteration 1 / 1', + 'Mean value loss: 0.5', + 'Mean surrogate loss: -0.2', + 'Mean entropy loss: -0.8', + 'Mean reward: 3', + 'Mean episode length: 30', + ], + message: '测试专用零动作策略', + deployment: fixture, + }; + await page.route('http://127.0.0.1:8765/**', (route) => { + const url = route.request().url(); + if (url.endsWith('/policy.onnx')) + return route.fulfill({ contentType: 'application/octet-stream', body: model }); + if (url.endsWith('/health')) + return route.fulfill({ + json: { ready: true, trainerRoot: '/test', tasks: [job.taskId], activeJobId: job.id }, + }); + if (url.includes('/presets')) return route.fulfill({ json: { presets: [] } }); + return route.fulfill({ json: job }); + }); + await page.goto('/'); + await page + .locator('input[type="file"]') + .first() + .setInputFiles({ name: 'go2.xml', mimeType: 'text/xml', buffer: Buffer.from(go2) }); + await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 }); + await page.getByRole('tab', { name: '控制台' }).click(); + const tools = page.getByRole('tabpanel', { name: '控制台' }); + await tools.getByRole('button', { name: /强化学习任务/ }).click(); + await tools.getByRole('button', { name: '连接', exact: true }).click(); + await expect(tools.getByRole('button', { name: '导入策略' })).toBeEnabled(); + await tools.getByRole('button', { name: '导入策略' }).click(); + await expect(page.getByText('训练配套物理地图', { exact: false })).toBeVisible({ + timeout: 30_000, + }); + await expect(page.getByLabel('摄像头画面', { exact: true })).toBeVisible(); + const rays = page.getByRole('checkbox', { name: '显示避障射线' }); + await expect(rays).toBeChecked(); + await rays.uncheck(); + await rays.check(); + const section = tools.getByRole('button', { name: /ONNX 策略运行/ }); + if ((await section.getAttribute('aria-expanded')) === 'false') await section.click(); + await expect(tools.getByText('81 / 12')).toBeVisible(); + await expect(tools.getByText('Go2 前视射线避障导航', { exact: true })).toBeVisible(); + const count = tools.getByText('推理次数', { exact: true }).locator('..'); + await expect + .poll(async () => Number((await count.textContent())?.replace(/\D/g, ''))) + .toBeGreaterThan(0); + const pause = page.getByRole('button', { name: '⏸ 暂停' }); + if (await pause.isVisible()) await pause.click(); + const stop = tools.getByRole('button', { name: '停止', exact: true }); + if (await stop.isVisible()) await stop.click(); + const targetRow = tools.getByText('当前目标 (X, Y)', { exact: true }).locator('..'); + const initialTarget = await targetRow.textContent(); + const selection = () => page.locator('[aria-selected="true"]').allTextContents(); + const beforeSelection = await selection(); + await tools.getByRole('button', { name: '设定目标', exact: true }).click(); + await page.keyboard.press('Escape'); + await expect(tools.getByRole('button', { name: '设定目标', exact: true })).toHaveAttribute( + 'aria-pressed', + 'false', + ); + await expect(targetRow).toHaveText(initialTarget!); + await tools.getByRole('button', { name: '设定目标', exact: true }).click(); + const canvas = page.locator('.viewport-shell canvas').first(); + const bounds = (await canvas.boundingBox())!; + await canvas.click({ position: { x: bounds.width * 0.5, y: bounds.height * 0.55 } }); + await expect(targetRow).not.toHaveText(initialTarget!); + await expect(tools.getByRole('button', { name: '设定目标', exact: true })).toHaveAttribute( + 'aria-pressed', + 'false', + ); + expect(await selection()).toEqual(beforeSelection); + await expect(tools.getByRole('button', { name: '启用', exact: true })).toBeVisible(); + await tools.getByRole('button', { name: '复位目标点' }).click(); + await expect(targetRow).toHaveText(initialTarget!); + await tools.getByRole('button', { name: '设定目标', exact: true }).click(); + await tools.getByRole('button', { name: '卸载', exact: true }).click(); + await expect(targetRow).toHaveCount(0); + await expect(tools.locator('.uplot')).toHaveCount(0); + await tools.getByRole('button', { name: /训练指标趋势/ }).click(); + await expect(tools.locator('.uplot canvas')).toHaveCount(1); + await tools.getByRole('tab', { name: '综合', exact: true }).click(); + await expect(tools.locator('.uplot canvas')).toHaveCount(5); + await tools.getByRole('button', { name: /训练指标趋势/ }).click(); + await expect(tools.locator('.uplot')).toHaveCount(0); +}); + +test('下载metadata与作业不匹配时拒绝,未更换地图或启用策略', async ({ page }) => { + await page.addInitScript(() => sessionStorage.setItem('mujoco-local-training-token', 'test')); + const job = { + id: 'b'.repeat(32), + taskId: fixture.taskId, + state: 'succeeded', + artifactReady: true, + progress: 1, + iteration: 1, + maxIterations: 1, + logs: [], + message: '', + deployment: { ...fixture, sensorCfg: { ...fixture.sensorCfg, fov: 60 } }, + }; + await page.route('http://127.0.0.1:8765/**', (route) => { + const url = route.request().url(); + if (url.endsWith('/policy.onnx')) + return route.fulfill({ contentType: 'application/octet-stream', body: model }); + if (url.endsWith('/health')) + return route.fulfill({ json: { ready: true, tasks: [job.taskId], activeJobId: job.id } }); + if (url.includes('/presets')) return route.fulfill({ json: { presets: [] } }); + return route.fulfill({ json: job }); + }); + await page.goto('/'); + await page + .locator('input[type="file"]') + .first() + .setInputFiles({ name: 'go2.xml', mimeType: 'text/xml', buffer: Buffer.from(go2) }); + await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 }); + await page.getByRole('tab', { name: '控制台' }).click(); + const tools = page.getByRole('tabpanel', { name: '控制台' }); + await tools.getByRole('button', { name: /强化学习任务/ }).click(); + await tools.getByRole('button', { name: '连接', exact: true }).click(); + await tools.getByRole('button', { name: '导入策略' }).click(); + await expect(page.getByText('下载的策略与训练作业部署配置不一致').first()).toBeVisible(); + await expect(page.getByText('训练配套物理地图', { exact: false })).toHaveCount(0); + await expect(page.getByLabel('摄像头画面', { exact: true })).toHaveCount(0); +}); + +for (const failure of ['wrong-graph', 'ort-init-failure']) { + test(`事务导入${failure}时保留旧地图、策略和暂停状态`, async ({ page }) => { + await page.addInitScript(() => sessionStorage.setItem('mujoco-local-training-token', 'test')); + let bytes = model; + const job = { + id: 'c'.repeat(32), + taskId: fixture.taskId, + state: 'succeeded', + artifactReady: true, + progress: 1, + iteration: 1, + maxIterations: 1, + logs: [], + message: '', + deployment: fixture, + }; + await page.route('http://127.0.0.1:8765/**', (route) => { + const url = route.request().url(); + if (url.endsWith('/policy.onnx')) + return route.fulfill({ contentType: 'application/octet-stream', body: bytes }); + if (url.endsWith('/health')) + return route.fulfill({ json: { ready: true, tasks: [job.taskId], activeJobId: job.id } }); + if (url.includes('/presets')) return route.fulfill({ json: { presets: [] } }); + return route.fulfill({ json: job }); + }); + await page.goto('/'); + await page + .locator('input[type="file"]') + .first() + .setInputFiles({ name: 'go2.xml', mimeType: 'text/xml', buffer: Buffer.from(go2) }); + await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 }); + await page.getByRole('tab', { name: '控制台' }).click(); + const tools = page.getByRole('tabpanel', { name: '控制台' }); + await tools.getByRole('button', { name: /强化学习任务/ }).click(); + await tools.getByRole('button', { name: '连接', exact: true }).click(); + await tools.getByRole('button', { name: '导入策略' }).click(); + await expect(page.getByLabel('摄像头画面', { exact: true })).toBeVisible(); + const pause = page.getByRole('button', { name: '⏸ 暂停' }); + if (await pause.isVisible()) await pause.click(); + const section = tools.getByRole('button', { name: /ONNX 策略运行/ }); + if ((await section.getAttribute('aria-expanded')) === 'false') await section.click(); + await expect(tools.getByText('81 / 12')).toBeVisible(); + const inference = tools.getByText('推理次数', { exact: true }).locator('..'); + // Stop policy as well to invalidate any pending inference, leaving a stable baseline. + const stop = tools.getByRole('button', { name: '停止', exact: true }); + if (await stop.isVisible()) await stop.click(); + const before = await inference.textContent(); + bytes = readFileSync(resolve(`web_platform/fixtures/obstacle/${failure}.onnx`)); + await tools.getByRole('button', { name: '导入策略' }).click(); + await expect(tools.getByRole('alert')).toContainText( + failure === 'wrong-graph' ? '维度' : 'MissingOperatorForTransactionTest', + ); + await expect(tools.getByText('81 / 12')).toBeVisible(); + await expect(inference).toHaveText(before!); + await expect(page.getByLabel('摄像头画面', { exact: true })).toBeVisible(); + await expect(page.getByText('训练配套物理地图', { exact: false })).toBeVisible(); + await expect(page.getByRole('button', { name: '▶ 播放' })).toBeVisible(); + }); +} + +test('新server默认Flat作业兼容旧无metadata47维ONNX,错误graph拒绝且保留策略', async ({ page }) => { + await page.addInitScript(() => sessionStorage.setItem('mujoco-local-training-token', 'test')); + const flat = { + ...fixture, + taskId: 'Unitree-Go2-Flat', + observationSize: 47, + observationTerms: fixture.observationTerms.slice(0, 7), + terrain: undefined, + terrainPreset: undefined, + terrainParams: undefined, + sensorCfg: undefined, + navigation: undefined, + }; + let bytes = readFileSync(resolve('web_platform/fixtures/obstacle/legacy-flat.onnx')); + const job = { + id: 'd'.repeat(32), + taskId: flat.taskId, + state: 'succeeded', + artifactReady: true, + progress: 1, + iteration: 1, + maxIterations: 1, + logs: [], + message: '', + deployment: flat, + }; + await page.route('http://127.0.0.1:8765/**', (route) => { + const url = route.request().url(); + if (url.endsWith('/policy.onnx')) + return route.fulfill({ contentType: 'application/octet-stream', body: bytes }); + if (url.endsWith('/health')) + return route.fulfill({ json: { ready: true, tasks: [job.taskId], activeJobId: job.id } }); + if (url.includes('/presets')) return route.fulfill({ json: { presets: [] } }); + return route.fulfill({ json: job }); + }); + const actuators = + '' + + fixture.jointNames + .map((name: string) => ``) + .join('') + + ''; + await page.goto('/'); + await page + .locator('input[type="file"]') + .first() + .setInputFiles({ + name: 'go2.xml', + mimeType: 'text/xml', + buffer: Buffer.from(go2.replace('', actuators + '')), + }); + await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 }); + await page.getByRole('tab', { name: '控制台' }).click(); + const tools = page.getByRole('tabpanel', { name: '控制台' }); + await tools.getByRole('button', { name: /强化学习任务/ }).click(); + await tools.getByRole('button', { name: '连接', exact: true }).click(); + await tools.getByRole('button', { name: '导入策略' }).click(); + const section = tools.getByRole('button', { name: /ONNX 策略运行/ }); + if ((await section.getAttribute('aria-expanded')) === 'false') await section.click(); + await expect(tools.getByText('47 / 12')).toBeVisible(); + bytes = readFileSync(resolve('web_platform/fixtures/obstacle/legacy-wrong-shape.onnx')); + await tools.getByRole('button', { name: '导入策略' }).click(); + await expect(tools.getByRole('alert')).toContainText('维度'); + await expect(tools.getByText('47 / 12')).toBeVisible(); + await expect(page.getByText('训练配套物理地图', { exact: false })).toHaveCount(0); +}); diff --git a/web_platform/e2e/pretrainedNavigation.spec.ts b/web_platform/e2e/pretrainedNavigation.spec.ts new file mode 100644 index 00000000..a2081995 --- /dev/null +++ b/web_platform/e2e/pretrainedNavigation.spec.ts @@ -0,0 +1,209 @@ +import { expect, test } from '@playwright/test'; +import { readFileSync, writeFileSync } from 'node:fs'; +import { resolve } from 'node:path'; + +// Opt-in real source-derived actor, not a constant fixture. Run against the Vite dev server. +const policyPath = process.env.GO2_PRETRAINED_NAV_POLICY; +test('真实warm-start策略连续行走超过20秒并换目标,不重载策略或自动启用', async ({ + page, +}, testInfo) => { + test.skip( + !policyPath, + '设置GO2_PRETRAINED_NAV_POLICY为真实47→81/97 warm-start导出,并运行Vite dev :4173', + ); + const xml = readFileSync( + resolve('training_server/rl/src/assets/robots/unitree_go2/xmls/go2.xml'), + 'utf8', + ) + .replace(/]*\/>/g, '') + .replace(/]*\bmesh="[^"]*"[^>]*\/>/g, ''); + const policy = Array.from(readFileSync(policyPath!)); + await page.goto(process.env.GO2_PRETRAINED_DEV_URL ?? 'http://127.0.0.1:4174'); + const result = await page.evaluate( + async ({ xml, policy }) => { + const adapterPath = '/src/simulation/PhysicsAdapter.ts'; + const deploymentPath = '/src/rl/deployment.ts'; + const navigationPath = '/src/rl/tasks/go2ObstacleAvoidance.ts'; + const { MainThreadPhysicsAdapter } = (await import( + adapterPath + )) as typeof import('../src/simulation/PhysicsAdapter'); + const { readPolicyDeployment } = (await import( + deploymentPath + )) as typeof import('../src/rl/deployment'); + const { obstacleNavigation } = (await import( + navigationPath + )) as typeof import('../src/rl/tasks/go2ObstacleAvoidance'); + const bytes = new Uint8Array(policy), + data = new TextEncoder().encode(xml); + const deployment = readPolicyDeployment(bytes)!; + const adapter = new MainThreadPhysicsAdapter(); + try { + await adapter.load( + { + id: 'real-pretrained', + name: 'go2', + files: [ + { path: 'go2.xml', data, size: data.length, source: 'file', mimeType: 'text/xml' }, + ], + entries: [{ path: 'go2.xml', format: 'mjcf', label: 'Go2' }], + maps: [], + totalBytes: data.length, + }, + 'go2.xml', + { + trainingDeployment: deployment, + trainingPolicy: { data: bytes, path: 'user-warmstart.onnx' }, + }, + ); + const session = adapter.session!; + const originalPolicy = (session as unknown as { rlPolicy: object }).rlPolicy; + const initial = adapter.snapshot()!; + const samples: { + time: number; + x: number; + y: number; + z: number; + enabled?: boolean; + error?: string; + }[] = []; + const run = async (seconds: number) => { + const target = Number(session.data.time) + seconds; + while (Number(session.data.time) < target) { + for (let i = 0; i < 10; i++) adapter.singleStep(); + await new Promise((resolve) => setTimeout(resolve, 0)); + if (!adapter.snapshot()?.rlPolicy?.enabled) break; + } + const s = adapter.snapshot()!; + samples.push({ + time: s.time, + x: s.qpos[0], + y: s.qpos[1], + z: s.qpos[2], + enabled: s.rlPolicy?.enabled, + error: s.rlPolicy?.error, + }); + return s; + }; + const before = await run(21); + const navBefore = obstacleNavigation( + before.qpos.slice(0, 3), + before.qpos.slice(3, 7), + before.rlPolicy!.navigation!.target, + deployment.terrain!.size, + deployment.navigation?.speed, + ); + // Nearby diagonal target exercises both heading and velocity command; no pose teleport. + adapter.setNavigationTarget([before.qpos[0] - 2, before.qpos[1] + 2]); + const changed = adapter.snapshot()!; + const navChanged = obstacleNavigation( + changed.qpos.slice(0, 3), + changed.qpos.slice(3, 7), + changed.rlPolicy!.navigation!.target, + deployment.terrain!.size, + deployment.navigation?.speed, + ); + const after = await run(4); + adapter.setRLPolicyEnabled(false); + adapter.setNavigationTarget([0, 0]); + const disabled = adapter.snapshot()!; + return { + initial: { time: initial.time, qpos: initial.qpos }, + before: { + time: before.time, + qpos: before.qpos, + ctrl: before.ctrl, + status: before.rlPolicy, + }, + after: { time: after.time, qpos: after.qpos, ctrl: after.ctrl, status: after.rlPolicy }, + samples, + navBefore, + navChanged, + sameSession: adapter.session === session, + samePolicy: (session as unknown as { rlPolicy: object }).rlPolicy === originalPolicy, + heldControlUnchangedOnTarget: changed.ctrl.every((v, i) => v === before.ctrl[i]), + countBeforeTarget: before.rlPolicy!.inferenceCount, + countAfterTarget: changed.rlPolicy!.inferenceCount, + disabledAfterNewTarget: !disabled.rlPolicy!.enabled, + }; + } finally { + adapter.dispose(); + } + }, + { xml, policy }, + ); + const evidence = testInfo.outputPath('real-navigation-measurements.json'); + writeFileSync(evidence, JSON.stringify(result, null, 2)); + await testInfo.attach('real-navigation-measurements.json', { + path: evidence, + contentType: 'application/json', + }); + expect(result.before.time).toBeGreaterThan(20); + expect(result.before.status?.enabled).toBe(true); + expect(result.after.status?.enabled).toBe(true); + expect( + Math.hypot( + result.before.qpos[0] - result.initial.qpos[0], + result.before.qpos[1] - result.initial.qpos[1], + ), + ).toBeGreaterThan(0.5); + expect(result.navChanged.command).not.toEqual(result.navBefore.command); + expect(result.navChanged.targetError).not.toEqual(result.navBefore.targetError); + expect(result.after.ctrl).not.toEqual(result.before.ctrl); + expect(result.after.status!.navigation!.distance).toBeLessThan(2 * Math.SQRT2); + expect(result.countAfterTarget).toBe(result.countBeforeTarget); + expect(result.after.status!.inferenceCount).toBeGreaterThan(result.countBeforeTarget); + expect(result.sameSession).toBe(true); + expect(result.samePolicy).toBe(true); + expect(result.heldControlUnchangedOnTarget).toBe(true); + expect(result.disabledAfterNewTarget).toBe(true); +}); + +test('用户原始47维ONNX通过普通Flat面板加载,不新增点击导航模式', async ({ page }) => { + const source = process.env.GO2_PRETRAINED_FLAT_POLICY; + const robotXml = process.env.GO2_PRETRAINED_FLAT_XML; + test.skip(!source || !robotXml, '设置原始47维ONNX与同源带actuator的机器人XML路径'); + const xml = readFileSync(robotXml!, 'utf8') + .replace(/]*\/>/g, '') + .replace(/]*\bmesh="[^"]*"[^>]*\/>/g, ''); + await page.goto(process.env.GO2_PRETRAINED_DEV_URL ?? 'http://127.0.0.1:4174'); + await page + .locator('input[type="file"]') + .first() + .setInputFiles({ name: 'go2.xml', mimeType: 'text/xml', buffer: Buffer.from(xml) }); + await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 }); + await page.getByRole('tab', { name: '控制台' }).click(); + const tools = page.getByRole('tabpanel', { name: '控制台' }); + const section = tools.getByRole('button', { name: /ONNX 策略运行/ }); + if ((await section.getAttribute('aria-expanded')) === 'false') await section.click(); + const chooser = page.waitForEvent('filechooser'); + await tools.getByRole('button', { name: '导入 ONNX' }).click(); + await (await chooser).setFiles(source!); + await expect + .poll( + async () => + (await tools.getByText('47 / 12').count()) + + (await page.getByRole('button', { name: '技术详情', exact: true }).count()), + ) + .toBeGreaterThan(0); + const details = page.getByRole('button', { name: '技术详情', exact: true }); + if (await details.isVisible()) { + await details.click(); + throw new Error((await page.getByRole('alert').allTextContents()).join('\n')); + } + await expect(tools.getByText('47 / 12')).toBeVisible({ timeout: 30_000 }); + const enable = tools.getByRole('button', { name: '启用', exact: true }); + if (await enable.isVisible()) await enable.click(); + const play = page.getByRole('button', { name: '▶ 播放', exact: true }); + if (await play.isVisible()) await play.click(); + await expect + .poll(async () => + Number( + (await tools.getByText('推理次数', { exact: true }).locator('..').textContent())?.replace( + /\D/g, + '', + ), + ), + ) + .toBeGreaterThan(2); + await expect(tools.getByRole('button', { name: '设定目标', exact: true })).toHaveCount(0); +}); diff --git a/web_platform/e2e/pretrainedUpload.spec.ts b/web_platform/e2e/pretrainedUpload.spec.ts new file mode 100644 index 00000000..d537baa4 --- /dev/null +++ b/web_platform/e2e/pretrainedUpload.spec.ts @@ -0,0 +1,177 @@ +import { expect, test } from '@playwright/test'; +import { spawn } from 'node:child_process'; +import { createHash } from 'node:crypto'; +import { mkdtempSync, readFileSync, writeFileSync } from 'node:fs'; +import { createServer } from 'node:net'; +import { tmpdir } from 'node:os'; +import { join, resolve } from 'node:path'; + +const realRoot = process.env.GO2_UPLOAD_REAL_DIR; +const token = 'local-upload-browser-test-token'; +async function unusedPort() { + const socket = createServer(); + await new Promise((done) => socket.listen(0, '127.0.0.1', done)); + const address = socket.address(); + const port = typeof address === 'object' && address ? address.port : 0; + await new Promise((done) => socket.close(() => done())); + return port; +} +for (const panel of ['ordinary', 'tuning'] as const) { + for (const format of ['pt', 'onnx'] as const) { + test(`${panel}真实单${format}文件:默认服务空目录上传并选择内容ID`, async ({ + page, + }, testInfo) => { + test.skip(!realRoot, '设置GO2_UPLOAD_REAL_DIR;Vite dev提供真实组件,启动独立默认训练服务'); + const root = mkdtempSync(join(tmpdir(), `go2-browser-upload-${panel}-${format}-`)); + const port = await unusedPort(); + const endpoint = `http://127.0.0.1:${port}`; + const backend = spawn( + resolve('.venv/bin/python'), + [ + '-u', + 'training_server/server.py', + '--port', + String(port), + '--token', + token, + '--tuning-data-root', + root, + ], + { + cwd: process.cwd(), + env: { ...process.env, DEEPSEEK_API_KEY: '', WANDB_MODE: 'disabled' }, + }, + ); + let log = ''; + backend.stdout.on('data', (data: Buffer) => { + log += data.toString(); + }); + backend.stderr.on('data', (data: Buffer) => { + log += data.toString(); + }); + try { + await expect + .poll( + async () => { + try { + return ( + await fetch(`${endpoint}/api/training/health`, { + headers: { Authorization: `Bearer ${token}` }, + }) + ).status; + } catch { + return 0; + } + }, + { timeout: 30_000 }, + ) + .toBe(200); + await page.goto(process.env.GO2_PRETRAINED_DEV_URL ?? 'http://127.0.0.1:4174'); + // Real production panel components, isolated from unrelated heavy scene import. + await page.evaluate(async (panel) => { + localStorage.clear(); + sessionStorage.clear(); + const entry = await (await fetch('/src/main.tsx')).text(); + const reactPath = entry.match(/from "([^"]+\/react\.js[^"]*)"/)![1]; + const domPath = entry.match(/from "([^"]+\/react-dom_client\.js[^"]*)"/)![1]; + const { createElement } = (await import(reactPath)).default; + const { createRoot } = (await import(domPath)).default; + const host = document.createElement('div'); + host.id = 'upload-browser-panel'; + host.style.cssText = + 'position:fixed;inset:0;z-index:10000;background:#111;overflow:auto;padding:20px'; + document.body.append(host); + if (panel === 'ordinary') { + const path = '/src/training/LocalTrainingPanel.tsx'; + const { LocalTrainingPanel } = await import(path); + createRoot(host).render(createElement(LocalTrainingPanel, { onPolicyReady: () => {} })); + } else { + const path = '/src/tuning/TuningApp.tsx'; + const { TuningApp } = await import(path); + createRoot(host).render(createElement(TuningApp)); + } + }, panel); + const ui = page.locator('#upload-browser-panel'); + await ui + .getByLabel(panel === 'ordinary' ? '本地训练服务地址' : '训练服务地址', { exact: true }) + .fill(endpoint); + await ui + .getByLabel(panel === 'ordinary' ? '训练服务访问令牌' : '访问令牌(仅当前标签页)', { + exact: true, + }) + .fill(token); + await ui + .getByRole('button', { name: panel === 'ordinary' ? /^连接$/ : '连接/刷新' }) + .click(); + await expect(ui.getByLabel('确认Go2 legacy47模板')).toBeEnabled(); + await expect(ui.getByLabel('基础策略', { exact: true })).toHaveValue(''); + await ui.getByLabel('确认Go2 legacy47模板').check(); + const path = join(realRoot!, format === 'pt' ? 'model_10000.pt' : 'policy.onnx'); + const sha = createHash('sha256').update(readFileSync(path)).digest('hex'); + const responsePromise = page.waitForResponse((response) => + response.url().startsWith(`${endpoint}/api/training/pretrained-sources/upload?`), + ); + await ui.getByLabel('选择基础策略文件').setInputFiles(path); + const response = await responsePromise; + expect(response.status()).toBe(201); + const catalog = await ( + await fetch(`${endpoint}/api/training/health`, { + headers: { Authorization: `Bearer ${token}` }, + }) + ).json(); + const record = catalog.pretrainedSources[0]; + await expect(ui.getByLabel('基础策略', { exact: true })).toHaveValue(record.id); + await expect(ui.getByText(/原文件 SHA256/)).toContainText(sha); + await expect(ui.getByText(/上传格式/)).toContainText(format); + await expect(ui.getByText(/已选择.*仅继承策略权重/)).toBeVisible(); + if (format === 'onnx') + await expect(ui.getByText(/ONNX统计count合成/)).toContainText('1000000'); + // Intercept only train dispatch; uploads and catalog are real authenticated HTTP. + // Never dispatch 4096-env training or an Agent/three-seed evaluation from this test. + let submitted: Record | undefined; + await page.route( + `${endpoint}${panel === 'ordinary' ? '/api/training/jobs' : '/api/tuning/sessions'}`, + async (route) => { + if (route.request().method() !== 'POST') return route.continue(); + submitted = route.request().postDataJSON() as Record; + await route.fulfill({ + status: 400, + contentType: 'application/json', + body: JSON.stringify({ error: '浏览器验收只截获启动请求,未训练或调用Agent' }), + }); + }, + ); + if (panel === 'tuning') await ui.getByLabel('Agent 失败时允许 Optuna fallback').check(); + await ui + .getByRole('button', { name: panel === 'ordinary' ? '发起本地训练' : '启动自调参' }) + .click(); + await expect.poll(() => submitted?.pretrainedSourceId).toBe(record.id); + if (panel === 'tuning') expect(submitted?.mode).toBe('approval'); + await expect(ui.getByLabel('基础策略', { exact: true })).toHaveValue(record.id); + const health = await ( + await fetch(`${endpoint}/api/training/health`, { + headers: { Authorization: `Bearer ${token}` }, + }) + ).json(); + expect(health.pretrainedSources).toHaveLength(1); + expect(createHash('sha256').update(readFileSync(path)).digest('hex')).toBe(sha); + await testInfo.attach('upload-contract', { + body: JSON.stringify( + { root, endpoint, record, submitted, sourceUnchanged: true }, + null, + 2, + ), + contentType: 'application/json', + }); + await page.screenshot({ path: testInfo.outputPath('upload.png') }); + } finally { + backend.kill('SIGTERM'); + await new Promise((done) => { + if (backend.exitCode !== null) done(); + else backend.once('exit', () => done()); + }); + writeFileSync(join(root, 'browser-server.log'), log); + } + }); + } +} diff --git a/web_platform/fixtures/obstacle/README.md b/web_platform/fixtures/obstacle/README.md new file mode 100644 index 00000000..e55fa968 --- /dev/null +++ b/web_platform/fixtures/obstacle/README.md @@ -0,0 +1,9 @@ +# 测试专用避障契约模型 + +`zero-action.onnx`:ONNX opset17/IR9,输入float32[1,81],Constant节点输出float32[1,12]全0,携带seed42避障deployment。仅验浏览器地图/感知/推理闭环,**不是已训练或可导航策略**。由onnx.helper.make_model/make_graph/make_node(Constant)生成,metadata来自src/rl/fixtures/obstacleDeployment.json。 + +事务回归fixture:`wrong-graph.onnx`包含合法避障metadata但实际输入47维;`ort-init-failure.onnx`包含合法81维metadata但不存在的算子(故意无法初始化ORT);`legacy-flat.onnx`为无metadata的47→12零动作旧导出;`legacy-wrong-shape.onnx`为无metadata的81→12零动作。均由相同onnx.helper入口生成,仅供成功/失败分支测试。 + +`multi-zero-action.onnx`同样为Constant零动作,仅将实际输入shape设为[1,97]并携带`multiRingDeployment.json`。不是训练策略,也没有将81网络冒充97。真实短训导出只保留在/tmp。 + +`multi-wrong-graph.onnx`携带合法97维metadata但真实输入81维,用于浏览器ORT明确shape拒绝及旧会话保留回归。 diff --git a/web_platform/fixtures/obstacle/legacy-flat.onnx b/web_platform/fixtures/obstacle/legacy-flat.onnx new file mode 100644 index 0000000000000000000000000000000000000000..92a5b578526a4b3123bbe92bd3a4aba1e2ba99dd GIT binary patch literal 177 zcmdZ?&wg@!1%JP~|Mkqt)9HHl>HD+yKK!Nn@9peYc_Q>SQ-{??=1KN= zkg2};?)E>k`P@56RR9|AlR5s4X=_pZ`-=@Uy}zWfR#GlrWMuL39)PYBd3(y@jY!wh zvxDzd_E1RmJUiq<$M=bZd0HK%?F4D1IsrXPU&rrER}|Qe5>Q3M#M0JTs#TW!?i~!{ z$dYy<)~}?BGb0UJ`j+@6i&On4agYTN{{aGn$e(|?+Q<3Ic3HVfn%>8|{RMR`l5K~0 z*E9WY)i?cabvga6XLX?UOb2Su^q$r4dRGT(V>&?pr`?f_+=^qOe{VMCRaV4izw5nO zpruvo95ygA!{_q_UWVn{Q+*>Q))H0v1!0ZSOC_Nt;cnIc&s!}{q2ui~mBra&h&8bJ z-2Jjndpx~fbR!OVobGD;l@_WWYxDEcwlK&aF)fp7lmAjbhkRm#5rfKvNv+A)Z<=S>Qa^oMw3AA>| zlUN@&#*1N62zxCL2(u_3^#Ot8)-jV?J;bTU)&l&s)1Ai!>+OTRqrJdEy z0c^*>gH3P2^ZgbcwxkyT8IPDxc|a-PqukhD`TT&+eh}gy#$!NsF}8HL=);Ht!nlt^ z5}}A8Oee!Qi18SV3(|H8bCpeE@wPhA5&WpuP=gqc zLPeLf{j^Rva{ykT#DEMj2>u*XPUv*TIEe8WlpS;IMa58r8TH|4BT>kv3yfqC<1wff zPwTEnbr>-E=HeAlj42%AQyJ(m#$!NUU~Jy}`V@&7=RxFiFiu1vVorC$4Prb7mF# z9UkH!0BsTEqv@P(5aThZo;1k5xIz#?2<|3fh$99noJ?_p7>|MYbX@6EBj@mG2M*_h z7I8SG&kbQb2BYz@_R~Jxd%h0{2{83TcsPLzCz?!?!x)c2*#+6UO84A6aX=#9XYh?~ zK3@HBI)5C*cnriRjkV8K1BR%NBgzP+fzKz?A|Awe48(3A?Nc1X5yA}la8D#rFd?cz zj3+?V7dknG1CRR*u8AQb{$x2ZJc#iaP)|A>TuLv$&O=Ct!POJ84unPFWH&N0gl`9- zJ(;RpoJm7JoNG={V&dmCQJq3XOIm)Nwq|{G6p%$vn>DMtAks+Pp1HIZF0=WKG*~v*qK&Cd>5SNITQIw1s@p~d8@Y4a58L`uZKN{fL_A8h zIEXh;CaN@@ZMS7T$khnIqTN^2**=%m#nkPtNa9^;3hsuhDTXgI)onvw|z)#LGnEcq6(=8hVcSwW5Bc4f^m$$iHAFIZ?&wg@!1%Ji;r!yx{r|a1l@6X=*=ojk0x3iD)MCfg%4y%pKlkD*z zQ+@NF+izy`xp$JP05sm`bNoxw)}r{=mm6q$e@6=>Yv7c1JdHD~^f&%52Q5tccBi*L$-- zORLs7Y+z)D&*uxg49mBt`bJEwC93ob!WyNQNT>czV6)MjY}u-PQOjEmS|&=I5oy)jvOiST7^D;MV%<_z>#`Ji4s*n$&OXf-*&4 zEl0!G|0LDnEQ@z93b_*LZgrB0bIx34Z{%9b&1$E_Db|m!W1SriBHg%VzK6-KVcmO` zXGNRPVC&*?X=fhn%!5!H6Z7=mP0~W`q+Us-GUe6c7gCxa!^&8sp;yG|DD}ZUkU@3#=1_~uCWD`&!_7r8bDt=&t5{T(yK#kTo-)h z`})OO+izo`St%Fm(#mi>XRBg+KG}hlr+X?p?)H@rpe3s3YniXgwe4`^#y^M?Xzh?E zu|95$7sI3w(0ywJbGU(AAyzUso0V!C?SbyyXK!V(IAOQ%A_t|vjSp_XRwRWzCJc(s ze)S@-41seQPUh9hcE62HoI4Du@@MnfX_&ZXwPbp105vaXAsl-{^Or3?o#yv3dK3>J z4XPHg8)lf5J7wOy9(k93>#^sg8EYPh*FojIB?0Wi$Bb~FGx!q*{!l=*Vby#}JFA-m z*p7h*o8E%w`z<_dNiP609xolPKI$1<1rW)r0o*sE^(ncP#mF{lOW<@z%ZK()*!~CV0CrZw$C1rIN$;D0Xt*Z zH*Zsc9L9JQ$gb|_61$sN9RM(5D8$4k5Ly_YPIW{0jsewGT+8YD0Qx)(QAA0|5s#v& zv_=N79RshUvc4cfXn<(M5erDjA~G4GA$-R`>V^=|;Dr?igfdJ+gnWWIXTfBk1~DFm ziY{sUX`OK90K7nn0U2Tt{5hta(CLhE5aTf@JLcGnilGQI>ci1SqL57&7|9^UV^A%g z)?JV4Fktk}#Veo~Q#iz@GSFd+$AG-R*u44mDH1czgUIJ#oQOiiobH4h#CQzGCl72+vUT`ST5KaZq;AF4{F&+b}8(DgeZVVqg zJj6i&+9JqD(>dKB#$!-DX^?$!g&=|u+)cs|M+{Urnc@a99s}{|xYDOa&f(Jz9L@(V z;&4iz8^U-DM&o1cr+v8hd>;@JVCskPZ~_-jG?^xcF&=}m3$k^U?zws5fJDB};2Yh1 zy!zpE{y2#77>G|AYoDzK3{f9Plo3h;pHHSmJc#iah}}Tir#OZqgcI)#dswEQ}4&HCd}Ko&i1)~xD+NF#N7=F(cY%;q@OjAQ&&JlugJx5zS&3ozx^ n5^DP!BZV_TFH3Jr0fFchI?LUjWa*u&xj{Y2Miqm6_U-I{SD{~m literal 0 HcmV?d00001 diff --git a/web_platform/fixtures/obstacle/ort-init-failure.onnx b/web_platform/fixtures/obstacle/ort-init-failure.onnx new file mode 100644 index 0000000000000000000000000000000000000000..96162dd8b309e6a431a68a56c475bc64435ed868 GIT binary patch literal 4510 zcmbtYTaVjB6sCYgSyc!Y@xWX0a}&+gm%f#v7Eoc^v6R-O;1tPt3&UWl=}p{FK*qUd&G_>x?hmv(mkB zrl{NXH@5c2pZh&w;6eHQtD}#~v-^x)f%k3RVP z3-@<)^kua)epb5GcxtOmzQ(~HjF`fTjquT*{p00(rNVU#dwid&Mo)(nXC~t)-w>louXklNT?Yt*r+!n&$PyS#QTUdXtI`*p|&W(p$+cx2Ps(*cQiOu&9d%)mHGaIxDdG?pp*3^StWg@)m6^0|;|=f3DS>=If+4pW`?sb$jfWsMU){o` zIR`hx5%)WX@vT>a5FVyUn_w|lLdRO+ohI>3(T$b2Uw7F8Ksekjpg0DG8p;kz=}g5^ zCsG1Q>a-Uu#2A1(M1x*~JsL}d(lJj(#xljY>_rPP2G9nAOicx@QwPcj$+bu_ohFK_ zUVso|2;fj3vfd7kCpuv2Q8AEWbK61#rBFmX#DHpLoFeu%xPJk~U{!ke> zXAc$YG-a6-sb)IM`ZN)D*#ba~8hf`3nJ}4YrV^2=O!UeQ@CAgoPKPd?jS?Xhmnmbh z;99GsSA2*uBwr5Y?H7DbX|(kTFqlFJYc92r{ifmp#sFUH^SyknM86fk)D0#3o-`a?x5W?9#>i?S*&3?CJ}A@853d*ph1sYFRC~D z{M1RBBsvp`fWu(aUfn{B0kp2(?2X4Y3JpSArig85)e8_}3;-PDKgT0Lr#wlZfnYJ~ zcTR{gAU|NPy)uFc#*qJM$}8Pz^nJbwFb2?q54wFU z2$66Zr+8_BLdSZw3o-`q0;b$mJf4UwjuoD(^q7m&e!qnn17HDD?k$-pCS#sSC8SJZ z-K!Z7F$TceAlxgXc*eM5G3uGflK03DG4{#dhn<#4ARQ}|vs8$-NMZ4oldd|uFB z03?A1pJpOyhmMnCc?ohi%IC#oi8lQ?x_gX0Et%9RFcQtT{Ws>0s-0foJvy;-`Nk|q zW}2_5({A;^4AW@)x*JKGA!T8~RLr+Dcv_*@qEA6Ln$WF$_TtCufo2elzDI uxM)YRRbH1<``Y3Ug()G@7#fwullCAWU}`NZ(y}Z@+o1{*MiU3u*Z%;Kymy`e literal 0 HcmV?d00001 diff --git a/web_platform/fixtures/obstacle/wrong-graph.onnx b/web_platform/fixtures/obstacle/wrong-graph.onnx new file mode 100644 index 0000000000000000000000000000000000000000..d87db31d14accc5505360b85781082e30b8204f0 GIT binary patch literal 4563 zcmbtYTaVjB6z)PRmZ(By5f8joeyo}>S6}*q+C`L0MN0}oP@yoJOcIyGw(LoEx2t{T zXTTG`lkeD0vYQ}3>TmSL%obgzfVs->1~d(w{-W(?cS|>AC`Xvy^o58anrn9jwZHPKD181t=zM9}6_#sPOX>I>$F8a?;px$_4AFuOtVzRLfzXNyQg=m6F9s%RU>xU~G^)yx-@8m~O8>{{sWbZmkVOa=z<{{)HRTD z08K=FRolCN`%DtA#OL?6TsDDqVk$c_*=)473ks~F%wOBFvy;)RH0#uz6LFj`mnNIk z>vL?Nn>rLnMPBV5S`nK90Vb$$YRU;^pPLh-Gv&;>kuA%-3~IAlOF^$FFC54w&0dtY zvJS*xl2+r=Iy=JA>r|}4Hf*LewGtOW)udaUTpQNWIA3KXMJR~gQs)*~<3`JrA`Wrk zTQg_%8l_=Xm~rDaUh_`Nv@U-&10pmi^VMwLh6Nk2dXKZksKKe(81bE&HSD}B7O7iJ zC^W$wUt*!v29%4Sv<2y&xw?pe6rAK|w(74r+0GPj&RM$LEKbd$qPlZc`y1!01whiP zi-3GZDS>=If+4qR`!}Y06%IhGFYUs)IR!Vv5qCQW;k8eM5FW%)lVBlNLWf%68%N<) z(Y2K~mz!)JARO)*P#6M34P*5#`FVToc~ zcB1(hJ!pMFrltbdssm+&yOWn2S?kc-$HV6lzpo6Ovx5qC9J54OB06Pz_&mZJr$d+4Mv;(;%b2lHaIIC;Dc;BElP?GI`UgIzG}`zC7)&9AHJ4h* zZc}j&qX)0?`A)uK30E>is)!_3-9aUMj2`)YpFUeS4z!TUOe8s%NY5_wc^N%$H_)yc zk1MUCB-Ah+lZdwNjPWsg(4fb)57nDpe(ETWBAtjxz+o_Ir*1w*4_aGqcE;lxg$AK5 zW5hPJ>ICpHdI0wFpW_vvV;)7&K(LT?JIBZ9k?%3rP8q=jW61wF=81w@oyzzaJz%~_ z*w!1QDqcNsObE?*{5Ics7(Hm-7u_xvgowBdV|=tgp<|uec^N%;9#d{B9*;y4h6?Xh zdd-D#x8Hn>9x#t7cb1G4lOa!}5>iH??$nI?7(HN35bl&wJYih15cNzX(OcyE7`x=} z!cI#hkPa2fSu8}@+3#>4qer~QF(`qXH)@Pl2PF^hY@w34v;GJV-v+DBi^7~XgO;Wi zguVmWyxC~wyOlvRnzh3iOG2onIj`8gv4)nYw?hJ9aSFBf@z>`-0M(k4@#M zEynEzTh{ITblQv$82)6dWmAOABAsRQ697q|!Ka=`nxW%)wm1j59p%$(yg-}&6x}_> zo`y_nWf+NOoBkVfN6}0#@E)DGEK~drGBlHPMV)q|2WFUB+n4P~S`R663#M$gp~2G% z#TI=Ey4fU3aO?n%dx~KQayU7Q4EK{c9l}Lhk}c9IpV*fcKPXHHk;c#<&mK1i0RdBM TUXYe~HrxzVm@w)%xV-!u>j!j9 literal 0 HcmV?d00001 diff --git a/web_platform/fixtures/obstacle/zero-action.onnx b/web_platform/fixtures/obstacle/zero-action.onnx new file mode 100644 index 0000000000000000000000000000000000000000..8916f0526d99ed2f4d30ebdec4c14ca69e3ed5c0 GIT binary patch literal 4568 zcmbtYTaVjB6lS3nOH`qGD%z#+p;Iw-LCeT zp8@auNxow{$!>OBKv-F?eSGG8bFMzehaWxq>!A1R(T8U2(mbnz!(*JC$=ri)d$-qS zv9b?7KRh_RbN3)%@BPQ)$?dnc%&R{SZXMlowsQSETb%ca)^Ff`zSsL8KdFvB?%g>& zI6CaTe|UhumnWa~Hv8Fobmz&f-g|dHDgOw1pB4+_rg^y>O>D8q&zClHN52Ob!P=G; zR0$s4=MV1(&Q!mBJ;6=zLzcSI+WymA^q-$V$MM2GUgzn=WMdnE^`y*CD_b7t%fh(y zWMQ$&l`D%!*rYAjpyt`D^esJzNqD10OW^{Pah|y{UwrizESiBZGpT!F%QUAwut5mY zn>@|jcV=lTkX{F0J{!%`VpLOU8@PA#bT+?68JorQ3SqdnWO!rAaCCjia1RCddnlmn zE!jgE-arB6%93BO`H7vH)xv#4D!$08lq^PC_SpaiV}s=3{XQSWbb0;$9~ek>Yi%Hs z175R%&?M*ook?)z(&;p_RqYP}uhE}s|G)dX*4ZqJG@H?$pc8%wgw6|EV5f*%_bgqe zu7Q*TXd>#f+Ft$Jdy;r1KEJo+vI(pcQ`wQpW}~%TP+%2h{>F}-os4FsS*Pxth~s>@ zG})wnKF0#Os6%m7_urS>p%=9X*Djbvm|wP78N^HoMtgo5ZTb#9S0ZnR7( z;t&_UXXdPaMrl|TX56@q*Sym*t!rS-fCvrBd^MZ5VZjEhUgK;rYH(^cMtrAc4LdK3 zMe0@)3QaJ_mw3=>1Ik5E+Jbb?TwO#!3QqDfTlLqRY-b8M=PX@r9!|}oqPlZc`y2bK z1whj4i-3GZDS>=If+4qR`!}Y09S%UOPwj(qa|&*TE$&th!fQu_5FW%)lVBlNLWf%6 z7f0b$(Y2K~r<-gZAZ+d$P#6M34P*5#`F zVToc~cB1(hJ!pMFrltbdssm+&yOWn2S?kc-$HV6lzpo6Ovx5qC9J54OB06Pz_&mZJr$d+4Mv;(;%b2lHaIIC;Dc;BElP?GI`T?I)8f|<645kpm znoBKYx2d>?(Sz6cd?#PAgew^$RYVf2?w}GrMvwfyPoFIu2U>q-U4;yo?^W z8)#RJ$CcJm5^9)^Nkm(B#`qXLXwc)@LG@;rpE`=8NGBo^a2Sl*shf|{gVxrYo$5=pbgWZ5FQW&~W6EvC?h zP~p8wuemVp_M4B<1LiU1&XSR0GUSO=Ldq!AotkkUqX(=B!ksdTCyXl=qMnH)dWU=; zW0(A0*lCFb(xE~*i-ia~`yK9M^oaL31|@LwMvc+xpyc75EmRVB)*s>F+hEmsQJB+a z(9+a`&~G4{Hyb^9xiV-*vvwGx$*cL%(xz{x_OV$?H1zo^vv}UN=EvK1`;vN6n^k#v zJf(r>0%O1${}adymgek>l)fEcoJQMEtY&K6`4u9gK_@Vqsr$#eV+V6RBAgerE-20U z*i??%Vq9*pW!=tCr_K0);ZL?&HbuxR(pg5I07wE2KJ`S>3?0w$)dA#oluxtq0&V(J zboUs08ZxPsVI-Pu`ftn~MKis?d$i-SOz|;fXeQ~3I_*Xe%rLdKuiBBc9#ZBOOxbKh zgQpdWE&3F6vq_X-+W{Q+6vGhYaB>zI?k96Pgp0N$TclM!v9ByXC`<^E#?T(null), notificationId = useRef(0), loadInFlight = useRef(false), + policyLoadInFlight = useRef(false), importInFlight = useRef(false), adapter = useRef(new MainThreadPhysicsAdapter()), root = useRef(null), @@ -319,6 +321,9 @@ export function App() { [controllerStatus, setControllerStatus] = useState(), [selectedPolicyPath, setSelectedPolicyPath] = useState(), [policyStatus, setPolicyStatus] = useState(), + [navigationTargetMode, setNavigationTargetMode] = useState(false), + [trainingDeployment, setTrainingDeployment] = useState(), + [showPerceptionRays, setShowPerceptionRays] = useState(true), [projectMaps, setProjectMaps] = useState([]), [editorDocument, setEditorDocument] = useState(null), [editorDrafts, setEditorDrafts] = useState>(() => new Map()), @@ -444,6 +449,13 @@ export function App() { setEditorSelection(selection ? { kind: 'body', bodyId: selection.bodyId } : null); if (selection) setRightOpen(true); }, + onNavigationMode: setNavigationTargetMode, + onNavigationTarget: (target) => { + adapter.current.setNavigationTarget(target); + const snapshot = adapter.current.snapshot(); + state.setSnapshot(snapshot ?? undefined); + setPolicyStatus(snapshot?.rlPolicy); + }, onFrame: (frame, fps, snapshot) => { const memory = (performance as Performance & { memory?: { usedJSHeapSize: number } }) .memory?.usedJSHeapSize; @@ -550,8 +562,14 @@ export function App() { viewer.current?.setMapDisplay(showVisualMap, showMapCollision); }, [showVisualMap, showMapCollision]); useEffect(() => { - viewer.current?.setParametricMapAssets(placedMapAssets, mapSceneDraft.changedIds); - }, [placedMapAssets, mapSceneDraft.changedIds]); + viewer.current?.setParametricMapAssets( + trainingDeployment ? [] : placedMapAssets, + mapSceneDraft.changedIds, + ); + }, [placedMapAssets, mapSceneDraft.changedIds, trainingDeployment]); + useEffect(() => { + viewer.current?.setShowPerceptionRays(showPerceptionRays); + }, [showPerceptionRays]); useEffect(() => { if (window.innerWidth < 900) return; try { @@ -585,6 +603,8 @@ export function App() { path: string, requestedMode?: UrdfLoadMode, requestedSceneAssets?: readonly PlacedMapAsset[], + requestedDeployment?: PolicyDeployment, + requestedPolicy?: { data: Uint8Array; path: string }, ) => { if (!manifest.current || loadInFlight.current) return false; // 普通模型重载始终使用上次成功应用的地图基线。只有场景提交入口可以显式 @@ -623,6 +643,8 @@ export function App() { baseMode: baseModeRef.current, enhancements: urdfEnhancementsRef.current, mapAssets: sceneAssets, + trainingDeployment: requestedDeployment, + trainingPolicy: requestedPolicy, onProgress: ({ value, label }) => setImportProgress({ title: '正在准备仿真', @@ -670,9 +692,10 @@ export function App() { await activeViewer.setVisualMaps([]); let visualMapWarning: string | undefined; try { - const assets = manifest.current - ? visualMapAssets(manifest.current, placedMapAssetsRef.current) - : []; + const assets = + manifest.current && !requestedDeployment + ? visualMapAssets(manifest.current, placedMapAssetsRef.current) + : []; await activeViewer.setVisualMaps(assets); } catch (error) { visualMapWarning = `视觉地图加载失败:${error instanceof Error ? error.message : String(error)}`; @@ -704,11 +727,12 @@ export function App() { setAppliedMapAssets(committedMapAssets); } activeViewer.setParametricMapAssets( - placedMapAssetsRef.current, + requestedDeployment ? [] : placedMapAssetsRef.current, summarizeMapSceneDraft(placedMapAssetsRef.current, appliedMapAssetsRef.current) .changedIds, ); adapter.current.releaseRetired(); + setTrainingDeployment(requestedDeployment); sessionSwapped = false; return true; } catch (error) { @@ -734,6 +758,7 @@ export function App() { if (retained && previousEntry) state.setEntry(previousEntry); adapter.current.setPaused(previousPaused); state.setPaused(previousPaused); + state.setSnapshot(retained ?? undefined); setControllerStatus(retained?.controller); setPolicyStatus(retained?.rlPolicy); return false; @@ -874,6 +899,7 @@ export function App() { setSelectedControllerPath(undefined); setControllerStatus(undefined); setSelectedPolicyPath(undefined); + setTrainingDeployment(undefined); setPolicyStatus(undefined); setProjectMaps([]); setCommittedEditorDocuments(new Map()); @@ -1962,26 +1988,62 @@ export function App() { setControllerStatus(undefined); state.setSnapshot(adapter.current.snapshot() ?? undefined); }; - const loadPolicyBytes = async (data: Uint8Array, path: string) => { + const loadPolicyBytes = async (data: Uint8Array, path: string, expected?: PolicyDeployment) => { + if (policyLoadInFlight.current || loadInFlight.current) + throw new Error('模型/策略正在加载,请稍后重试'); + viewer.current?.setNavigationTargetMode(false); + const previousPaused = useAppStore.getState().paused; + const previousSession = adapter.current.session; + policyLoadInFlight.current = true; state.setLoading(true); - setImportProgress({ - title: '正在加载强化学习策略', - label: '初始化 ONNX Runtime', - detail: path, - value: 0.55, - }); state.setDiagnostic(undefined); + adapter.current.setPaused(true); + state.setPaused(true); try { - const status = await adapter.current.loadRLPolicy(data, path); - setPolicyStatus(status); + const deployment = resolvePolicyDeployment(data, expected); + if (deployment?.terrain) { + const entry = useAppStore.getState().selectedEntry; + if (!entry || !manifest.current) throw new Error('请先导入并加载Go2机器人'); + if (!(await loadEntry(entry, 'mjcf', undefined, deployment, { data, path }))) + throw new Error( + useAppStore.getState().diagnostic?.detail ?? '配套训练地图加载失败,策略未启用', + ); + } + setImportProgress({ + title: '正在加载强化学习策略', + label: '初始化 ONNX Runtime', + detail: path, + value: 0.55, + }); + state.setLoading(true); + const session = adapter.current.session; + const status = deployment?.terrain + ? adapter.current.snapshot()!.rlPolicy! + : await adapter.current.loadRLPolicy(data, path, deployment); + if (session !== adapter.current.session) throw new Error('模型已切换,策略加载取消'); + if (deployment?.terrain) { + adapter.current.setPaused(false); + state.setPaused(false); + setShowSensorCamera(true); + viewer.current?.setShowSensorCamera(true); + } + setPolicyStatus(adapter.current.snapshot()?.rlPolicy ?? status); state.setSnapshot(adapter.current.snapshot() ?? undefined); notify( 'ONNX 策略已加载', `${status.taskName} · ${status.observationSize} → ${status.actionSize}`, ); } catch (error) { + if (adapter.current.session === previousSession) { + adapter.current.setPaused(previousPaused); + state.setPaused(previousPaused); + state.setSnapshot(adapter.current.snapshot() ?? undefined); + setPolicyStatus(adapter.current.snapshot()?.rlPolicy); + } state.setDiagnostic(diagnostic('仿真', error, path)); + throw error; } finally { + policyLoadInFlight.current = false; setImportProgress(undefined); state.setLoading(false); } @@ -1993,45 +2055,49 @@ export function App() { return; } setSelectedPolicyPath(path); - void loadPolicyBytes(file.data, path); + void loadPolicyBytes(file.data, path).catch(() => {}); }; - const importPolicy = (file: File) => { - void (async () => { - try { - if (!/\.onnx$/i.test(file.name)) throw new Error('请选择 .onnx 文件'); - if (file.size > 64 * 1024 * 1024) throw new Error('ONNX 策略不能超过 64 MiB'); - const path = normalizeProjectPath(file.name), - data = new Uint8Array(await file.arrayBuffer()); - if (manifest.current) { - const index = manifest.current.files.findIndex((candidate) => candidate.path === path), - files = manifest.current.files.slice(), - entry = { - path, - data, - size: data.byteLength, - source: 'file' as const, - mimeType: file.type || 'application/octet-stream', - }; - if (index >= 0) files[index] = entry; - else files.push(entry); - const totalBytes = files.reduce((total, item) => total + item.size, 0); - if (totalBytes > DEFAULT_IMPORT_LIMITS.maxTotalBytes) - throw new Error('加入 ONNX 后工程总大小超过 512 MiB'); - manifest.current = { ...manifest.current, files, totalBytes }; - state.setProject( - manifest.current.name, - files.map(({ path: filePath, size }) => ({ path: filePath, size })), - manifest.current.entries, - manifest.current.selectedEntry, - ); - state.setSnapshot(adapter.current.snapshot() ?? undefined); - } - setSelectedPolicyPath(path); - await loadPolicyBytes(data, path); - } catch (error) { - state.setDiagnostic(diagnostic('仿真', error, file.name)); + const importPolicy = async (file: File, expected?: PolicyDeployment) => { + try { + if (!/\.onnx$/i.test(file.name)) throw new Error('请选择 .onnx 文件'); + if (file.size > 64 * 1024 * 1024) throw new Error('ONNX 策略不能超过 64 MiB'); + const project = manifest.current; + const path = normalizeProjectPath(file.name), + data = new Uint8Array(await file.arrayBuffer()); + if (project !== manifest.current) throw new Error('工程已切换,策略导入取消'); + const projectedBytes = + (project?.files + .filter((item) => item.path !== path) + .reduce((total, item) => total + item.size, 0) ?? 0) + data.byteLength; + if (projectedBytes > DEFAULT_IMPORT_LIMITS.maxTotalBytes) + throw new Error('加入 ONNX 后工程总大小超过 512 MiB'); + await loadPolicyBytes(data, path, expected); + if (project && manifest.current?.id === project.id) { + const files = manifest.current.files.filter((candidate) => candidate.path !== path); + files.push({ + path, + data, + size: data.byteLength, + source: 'file', + mimeType: file.type || 'application/octet-stream', + }); + const totalBytes = files.reduce((total, item) => total + item.size, 0); + if (totalBytes > DEFAULT_IMPORT_LIMITS.maxTotalBytes) + throw new Error('加入 ONNX 后工程总大小超过 512 MiB'); + manifest.current = { ...manifest.current, files, totalBytes }; + state.setProject( + manifest.current.name, + files.map(({ path, size }) => ({ path, size })), + manifest.current.entries, + useAppStore.getState().selectedEntry, + ); + state.setSnapshot(adapter.current.snapshot() ?? undefined); } - })(); + setSelectedPolicyPath(path); + } catch (error) { + state.setDiagnostic(diagnostic('仿真', error, file.name)); + throw error; + } }; const togglePolicy = (enabled: boolean) => { adapter.current.setRLPolicyEnabled(enabled); @@ -2372,8 +2438,29 @@ export function App() { onSelectPolicyPath={setSelectedPolicyPath} onLoadPolicyPath={loadPolicyPath} onImportPolicy={importPolicy} + compileTrainingScene={(coordinates) => { + if ( + mapSceneDirty || + trainingDeployment || + useAppStore.getState().loading || + loadInFlight.current + ) + throw new Error('请先应用地图草稿;训练部署/加载中的场景不能同步'); + return adapter.current.exportTrainingTerrain(appliedMapAssets, coordinates); + }} + trainingSceneMaps={appliedMapAssets} + trainingSceneDirty={mapSceneDirty || Boolean(trainingDeployment)} onTogglePolicy={togglePolicy} onPolicyCommand={setPolicyCommand} + navigationTargetMode={navigationTargetMode} + onNavigationTargetMode={(active) => viewer.current?.setNavigationTargetMode(active)} + onResetNavigationTarget={() => { + viewer.current?.setNavigationTargetMode(false); + adapter.current.resetNavigationTarget(); + const snapshot = adapter.current.snapshot(); + state.setSnapshot(snapshot ?? undefined); + setPolicyStatus(snapshot?.rlPolicy); + }} onRemovePolicy={removePolicy} onDataRecorderConfigure={configureDataRecorder} onDataRecordingStart={startDataRecording} @@ -2570,6 +2657,22 @@ export function App() { progress={importProgress} /> setToast(undefined)} /> + {trainingDeployment && ( +
+ 训练配套物理地图(编辑器地图未更改,重载模型恢复)。 + {trainingDeployment.terrain?.approximation && '训练专用离散近似。'} + {(policyStatus?.observationSize === 81 || policyStatus?.observationSize === 97) && ( + + )} +
+ )} {Boolean(state.snapshot?.model.ncam) && (showSensorCamera ? (
void; + onResetNavigationTarget?: () => void; onResetJoints: () => void; onToggleJointLimits: () => void; onToggleAdvanced: () => void; @@ -101,7 +113,10 @@ export function WorkspaceToolsPanel({ onRemoveController: () => void; onSelectPolicyPath: (path: string) => void; onLoadPolicyPath: (path: string) => void; - onImportPolicy: (file: File) => void; + onImportPolicy: (file: File, deployment?: PolicyDeployment) => void | Promise; + compileTrainingScene?: TrainingSceneCompiler; + trainingSceneMaps?: readonly PlacedMapAsset[]; + trainingSceneDirty?: boolean; onTogglePolicy: (enabled: boolean) => void; onPolicyCommand: (command: RLCommand) => void; onRemovePolicy: () => void; @@ -251,9 +266,14 @@ export function WorkspaceToolsPanel({ loading={loading} onSelectPath={onSelectPolicyPath} onLoadPath={onLoadPolicyPath} - onImport={onImportPolicy} + onImport={(file) => { + void Promise.resolve(onImportPolicy(file)).catch(() => {}); + }} onToggle={onTogglePolicy} onCommand={onPolicyCommand} + navigationTargetMode={navigationTargetMode} + onNavigationTargetMode={onNavigationTargetMode} + onResetNavigationTarget={onResetNavigationTarget} onRemove={onRemovePolicy} /> @@ -262,7 +282,12 @@ export function WorkspaceToolsPanel({ defaultOpen={false} badge={} > - + ) : ( diff --git a/web_platform/src/components/charts/ScalarChart.test.tsx b/web_platform/src/components/charts/ScalarChart.test.tsx new file mode 100644 index 00000000..c72457ba --- /dev/null +++ b/web_platform/src/components/charts/ScalarChart.test.tsx @@ -0,0 +1,47 @@ +import { fireEvent, render, screen } from '@testing-library/react'; +import { ScalarChart } from './ScalarChart'; +import { ScalarChart as TuningChart } from '../../tuning/ScalarChart'; +import type uPlot from 'uplot'; +const spies = vi.hoisted(() => ({ + create: vi.fn(), + destroy: vi.fn(), + setData: vi.fn(), + setScale: vi.fn(), +})); +vi.mock('uplot', () => ({ + default: class { + over = document.createElement('div'); + scales = { x: { min: 0, max: 10 }, y: { min: 0, max: 10 } }; + constructor( + options: unknown, + public data: unknown, + ) { + spies.create(options, data); + } + setScale = spies.setScale; + setData = spies.setData; + setSize() {} + destroy = spies.destroy; + }, +})); +it('共享原tuning API、EMA只修改曲线、悬停保留原值、缩放及卸载清理', () => { + expect(TuningChart).toBe(ScalarChart); + const points = [ + { step: 1, wallTime: 0, value: 0 }, + { step: 2, wallTime: 0, value: 10 }, + ]; + const view = render(); + const [options, data] = spies.create.mock.lastCall!; + expect(data).toEqual([ + [1, 2], + [0, 6], + ]); + expect(options.series[1].value({} as uPlot, 6, 1, 1)).toBe('10'); + expect(points[1].value).toBe(10); + fireEvent.click(screen.getByRole('button', { name: '训练与评估 Scalars 放大' })); + expect(spies.setScale).toHaveBeenCalled(); + fireEvent.click(screen.getByRole('button', { name: '训练与评估 Scalars 重置缩放' })); + expect(spies.setData).toHaveBeenCalled(); + view.unmount(); + expect(spies.destroy).toHaveBeenCalled(); +}); diff --git a/web_platform/src/components/charts/ScalarChart.tsx b/web_platform/src/components/charts/ScalarChart.tsx new file mode 100644 index 00000000..90b07ac7 --- /dev/null +++ b/web_platform/src/components/charts/ScalarChart.tsx @@ -0,0 +1,216 @@ +import { useEffect, useMemo, useRef } from 'react'; +import { RotateCcw, ZoomIn, ZoomOut } from 'lucide-react'; +import uPlot from 'uplot'; +import 'uplot/dist/uPlot.min.css'; +import type { ScalarSeries } from '../../training/types'; + +const COLORS = ['#38d39f', '#60a5fa', '#f59e0b', '#f472b6', '#a78bfa', '#fb7185']; + +function smooth(values: (number | null)[], factor: number): (number | null)[] { + if (factor <= 0) return values; + let previous: number | undefined; + return values.map((value) => { + if (value === null) return null; + previous = previous === undefined ? value : factor * previous + (1 - factor) * value; + return previous; + }); +} + +export function ScalarChart({ + series, + smoothing, + title = '训练与评估 Scalars', + xLabel = 'Step', +}: { + series: ScalarSeries[]; + smoothing: number; + title?: string; + xLabel?: string; +}) { + const host = useRef(null); + const chartRef = useRef(null); + const trackZoom = useRef(false); + const zoomRanges = useRef>>({}); + const prepared = useMemo(() => { + const steps = Array.from( + new Set(series.flatMap((item) => item.points.map((point) => point.step))), + ).sort((a, b) => a - b); + const columns: uPlot.AlignedData = [steps]; + for (const item of series) { + const byStep = new Map(item.points.map((point) => [point.step, point.value])); + columns.push( + smooth( + steps.map((step) => byStep.get(step) ?? null), + smoothing, + ), + ); + } + return columns; + }, [series, smoothing]); + + useEffect(() => { + if (!host.current || series.length === 0 || prepared[0].length === 0) return; + const element = host.current; + const width = Math.max(280, Math.floor(element.getBoundingClientRect().width)); + trackZoom.current = false; + const style = getComputedStyle(element); + const axisColor = style.getPropertyValue('--ui-text-tertiary').trim() || '#8fa0b5'; + const gridColor = style.getPropertyValue('--ui-border').trim() || '#213044'; + const chart = new uPlot( + { + width, + height: 280, + cursor: { drag: { x: true, y: true, setScale: true } }, + scales: { x: { time: false } }, + hooks: { + setScale: [ + (instance, key) => { + if (!trackZoom.current || (key !== 'x' && key !== 'y')) return; + const scale = instance.scales[key]; + if (typeof scale.min === 'number' && typeof scale.max === 'number') + zoomRanges.current[key] = { min: scale.min, max: scale.max }; + }, + ], + }, + axes: [ + { stroke: axisColor, grid: { stroke: gridColor } }, + { stroke: axisColor, grid: { stroke: gridColor } }, + ], + series: [ + { label: xLabel }, + ...series.map((item, index) => ({ + label: item.tag, + stroke: COLORS[index % COLORS.length], + width: 2, + spanGaps: true, + value: (_chart: uPlot, _value: number, _seriesIndex: number, index: number | null) => { + if (index === null) return '—'; + const raw = item.points.find((point) => point.step === prepared[0][index]); + return raw ? String(raw.value) : '—'; + }, + })), + ], + }, + prepared, + element, + ); + chartRef.current = chart; + trackZoom.current = true; + for (const key of ['x', 'y'] as const) { + const range = zoomRanges.current[key]; + if (range) chart.setScale(key, range); + } + const wheelZoom = (event: WheelEvent) => { + if (!event.deltaY) return; + event.preventDefault(); + event.stopPropagation(); + const bounds = chart.over.getBoundingClientRect(); + if (!bounds.width || !bounds.height) return; + const xRatio = Math.min(1, Math.max(0, (event.clientX - bounds.left) / bounds.width)); + const yRatio = Math.min(1, Math.max(0, (event.clientY - bounds.top) / bounds.height)); + const unit = event.deltaMode === 1 ? 16 : event.deltaMode === 2 ? window.innerHeight : 1; + const factor = Math.min(2, Math.max(0.5, Math.exp(event.deltaY * unit * 0.002))); + for (const key of ['x', 'y'] as const) { + const scale = chart.scales[key]; + if (typeof scale.min !== 'number' || typeof scale.max !== 'number') continue; + const ratio = key === 'x' ? xRatio : 1 - yRatio; + const anchor = scale.min + (scale.max - scale.min) * ratio; + chart.setScale(key, { + min: anchor - (anchor - scale.min) * factor, + max: anchor + (scale.max - anchor) * factor, + }); + } + }; + chart.over.addEventListener('wheel', wheelZoom, { passive: false }); + let frame = 0; + let lastWidth = width; + const observer = new ResizeObserver(() => { + window.cancelAnimationFrame(frame); + frame = window.requestAnimationFrame(() => { + const nextWidth = Math.max(280, Math.floor(element.getBoundingClientRect().width)); + if (nextWidth !== lastWidth) { + lastWidth = nextWidth; + chart.setSize({ width: nextWidth, height: 280 }); + } + }); + }); + observer.observe(element); + return () => { + observer.disconnect(); + chart.over.removeEventListener('wheel', wheelZoom); + window.cancelAnimationFrame(frame); + trackZoom.current = false; + chartRef.current = null; + chart.destroy(); + }; + }, [prepared, series, xLabel]); + + const zoom = (factor: number) => { + const chart = chartRef.current; + if (!chart) return; + for (const key of ['x', 'y']) { + const scale = chart.scales[key]; + if (typeof scale.min !== 'number' || typeof scale.max !== 'number') continue; + const center = (scale.min + scale.max) / 2; + const radius = ((scale.max - scale.min) * factor) / 2 || 1; + chart.setScale(key, { min: center - radius, max: center + radius }); + } + }; + const resetZoom = () => { + const chart = chartRef.current; + if (!chart) return; + zoomRanges.current = {}; + trackZoom.current = false; + chart.setData(chart.data, true); + trackZoom.current = true; + }; + + if (series.length === 0) + return ( +
+ 当前 trial 尚无 scalar 数据 +
+ ); + return ( +
+
+

+ {title} +

+
+ + + +
+
+
+

+ 悬停图例显示原始值;曲线可平滑。滚轮或拖拽缩放,右上角复位。 +

+
+ ); +} diff --git a/web_platform/src/map/customTrainingMap.test.ts b/web_platform/src/map/customTrainingMap.test.ts new file mode 100644 index 00000000..ea34e9a2 --- /dev/null +++ b/web_platform/src/map/customTrainingMap.test.ts @@ -0,0 +1,137 @@ +import { describe, expect, it } from 'vitest'; +import { + collisionHalfBounds, + trainingTerrainFromCompiledScene, + type CompiledMapGeometry, +} from './trainingMap'; +import { validateCustomTerrain, validatePolicyDeployment } from '../rl/deployment'; +import deployment from '../rl/fixtures/obstacleDeployment.json'; +import shared from '../../../training_server/tests/fixtures/custom-boxes.json'; + +const identity = [1, 0, 0, 0, 1, 0, 0, 0, 1]; +function geometry(patch: Partial = {}): CompiledMapGeometry { + return { + name: '__platform_map_one__obstacle', + type: 6, + position: [1, 2, 0.5], + rotation: identity, + size: [0.4, 0.3, 0.5], + friction: [0.8, 0.005, 0.0001], + collision: true, + static: true, + map: true, + ...patch, + }; +} +const pose = [...shared.spawn, ...shared.spawnQuaternion]; +const coords = { spawn: [-2, -1] as [number, number], target: [2, -1] as [number, number] }; +describe('customTrainingMap compiled collision export', () => { + it('共享payload逐字段roundtrip,不从种子生成布局,custom metadata严格一致', () => { + const layout = trainingTerrainFromCompiledScene([geometry()], pose, 6, coords); + expect(validateCustomTerrain(layout)).toEqual(shared); + const d = { + ...deployment, + terrainPreset: 'custom_boxes', + terrain: layout, + terrainParams: { size: layout.size, friction: layout.friction }, + }; + expect(validatePolicyDeployment(d).terrain).toEqual(shared); + expect(() => + validatePolicyDeployment({ ...d, terrainParams: { size: 8, friction: 0.8 } }), + ).toThrow(); + expect(() => validatePolicyDeployment({ ...d, terrain: undefined })).toThrow(); + }); + it('abs(R)*halfsize支持实例旋转/平移;排除机器人和装饰,所有碰撞障碍保留', () => { + const c = Math.SQRT1_2, + r = [c, -c, 0, c, c, 0, 0, 0, 1]; + const layout = trainingTerrainFromCompiledScene( + [ + geometry({ + name: '__platform_map_one__ground', + position: [0, 0, -0.05], + size: [5, 5, 0.05], + }), + geometry({ position: [2, 2, 0.5], rotation: r }), + geometry({ position: [-2, 2, 0.5] }), + geometry({ static: false, map: false, position: [0, 0, 0.32], size: [20, 20, 20] }), + geometry({ collision: false, type: 7, size: [30, 30, 30] }), + ], + pose, + 4, + coords, + ); + expect(layout.boxes).toHaveLength(3); + expect(layout.size).toBe(10); + expect(layout.boxes[1].pos).toEqual([2, 2, 0.5]); + expect(layout.boxes[1].size[0]).toBeCloseTo(0.7 * c, 12); + expect(layout.boxes[1].size[1]).toBeCloseTo(0.7 * c, 12); + expect(layout.approximation).toBe(true); + validateCustomTerrain(layout); + }); + it('primitive bounds有解析证明;mesh/hfield、不明静态/坑底/混合摩擦拒绝', () => { + expect(collisionHalfBounds(2, [0.4, 0, 0], identity)).toEqual([0.4, 0.4, 0.4]); + expect(collisionHalfBounds(3, [0.2, 0.5, 0], identity)).toEqual([0.2, 0.2, 0.7]); + expect(collisionHalfBounds(4, [0.2, 0.3, 0.4], identity)).toEqual([0.2, 0.3, 0.4]); + expect(collisionHalfBounds(5, [0.2, 0.5, 0], identity)).toEqual([0.2, 0.2, 0.5]); + for (const patch of [ + { type: 1 }, + { type: 7 }, + { map: false }, + { position: [1, 2, -1] }, + { friction: [0.8, 0.01, 0.0001] }, + { position: [15, 2, 0.5] }, + { size: [0, 0.3, 0.5] }, + { size: [NaN, 0.3, 0.5] }, + { name: 'ground', type: 0, position: [0, 0, 1] }, + { name: 'ground', type: 0, position: [0, 0, 0], rotation: [1, 0, 0, 0, 0, -1, 0, 1, 0] }, + ]) + expect(() => trainingTerrainFromCompiledScene([geometry(patch)], pose, 6, coords)).toThrow(); + expect(() => + trainingTerrainFromCompiledScene( + [geometry(), geometry({ friction: [1, 0.005, 0.0001] })], + pose, + 6, + coords, + ), + ).toThrow(/混合摩擦/); + expect(() => + trainingTerrainFromCompiledScene( + Array.from({ length: 257 }, () => geometry()), + pose, + 6, + coords, + ), + ).toThrow(/256/); + // Only explicitly named horizontal surface at z=0 may become standard floor. + expect( + trainingTerrainFromCompiledScene( + [geometry({ name: 'floor', type: 0, position: [0, 0, 0], map: false })], + pose, + 6, + coords, + ).boxes, + ).toHaveLength(1); + }); + it('严格有限数值/正半尺寸/字段与数量校验,不扩尺寸、不偷偷清除安全区障碍', () => { + for (const bad of [NaN, Infinity, -Infinity, -1, 0, -0, true]) { + const t = structuredClone(shared); + t.boxes[1].size[0] = bad as number; + expect(() => validateCustomTerrain(t)).toThrow(); + } + for (const patch of [ + { approximation: false }, + { actualObstacleCount: 5 }, + { path: '../../etc/passwd' }, + { spawnQuaternion: [0, 0, 0, 0] }, + { target: undefined }, + { size: 25 }, + ]) + expect(() => validateCustomTerrain({ ...shared, ...patch })).toThrow(); + expect(() => validateCustomTerrain({ ...shared, target: [1.9, 2] })).toThrow(/安全区/); + validateCustomTerrain({ ...shared, target: [1.91, 2] }); + validateCustomTerrain({ ...shared, target: [1.8, 2.7] }); + const t = trainingTerrainFromCompiledScene([geometry()], pose, 6, { target: [1, 2] }); + expect(t.boxes).toHaveLength(2); + expect(() => validateCustomTerrain(t)).toThrow(/安全区/); + }); +}); diff --git a/web_platform/src/map/trainingMap.test.ts b/web_platform/src/map/trainingMap.test.ts new file mode 100644 index 00000000..b43bd5c8 --- /dev/null +++ b/web_platform/src/map/trainingMap.test.ts @@ -0,0 +1,191 @@ +import dynamics from '../rl/fixtures/go2CompiledDynamics.json'; +import { readFileSync } from 'node:fs'; +import loadMujoco from '@mujoco/mujoco'; +import { describe, it, expect } from 'vitest'; +import { composeTrainingMap } from './trainingMap'; +import fixture from '../rl/fixtures/obstacleDeployment.json'; +import golden from '../rl/fixtures/obstacleRayGolden.json'; +import multiGolden from '../rl/fixtures/multiRingGolden.json'; +import multiFixture from '../rl/fixtures/multiRingDeployment.json'; +import { validatePolicyDeployment } from '../rl/deployment'; +import { SimulationSession } from '../simulation/SimulationSession'; +import { Go2ObstacleAvoidanceBindings } from '../rl/runtime/Go2ObstacleAvoidanceBindings'; +const deployment = validatePolicyDeployment(fixture); +const simple = new TextEncoder().encode( + '' + + deployment.jointNames + .map((name) => ``) + .join('') + + '', +); +describe('trainingMap', () => { + it('替换无限地面,保留精确halfSize/世界坐标/摩擦,添加PiP camera', () => { + const xml = new TextDecoder().decode(composeTrainingMap(simple, deployment)); + const doc = new DOMParser().parseFromString(xml, 'application/xml'); + expect(doc.querySelector('[type="plane"], keyframe')).toBeNull(); + expect(doc.querySelectorAll('geom[name^="__training_terrain_"]')).toHaveLength(25); + expect(doc.querySelector('geom[name="__training_terrain_0"]')?.getAttribute('size')).toBe( + '6 6 0.1', + ); + expect(doc.querySelector('camera[name="__platform_camera__"]')).not.toBeNull(); + expect(doc.querySelector('body')?.getAttribute('pos')).toBe('-5 0 0.32'); + }); + it('真实WASM编译Go2+共享地图;出生姿态、81维观测、32ray golden、reset和资源生命周期', async () => { + const module = await loadMujoco({ + wasmBinary: readFileSync('node_modules/@mujoco/mujoco/mujoco.wasm'), + }); + const original = readFileSync( + 'training_server/rl/src/assets/robots/unitree_go2/xmls/go2.xml', + 'utf8', + ); + const doc = new DOMParser().parseFromString(original, 'application/xml'); + // Strip visual meshes only for the lightweight physics test; collision/inertia/joints remain original Go2. + doc.querySelectorAll('mesh, geom[mesh]').forEach((e) => e.remove()); + const bytes = composeTrainingMap( + new TextEncoder().encode(new XMLSerializer().serializeToString(doc)), + deployment, + ); + const path = '/training-map-test.xml'; + module.FS.writeFile(path, bytes); + const session = new SimulationSession(module, path); + try { + session.configureDeployment(deployment); + expect(Array.from(session.data.qpos.slice(0, 7))).toEqual([-5, 0, 0.32, 1, 0, 0, 0]); + expect(session.model.ncam).toBeGreaterThan(0); + for (const [name, expected] of Object.entries(dynamics.joints)) { + const joint = session.model.jnt(name); + try { + const address = Number(joint.dofadr); + expect(Number(session.model.dof_armature[address])).toBe(expected.dof_armature); + expect(Number(session.model.dof_damping[address])).toBe(expected.dof_damping); + expect(Number(session.model.dof_frictionloss[address])).toBe(expected.dof_frictionloss); + } finally { + joint.delete(); + } + } + for (const [name, expected] of Object.entries(dynamics.geoms)) { + const geom = session.model.geom(name); + try { + expect(Number(geom.contype)).toBe(expected.geom_contype); + expect(Number(geom.conaffinity)).toBe(expected.geom_conaffinity); + expect(Number(geom.condim)).toBe(expected.geom_condim); + expect(Number(geom.priority)).toBe(expected.geom_priority); + expect(Number(geom.group)).toBe(expected.geom_group); + expect(Array.from(geom.solimp)).toEqual(expected.geom_solimp); + expect(Array.from(geom.friction)).toEqual( + name.includes('_foot_') + ? [deployment.terrain!.friction, 0.005, 0.0001] + : expected.geom_friction, + ); + } finally { + geom.delete(); + } + } + + expect(Array.from(session.model.qpos0.slice(7))).toEqual(Array(12).fill(0)); + expect(Array.from(session.data.qpos.slice(7))).toEqual(deployment.defaultJointPosition); + const bindings = new Go2ObstacleAvoidanceBindings( + session.model, + session.data, + (id, value) => session.setActuator(id, value), + deployment, + ); + for (const item of golden.cases) { + session.data.qpos.set([...item.position, ...item.quaternion]); + module.mj_forward(session.model, session.data); + const obs = bindings.observe(0, new Float32Array(12)); + expect(obs).toHaveLength(81); + expect(obs.every(Number.isFinite)).toBe(true); + item.depth.forEach((distance, i) => expect(obs[47 + i]).toBeCloseTo(distance, 5)); + expect(Array.from(obs.slice(11, 23))).toEqual(Array(12).fill(0)); + } + session.reset(); + const original = [...deployment.terrain!.target]; + const target: [number, number] = [-5, 0]; + bindings.setNavigationTarget(target); + target[0] = 999; + expect(bindings.navigationStatus().target).toEqual([-5, 0]); + expect(bindings.observe(0, new Float32Array(12)).slice(6, 9)).toEqual(new Float32Array(3)); + bindings.setNavigationTarget([-5, 3]); + const changed = bindings.observe(0, new Float32Array(12)); + expect(changed[79]).toBeCloseTo(0.5); + expect(changed[80]).toBeCloseTo(3 / deployment.terrain!.size); + expect(changed[8]).toBe(1); + expect(bindings.navigationStatus().distance).toBeCloseTo(3); + bindings.setNavigationTarget([999, -999]); + expect(bindings.navigationStatus().target).toEqual([5.5, -5.5]); + for (const value of [NaN, Infinity, -Infinity]) + expect(() => bindings.setNavigationTarget([value, 0])).toThrow(/有限/); + const status = bindings.navigationStatus(); + status.defaultTarget[0] = 999; + bindings.resetNavigationTarget(); + expect(bindings.navigationStatus().target).toEqual(original); + expect(deployment.terrain!.target).toEqual(original); + bindings.setNavigationTarget([0, 2]); + bindings.reset(0); + expect(bindings.navigationStatus().target).toEqual(original); + const box = deployment.terrain!.boxes[1]; + bindings.setNavigationTarget([box.pos[0], box.pos[1]]); + expect(bindings.navigationStatus().targetHeight).toBeCloseTo(box.pos[2] + box.size[2]); + bindings.apply(new Float32Array(12)); + module.mj_step(session.model, session.data); + session.reset(); + bindings.reset(0); + expect(Array.from(session.data.qpos.slice(0, 3))).toEqual(deployment.terrain!.spawn); + expect(bindings.observe(21, new Float32Array(12))).toHaveLength(81); + expect(bindings.rays).toHaveLength(32); + session.data.qpos[2] = 0.1; + module.mj_forward(session.model, session.data); + expect(() => bindings.observe(22, new Float32Array(12))).toThrow(/导航安全停止/); + } finally { + session.dispose(); + module.FS.unlink(path); + } + }, 60_000); +}); + +it('真实WASM完整48ray/97obs:CPU完整姿态/低障碍/坑边golden,绑定缓存按模型隔离', async () => { + const module = await loadMujoco({ + wasmBinary: readFileSync('node_modules/@mujoco/mujoco/mujoco.wasm'), + }); + const original = readFileSync( + 'training_server/rl/src/assets/robots/unitree_go2/xmls/go2.xml', + 'utf8', + ); + const doc = new DOMParser().parseFromString(original, 'application/xml'); + doc.querySelectorAll('mesh, geom[mesh]').forEach((e) => e.remove()); + const robot = new TextEncoder().encode(new XMLSerializer().serializeToString(doc)); + for (const layout of ['default', 'low', 'edge'] as const) { + const boxes = multiGolden.layouts[layout]; + const d = validatePolicyDeployment({ + ...multiFixture, + terrain: { ...multiFixture.terrain, boxes, actualObstacleCount: boxes.length - 1 }, + }); + const path = '/multi-ring-wasm.xml'; + module.FS.writeFile(path, composeTrainingMap(robot, d)); + const session = new SimulationSession(module, path); + try { + session.configureDeployment(d); + const bindings = new Go2ObstacleAvoidanceBindings(session.model, session.data, () => {}, d); + for (const c of multiGolden.cases.filter((c) => c.layout === layout)) { + session.data.qpos.set([...c.position, ...c.quaternion]); + module.mj_forward(session.model, session.data); + const obs = bindings.observe(0, new Float32Array(12)); + expect(obs).toHaveLength(97); + expect(bindings.rays).toHaveLength(48); + c.depth.forEach((v, i) => expect(obs[47 + i]).toBeCloseTo(v, 5)); + expect(Array.from(obs.slice(11, 23))).toEqual(Array(12).fill(0)); + } + bindings.setNavigationTarget([-5, 3]); + session.reset(); + const obs = bindings.observe(0, new Float32Array(12)); + expect(obs[95]).toBeCloseTo(0.5); + expect(obs[96]).toBeCloseTo(0.25); + bindings.clear(); + expect(bindings.rays).toEqual([]); + } finally { + session.dispose(); + module.FS.unlink(path); + } + } +}, 60_000); diff --git a/web_platform/src/map/trainingMap.ts b/web_platform/src/map/trainingMap.ts new file mode 100644 index 00000000..00eb0430 --- /dev/null +++ b/web_platform/src/map/trainingMap.ts @@ -0,0 +1,199 @@ +import { + validatePolicyDeployment, + type PolicyDeployment, + type TrainingTerrain, +} from '../rl/deployment'; + +/** Input must be flattened MJCF (mj_saveLastXML), so included ground cannot survive composition. */ +export function composeTrainingMap(source: Uint8Array, input: PolicyDeployment): Uint8Array { + const deployment = validatePolicyDeployment(input), + terrain = deployment.terrain; + if (!terrain) throw new Error('缺少训练地图'); + const doc = new DOMParser().parseFromString(new TextDecoder().decode(source), 'application/xml'); + if (doc.querySelector('parsererror, include')) throw new Error('训练地图需要已展开的有效MJCF'); + const world = doc.querySelector('mujoco > worldbody'); + if (!world) throw new Error('缺少worldbody'); + const robots = Array.from(world.children).filter( + (e) => e.tagName === 'body' && e.querySelector('freejoint, joint[type="free"]'), + ); + if (robots.length !== 1) throw new Error('训练评测需要唯一浮动基座机器人'); + const robot = robots[0]; + for (const child of Array.from(world.children)) + if (child !== robot && child.tagName !== 'light') child.remove(); + doc.querySelector('keyframe')?.remove(); + for (const key of ['euler', 'axisangle', 'xyaxes', 'zaxis']) robot.removeAttribute(key); + robot.setAttribute('pos', terrain.spawn.join(' ')); + robot.setAttribute('quat', terrain.spawnQuaternion.join(' ')); + // Match go2_constants.py FULL_COLLISION + obstacle terrain startup foot friction. + // Raw go2.xml does not contain these mjlab Entity compilation overrides. + for (const geom of robot.querySelectorAll('geom')) { + const name = geom.getAttribute('name') ?? ''; + if (!name.endsWith('_collision')) continue; + const foot = /^[FR][LR]_foot_collision$/.test(name); + for (const [key, value] of Object.entries({ + contype: '1', + conaffinity: '0', + condim: foot ? '3' : '1', + priority: foot ? '1' : '0', + group: '3', + friction: `${foot ? terrain.friction : 1} 0.005 0.0001`, + solimp: `0.9 0.95 ${foot ? 0.023 : 0.001} 0.5 2`, + })) + geom.setAttribute(key, value); + } + let actuators = doc.querySelector('mujoco > actuator'); + if (!actuators) { + actuators = doc.createElement('actuator'); + doc.documentElement.append(actuators); + } + for (const name of deployment.jointNames) { + const joint = robot.querySelector(`joint[name="${name}"]`); + if (!joint) throw new Error(`机器人缺少训练关节:${name}`); + joint.setAttribute('armature', name.includes('_calf_') ? '0.02' : '0.01'); + joint.setAttribute('damping', '0'); + joint.setAttribute('frictionloss', '0'); + if (!actuators.querySelector(`[joint="${name}"]`)) { + const motor = doc.createElement('motor'); + motor.setAttribute('name', `${name}_motor`); + motor.setAttribute('joint', name); + motor.setAttribute('gear', '1'); + actuators.append(motor); + } + } + for (let i = 0; i < terrain.boxes.length; i++) { + const box = terrain.boxes[i], + geom = doc.createElement('geom'); + for (const [key, value] of Object.entries({ + name: `__training_terrain_${i}`, + type: 'box', + pos: box.pos.join(' '), + size: box.size.join(' '), + quat: '1 0 0 0', + group: '2', + contype: '1', + conaffinity: '1', + friction: `${terrain.friction} 0.005 0.0001`, + priority: '1', + condim: '3', + rgba: i === 0 ? '0.35 0.4 0.45 1' : '0.65 0.35 0.2 1', + })) + geom.setAttribute(key, value); + world.append(geom); + } + // Camera's -Z optical axis looks along robot +X, +Y is image-up along robot +Z. + doc.querySelector('camera[name="__platform_camera__"]')?.remove(); + const camera = doc.createElement('camera'); + for (const [key, value] of Object.entries({ + name: '__platform_camera__', + pos: '0.3 0 0.05', + xyaxes: '0 -1 0 0 0 1', + fovy: '60', + })) + camera.setAttribute(key, value); + robot.append(camera); + return new TextEncoder().encode(new XMLSerializer().serializeToString(doc)); +} + +export interface TrainingSceneCoordinates { + spawn?: [number, number]; + target?: [number, number]; +} +export type TrainingSceneCompiler = (coordinates?: TrainingSceneCoordinates) => TrainingTerrain; + +export interface CompiledMapGeometry { + name: string; + type: number; + position: number[]; + rotation: number[]; + size: number[]; + friction: number[]; + collision: boolean; + static: boolean; + map: boolean; +} + +/** Exact analytic world bounds for primitive MuJoCo collision shapes (not visual bounds). */ +export function collisionHalfBounds( + type: number, + s: readonly number[], + r: readonly number[], +): number[] { + if (![2, 3, 4, 5, 6].includes(type)) + throw new Error('不支持的静态碰撞形状:mesh/hfield无法精确提取;请改用box等基本几何'); + return [0, 1, 2].map((i) => { + const row = r.slice(i * 3, i * 3 + 3); + if (type === 2) return s[0]; // sphere + if (type === 3) return s[0] + Math.abs(row[2]) * s[1]; // capsule, axial half-length + if (type === 4) return Math.hypot(...row.map((v, j) => v * s[j])); // ellipsoid + if (type === 5) return Math.hypot(row[0], row[1]) * s[0] + Math.abs(row[2]) * s[1]; + return row.reduce((sum, v, j) => sum + Math.abs(v) * s[j], 0); // box: abs(R)*halfsize + }); +} + +/** Uses already compiled, applied physics. Deliberately does not run terrain RNG or parse XML. + * Returns a candidate so the UI can expose unsafe suggested coordinates for explicit correction. + * validateCustomTerrain is mandatory before declaring synchronization/upload success. + */ +export function trainingTerrainFromCompiledScene( + geometries: readonly CompiledMapGeometry[], + initialPose: readonly number[], + mapExtent: number, + coordinates: TrainingSceneCoordinates = {}, +): TrainingTerrain { + const obstacles: TrainingTerrain['boxes'] = []; + let friction: number | undefined; + let extent = Math.max(4, mapExtent); + for (const g of geometries) { + if (!g.static || !g.collision) continue; + if (![...g.position, ...g.rotation, ...g.size, ...g.friction].every(Number.isFinite)) + throw new Error(`静态碰撞几何 ${g.name} 含非有限数值`); + const supportName = /(?:^|[_-])(?:floor|ground|flat)(?:[_-]|$)/i.test(g.name); + if (!g.map && !(g.type === 0 && supportName)) + throw new Error(`发现未声明为已应用地图的静态碰撞几何 ${g.name},不能静默漏障碍`); + if ( + g.friction.length !== 3 || + Math.abs(g.friction[1] - 0.005) > 1e-9 || + Math.abs(g.friction[2] - 0.0001) > 1e-9 || + (friction !== undefined && Math.abs(g.friction[0] - friction) > 1e-9) + ) + throw new Error( + '地图混合摩擦无法用单一friction表示;请先显式统一各碰撞几何摩擦为 [f,0.005,0.0001]', + ); + friction = g.friction[0]; + const horizontal = Math.abs(Math.abs(g.rotation[8]) - 1) < 1e-9; + if (g.type === 0) { + if (!supportName || !horizontal || g.rotation[8] < 0 || Math.abs(g.position[2]) > 1e-6) + throw new Error('仅支持明确命名floor/ground的水平z=0支撑plane;倾斜/非零高度plane无法导出'); + continue; + } + const half = collisionHalfBounds(g.type, g.size, g.rotation); + if (half.some((v) => !Number.isFinite(v) || v <= 0)) + throw new Error('碰撞半尺寸必须严格大于零'); + extent = Math.max(extent, ...[0, 1].map((i) => Math.abs(g.position[i]) + half[i])); + const top = g.position[2] + half[2], + bottom = g.position[2] - half[2]; + if (g.type === 6 && supportName && horizontal && Math.abs(top) <= 1e-6 && bottom < 0) continue; + if (bottom < -1e-6) + throw new Error(`几何 ${g.name} 含地下/坑底结构;标准floor会填平,拒绝同步`); + obstacles.push({ pos: [...g.position], size: half, yaw: 0 }); + if (obstacles.length > 256) throw new Error('自定义障碍物超过256,不能截断'); + } + if (friction === undefined) throw new Error('已应用地图没有可导出的静态碰撞几何'); + if (extent > 12 + 1e-6) throw new Error('世界地图范围超过24m上限;不平移或裁剪几何'); + const size = Math.max(8, extent * 2); + const spawn = coordinates.spawn ?? [initialPose[0], initialPose[1]]; + const q = initialPose.slice(3, 7); + const yaw = Math.atan2(2 * (q[0] * q[3] + q[1] * q[2]), 1 - 2 * (q[2] ** 2 + q[3] ** 2)); + const target = coordinates.target ?? [spawn[0] + 3 * Math.cos(yaw), spawn[1] + 3 * Math.sin(yaw)]; + return { + representation: 'boxes-v1', + approximation: true, + size, + friction, + boxes: [{ pos: [0, 0, -0.1], size: [size / 2, size / 2, 0.1], yaw: 0 }, ...obstacles], + spawn: [spawn[0], spawn[1], 0.32], + spawnQuaternion: q, + target: [...target], + actualObstacleCount: obstacles.length, + }; +} diff --git a/web_platform/src/rl/RLPolicyPanel.test.tsx b/web_platform/src/rl/RLPolicyPanel.test.tsx new file mode 100644 index 00000000..31106446 --- /dev/null +++ b/web_platform/src/rl/RLPolicyPanel.test.tsx @@ -0,0 +1,64 @@ +import { fireEvent, render, screen } from '@testing-library/react'; +import { RLPolicyPanel } from './RLPolicyPanel'; +import type { RLPolicyStatus } from './types'; + +it('仅避障策略显示实时导航状态,支持暂停时设定/复位', () => { + const status: RLPolicyStatus = { + taskId: 'Unitree-Go2-ObstacleAvoidance', + taskName: '避障', + path: 'policy.onnx', + loaded: true, + enabled: false, + controlHz: 50, + observationSize: 81, + actionSize: 12, + inputName: 'in', + outputName: 'out', + command: { linearX: 0, linearY: 0, angularZ: 0 }, + inferenceCount: 0, + lastInferenceMs: 0, + navigation: { target: [2, -1], defaultTarget: [5, 0], distance: 3.25, targetHeight: 0 }, + }; + const props = { + paths: [], + loading: false, + status, + onSelectPath: vi.fn(), + onLoadPath: vi.fn(), + onImport: vi.fn(), + onToggle: vi.fn(), + onCommand: vi.fn(), + onRemove: vi.fn(), + onNavigationTargetMode: vi.fn(), + onResetNavigationTarget: vi.fn(), + }; + const view = render(); + expect(screen.getByText('(2.00, -1.00)')).toBeInTheDocument(); + expect(screen.getByText('3.25 m')).toBeInTheDocument(); + fireEvent.click(screen.getByRole('button', { name: /^设定目标$/ })); + expect(props.onNavigationTargetMode).toHaveBeenCalledWith(true); + fireEvent.click(screen.getByRole('button', { name: '复位目标点' })); + expect(props.onResetNavigationTarget).toHaveBeenCalledOnce(); + expect(props.onToggle).not.toHaveBeenCalled(); + expect(screen.getByText(/策略未启用:设定目标不会自动启动/)).toBeVisible(); + view.rerender(); + expect(screen.getByText(/仿真已暂停/)).toBeVisible(); + view.rerender(); + expect(screen.getByText(/导航安全停止:请重置并重新启用/)).toBeVisible(); + view.rerender( + , + ); + expect(screen.getByText('1.10 m')).toBeInTheDocument(); + expect(screen.getByText(/Esc 取消/)).toBeInTheDocument(); + view.rerender( + , + ); + expect(screen.queryByText('复位目标点')).not.toBeInTheDocument(); +}); diff --git a/web_platform/src/rl/RLPolicyPanel.tsx b/web_platform/src/rl/RLPolicyPanel.tsx index 1f05a697..1077be7f 100644 --- a/web_platform/src/rl/RLPolicyPanel.tsx +++ b/web_platform/src/rl/RLPolicyPanel.tsx @@ -1,6 +1,7 @@ import { useRef, type ChangeEvent } from 'react'; import { BrainCircuit, FileUp, Power, RotateCw, Trash2 } from 'lucide-react'; import type { RLCommand, RLPolicyStatus } from './types'; +import { useAppStore } from '../stores/useAppStore'; import { Badge, Button, PropertyRow, Select } from '../components/ui'; export interface RLPolicyPanelProps { @@ -14,6 +15,9 @@ export interface RLPolicyPanelProps { onToggle(enabled: boolean): void; onCommand(command: RLCommand): void; onRemove(): void; + navigationTargetMode?: boolean; + onNavigationTargetMode?(active: boolean): void; + onResetNavigationTarget?(): void; } export function RLPolicyPanel({ @@ -27,8 +31,12 @@ export function RLPolicyPanel({ onToggle, onCommand, onRemove, + navigationTargetMode = false, + onNavigationTargetMode, + onResetNavigationTarget, }: RLPolicyPanelProps) { const input = useRef(null); + const paused = useAppStore((state) => state.paused); const importFile = (event: ChangeEvent) => { const file = event.target.files?.[0]; if (file) onImport(file); @@ -98,36 +106,77 @@ export function RLPolicyPanel({ /> -
-

速度指令(机身坐标系)

- onCommand({ ...command, linearX })} - /> - onCommand({ ...command, linearY })} - /> - onCommand({ ...command, angularZ })} - /> - -
+ {status.taskId === 'Unitree-Go2-ObstacleAvoidance' && status.navigation && ( +
+ + +

+ {status.error + ? '导航安全停止:请重置并重新启用策略。' + : !status.enabled + ? '策略未启用:设定目标不会自动启动,请手动启用。' + : paused + ? '仿真已暂停:目标已设定,请恢复仿真后运动。' + : '导航控制运行中(不保证到达);可持续更换目标。'} +

+
+ + +
+ {navigationTargetMode && ( +

+ 点击主视口地形设定目标,Esc 取消;不会选择或移动地图对象。 +

+ )} +
+ )} + {status.observationSize === 81 || status.observationSize === 97 ? ( +

+ 自动导航到配套目标,到达后停止指令。评测20秒/跌倒/越界时停止(非训练端自动重置);请重置后启用。水平射线有矮障碍/跌落盲区,Go2-W不是Go2同构模型。 +

+ ) : ( +
+

速度指令(机身坐标系)

+ onCommand({ ...command, linearX })} + /> + onCommand({ ...command, linearY })} + /> + onCommand({ ...command, angularZ })} + /> + +
+ )} {status.error && (

Array.from(new TextEncoder().encode(s)); +export function metadataModel(value: unknown): Uint8Array { + return new Uint8Array( + field(14, [ + ...field(1, string('platform_deployment')), + ...field(2, string(JSON.stringify(value))), + ]), + ); +} +describe('policy deployment metadata', () => { + it('解析ONNX ModelProto的部署元数据且忽略graph,不依赖ORT', () => { + const bytes = metadataModel(fixture); + expect(readPolicyDeployment(bytes)).toEqual(validatePolicyDeployment(fixture)); + expect(readPolicyDeployment(new Uint8Array([8, 9]))).toBeUndefined(); + expect(() => readPolicyDeployment(bytes.slice(0, -1))).toThrow(/截断|长度/); + expect(() => readPolicyDeployment(new Uint8Array([...bytes, ...bytes]))).toThrow(/重复/); + }); + it.each([ + ['version', 2], + ['taskId', 'Unitree-Go2-Rough'], + ['browserCompatible', false], + ['observationSize', 47], + ['actionSize', 16], + ['controlHz', 200], + ['seed', NaN], + ['jointNames', []], + ['effortLimits', []], + ['observationTerms', []], + ])('拒绝不兼容%s', (key, value) => + expect(() => validatePolicyDeployment({ ...fixture, [key]: value })).toThrow(), + ); + it('拒绝不支持sensor、超限box、不有限数字及出生点偏移', () => { + expect(() => + validatePolicyDeployment({ ...fixture, sensorCfg: { ...fixture.sensorCfg, rayCount: 64 } }), + ).toThrow(); + expect(() => + validatePolicyDeployment({ + ...fixture, + sensorCfg: { ...fixture.sensorCfg, type: 'camera_depth' }, + }), + ).toThrow(); + expect(() => + validatePolicyDeployment({ + ...fixture, + terrain: { ...fixture.terrain, boxes: Array(258).fill(fixture.terrain.boxes[0]) }, + }), + ).toThrow(); + expect(() => + validatePolicyDeployment({ ...fixture, terrain: { ...fixture.terrain, friction: Infinity } }), + ).toThrow(); + expect(() => + validatePolicyDeployment({ + ...fixture, + terrain: { ...fixture.terrain, spawn: [0, 0, 0.32] }, + }), + ).toThrow(); + }); +}); + +it('新server默认Flat允许旧无metadata导出,仅返回需要真实graph47→12校验的预期契约', () => { + const { terrain, terrainPreset, terrainParams, sensorCfg, navigation, ...base } = + validatePolicyDeployment(fixture); + const flat = validatePolicyDeployment({ + ...base, + taskId: 'Unitree-Go2-Flat', + observationSize: 47, + observationTerms: base.observationTerms.slice(0, 7), + }); + const legacy = new Uint8Array([8, 9]); + expect(resolvePolicyDeployment(legacy, flat)).toEqual(flat); + expect(() => resolvePolicyDeployment(legacy, validatePolicyDeployment(fixture))).toThrow( + /不一致/, + ); + expect(() => + resolvePolicyDeployment(legacy, { ...flat, terrain, terrainPreset, terrainParams }), + ).toThrow(/不一致/); + expect(() => resolvePolicyDeployment(legacy, { ...flat, sensorCfg })).toThrow(/不一致/); + expect(() => resolvePolicyDeployment(legacy, { ...flat, navigation })).toThrow(/不一致/); +}); + +it('custom_boxes ONNX metadata保持完整布局且绝不降级旧Flat或preset', async () => { + const { default: terrain } = + await import('../../../training_server/tests/fixtures/custom-boxes.json'); + const custom = validatePolicyDeployment({ + ...fixture, + terrainPreset: 'custom_boxes', + terrain, + terrainParams: { size: terrain.size, friction: terrain.friction }, + }); + const bytes = metadataModel(custom); // ModelProto metadata fixture, not a trained graph. + expect(readPolicyDeployment(bytes)).toEqual(custom); + expect(resolvePolicyDeployment(bytes, custom)).toEqual(custom); + expect(() => resolvePolicyDeployment(new Uint8Array([8, 9]), custom)).toThrow(/不一致/); + expect(() => resolvePolicyDeployment(metadataModel(fixture), custom)).toThrow(/不一致/); +}); + +it('调参导航速度有限且在专属范围,旧0.6契约仍有效', () => { + for (const speed of [0.3, 0.6, 0.9, 1.2]) { + const value = structuredClone(fixture); + value.navigation.speed = speed; + expect(readPolicyDeployment(metadataModel(value))?.navigation?.speed).toBe(speed); + } + for (const speed of [0, 1.21, NaN, Infinity]) { + const value = structuredClone(fixture); + value.navigation.speed = speed; + expect(() => validatePolicyDeployment(value)).toThrow(); + } +}); diff --git a/web_platform/src/rl/deployment.ts b/web_platform/src/rl/deployment.ts new file mode 100644 index 00000000..6b27edfd --- /dev/null +++ b/web_platform/src/rl/deployment.ts @@ -0,0 +1,406 @@ +import { GO2W_VELOCITY_TASK } from './tasks/go2wVelocity'; + +export const OBSTACLE_TASK_ID = 'Unitree-Go2-ObstacleAvoidance'; +export interface TrainingTerrain { + representation: 'boxes-v1'; + approximation: boolean; + size: number; + friction: number; + boxes: { pos: number[]; size: number[]; yaw: number }[]; + spawn: number[]; + spawnQuaternion: number[]; + target: number[]; + actualObstacleCount: number; +} +export interface ObstacleSensorConfig { + type: 'raycast'; + fov: number; + maxDistance: number; + safetyDistance: number; + avoidanceWeight: number; + rayCount: 32 | 48; + sensorMode?: 'single_ring_raycast' | 'multi_ring_raycast'; + pitchAngles?: number[]; + yawCount?: number; + yawAngles?: number[]; + angleUnit?: 'deg'; + rayOrder?: 'layer-major'; + offset: number[]; + alignment: 'base'; + terrainOnly: true; + includeGround: true; +} +export interface PolicyDeployment { + version: 1; + taskId: string; + browserCompatible: true; + observationSize: number; + actionSize: 12; + controlHz: 50; + gaitPeriod: number; + jointNames: string[]; + defaultJointPosition: number[]; + actionScale: number[]; + stiffness: number[]; + damping: number[]; + effortLimits: number[]; + observationTerms: string[]; + seed: number; + terrainPreset?: string; + terrainParams?: Record; + terrain?: TrainingTerrain; + sensorCfg?: ObstacleSensorConfig; + navigation?: { + speed: number; + arrivalRadius: number; + distanceScale: number; + headingScale: number; + yawGain: number; + maxYawRate: number; + episodeSeconds: number; + onArrival: 'stop'; + onReset: 'respawn'; + }; +} +function requireValue(condition: unknown, label: string): asserts condition { + if (!condition) throw new Error(`不支持或无效的策略部署配置:${label}`); +} +function record(value: unknown): Record { + requireValue(value !== null && typeof value === 'object' && !Array.isArray(value), '对象'); + return value as Record; +} +function number(value: unknown, min: number, max: number): asserts value is number { + requireValue( + typeof value === 'number' && Number.isFinite(value) && value >= min && value <= max, + '有限数值范围', + ); +} +function vector(value: unknown, length: number, min = -100, max = 100): asserts value is number[] { + requireValue(Array.isArray(value) && value.length === length, `向量长度 ${length}`); + value.forEach((v) => number(v, min, max)); +} +function canonical(value: unknown): string { + if (Array.isArray(value)) return `[${value.map(canonical).join(',')}]`; + if (value !== null && typeof value === 'object') + return `{${Object.entries(value) + .sort(([a], [b]) => a.localeCompare(b)) + .map(([key, v]) => `${JSON.stringify(key)}:${canonical(v)}`) + .join(',')}}`; + return JSON.stringify(value); +} +export function policyDeploymentsMatch(a: PolicyDeployment, b: PolicyDeployment): boolean { + return canonical(validatePolicyDeployment(a)) === canonical(validatePolicyDeployment(b)); +} +function equal(actual: unknown, expected: unknown, label: string) { + requireValue(canonical(actual) === canonical(expected), label); +} +/** Strict custom layout boundary shared by scene export, upload and ONNX metadata. */ +export function validateCustomTerrain(value: unknown): TrainingTerrain { + const t = record(value); + const keys = (v: Record, expected: string[]) => + equal(Object.keys(v).sort(), expected.sort(), '自定义布局字段(不接受路径/MJCF)'); + keys(t, [ + 'representation', + 'approximation', + 'size', + 'friction', + 'boxes', + 'spawn', + 'spawnQuaternion', + 'target', + 'actualObstacleCount', + ]); + equal(t.representation, 'boxes-v1', '地图表示'); + equal(t.approximation, true, 'custom_boxes 必须标记AABB/底板标准化近似'); + number(t.size, 8, 24); + number(t.friction, 0.2, 2); + requireValue( + Array.isArray(t.boxes) && t.boxes.length >= 1 && t.boxes.length <= 257, + '最多256障碍物', + ); + for (const item of t.boxes) { + const b = record(item); + keys(b, ['pos', 'size', 'yaw']); + vector(b.pos, 3, -12, 12); + vector(b.size, 3, 0, 12); + requireValue( + b.size.every((v) => v > 0), + '半尺寸必须严格大于零', + ); + equal(b.yaw, 0, 'box yaw'); + for (let i = 0; i < 2; i++) + requireValue(Math.abs(b.pos[i]) + b.size[i] <= t.size / 2 + 1e-6, '世界地图边界'); + requireValue(b.pos[2] - b.size[2] >= -0.2 - 1e-6 && b.pos[2] + b.size[2] <= 12, 'box高度边界'); + } + equal( + t.boxes[0], + { pos: [0, 0, -0.1], size: [t.size / 2, t.size / 2, 0.1], yaw: 0 }, + '标准floor', + ); + equal(t.actualObstacleCount, t.boxes.length - 1, '实际障碍数'); + vector(t.spawn, 3, -12, 12); + equal(t.spawn[2], 0.32, '出生高度'); + vector(t.target, 2, -12, 12); + vector(t.spawnQuaternion, 4, -1, 1); + requireValue( + Math.abs(t.spawnQuaternion.reduce((sum, v) => sum + v * v, 0) - 1) <= 1e-6, + '归一化出生四元数', + ); + for (const point of [t.spawn, t.target]) { + requireValue( + point.slice(0, 2).every((v) => Math.abs(v) <= (t.size as number) / 2 - 0.5), + '起终点0.5m安全区越界', + ); + for (const item of t.boxes.slice(1)) { + const b = item as TrainingTerrain['boxes'][number]; + const distanceSquared = [0, 1].reduce( + (sum, i) => sum + Math.max(Math.abs(point[i] - b.pos[i]) - b.size[i], 0) ** 2, + 0, + ); + requireValue( + distanceSquared > 0.25, + '障碍物侵占起终点0.5m圆形安全区;请修改坐标,不会清除障碍', + ); + } + } + return structuredClone(t) as unknown as TrainingTerrain; +} + +/** A single spelling for the supported angular patterns; legacy metadata defaults to 32. */ +export function obstacleSensorPattern(mode: unknown, fov: unknown) { + requireValue( + mode === undefined || mode === 'single_ring_raycast' || mode === 'multi_ring_raycast', + 'sensorMode', + ); + number(fov, 30, 120); + const multi = mode === 'multi_ring_raycast'; + const yawCount = multi ? 16 : 32; + return { + sensorMode: multi ? ('multi_ring_raycast' as const) : ('single_ring_raycast' as const), + rayCount: multi ? (48 as const) : (32 as const), + pitchAngles: multi ? [0, -20, -45] : [0], + yawCount, + yawAngles: Array.from({ length: yawCount }, (_, i) => -fov / 2 + (i * fov) / (yawCount - 1)), + angleUnit: 'deg' as const, + rayOrder: 'layer-major' as const, + }; +} + +/** v1 only: reject incompatible contracts before altering the simulation or allocating ORT. */ +export function validatePolicyDeployment(value: unknown): PolicyDeployment { + const d = structuredClone(record(value)); + requireValue(JSON.stringify(d).length <= 100_000, '配置大小'); + requireValue( + d.version === 1 && d.browserCompatible === true, + '版本/浏览器兼容性(旧 Rough 不支持)', + ); + requireValue(d.taskId === OBSTACLE_TASK_ID || d.taskId === 'Unitree-Go2-Flat', '任务'); + if (d.taskId !== OBSTACLE_TASK_ID) equal(d.observationSize, 47, '观测维数'); + equal(d.actionSize, 12, '动作维数'); + equal(d.controlHz, 50, '控制频率'); + equal(d.gaitPeriod, 0.6, '步态周期'); + for (const key of [ + 'jointNames', + 'defaultJointPosition', + 'actionScale', + 'stiffness', + 'damping', + ] as const) + equal(d[key], GO2W_VELOCITY_TASK[key], key); + equal( + d.effortLimits, + [23.5, 23.5, 45, 23.5, 23.5, 45, 23.5, 23.5, 45, 23.5, 23.5, 45], + '力矩限幅', + ); + equal( + d.observationTerms, + [ + 'base_ang_vel', + 'projected_gravity', + 'command', + 'phase', + 'joint_pos', + 'joint_vel', + 'actions', + ...(d.taskId === OBSTACLE_TASK_ID ? ['forward_depth', 'target_error'] : []), + ], + '观测顺序', + ); + number(d.seed, 0, 2147483647); + requireValue(Number.isInteger(d.seed), 'seed'); + if (d.terrain !== undefined) { + if (d.terrainPreset === 'custom_boxes') { + const t = validateCustomTerrain(d.terrain); + equal(d.terrainParams, { size: t.size, friction: t.friction }, 'custom_boxes参数与布局一致'); + } else { + const t = record(d.terrain); + equal(t.representation, 'boxes-v1', '地图表示'); + requireValue(typeof t.approximation === 'boolean', '近似标记'); + number(t.size, 8, 24); + number(t.friction, 0.2, 2); + requireValue( + Array.isArray(t.boxes) && t.boxes.length >= 1 && t.boxes.length <= 257, + 'box数量', + ); + for (const item of t.boxes) { + const box = record(item); + vector(box.pos, 3, -12, 12); + vector(box.size, 3, 0.0001, 12); + equal(box.yaw, 0, 'box yaw'); + for (let i = 0; i < 2; i++) + requireValue(Math.abs(box.pos[i]) + box.size[i] <= t.size / 2 + 1e-6, 'box边界'); + } + equal(t.boxes[0], { pos: [0, 0, -0.1], size: [t.size / 2, t.size / 2, 0.1], yaw: 0 }, '地板'); + equal(t.spawn, [-t.size / 2 + 1, 0, 0.32], '出生点'); + equal(t.spawnQuaternion, [1, 0, 0, 0], '出生姿态'); + equal(t.target, [t.size / 2 - 1, 0], '目标'); + equal(t.actualObstacleCount, t.boxes.length - 1, '实际障碍数'); + requireValue( + ['plane', 'discrete_obstacles', 'rough', 'wave', 'pyramid_stairs'].includes( + String(d.terrainPreset), + ), + '地形预设', + ); + const params = record(d.terrainParams); + for (const value of Object.values(params)) number(value, 0, 100); + equal(params.size, t.size, '地图尺寸'); + equal(params.friction, t.friction, '摩擦'); + } + } else { + requireValue(d.terrainPreset !== 'custom_boxes', '缺失custom_boxes布局'); + } + if (d.taskId === OBSTACLE_TASK_ID) { + requireValue(d.terrain, '避障地图'); + const s = record(d.sensorCfg); + equal(s.type, 'raycast', '传感器'); + const pattern = obstacleSensorPattern(s.sensorMode, s.fov); + const allowed = new Set([ + 'type', + 'fov', + 'maxDistance', + 'safetyDistance', + 'avoidanceWeight', + 'offset', + 'alignment', + 'terrainOnly', + 'includeGround', + ...Object.keys(pattern), + ]); + requireValue( + Object.keys(s).every((key) => allowed.has(key)), + '传感器未知字段', + ); + for (const [key, expected] of Object.entries(pattern)) { + if (s[key] === undefined) continue; + if (key === 'yawAngles') { + vector(s[key], (expected as number[]).length, -60, 60); + requireValue( + (s[key] as number[]).every((v, i) => Math.abs(v - (expected as number[])[i]) <= 1e-10), + 'yawAngles与FOV矛盾', + ); + } else equal(s[key], expected, `sensorCfg.${key}`); + } + equal(s.rayCount, pattern.rayCount, '射线数量'); + equal(d.observationSize, 49 + pattern.rayCount, '观测维数'); + d.sensorCfg = { ...s, ...pattern }; + equal(s.offset, [0.3, 0, 0.05], '射线偏移'); + equal(s.alignment, 'base', '射线姿态'); + equal(s.terrainOnly, true, '仅地形'); + equal(s.includeGround, true, '包含地面'); + number(s.fov, 30, 120); + number(s.maxDistance, 1, 5); + number(s.safetyDistance, 0.1, 1); + number(s.avoidanceWeight, 0, 10); + requireValue(s.safetyDistance < s.maxDistance, '安全距离'); + const n = record(d.navigation); + number(n.speed, 0.3, 1.2); + for (const [key, expected] of Object.entries({ + arrivalRadius: 0.5, + distanceScale: record(d.terrain).size, + headingScale: Math.PI, + yawGain: 1, + maxYawRate: 1, + episodeSeconds: 20, + onArrival: 'stop', + onReset: 'respawn', + })) + equal(n[key], expected, `导航 ${key}`); + } + return structuredClone(d) as unknown as PolicyDeployment; +} + +/** Read only ModelProto.metadata_props (field 14); skip graph/tensors without decoding them. */ +export function readPolicyDeployment(bytes: Uint8Array): PolicyDeployment | undefined { + requireValue(bytes.length <= 64 * 1024 * 1024, 'ONNX超过64MiB'); + const entries: Uint8Array[] = []; + function fields(data: Uint8Array, visit: (field: number, value: Uint8Array) => void) { + let p = 0; + const varint = () => { + let value = 0; + for (let i = 0; i < 10; i++) { + requireValue(p < data.length, 'ONNX截断'); + const b = data[p++]; + value += (b & 127) * 2 ** (i * 7); + requireValue(Number.isSafeInteger(value), 'ONNX整数'); + if (!(b & 128)) return value; + } + throw new Error('无效ONNX varint'); + }; + while (p < data.length) { + const tag = varint(), + field = Math.floor(tag / 8), + wire = tag % 8; + requireValue(field > 0, 'ONNX字段'); + if (wire === 0) { + varint(); + continue; + } + const length = wire === 2 ? varint() : wire === 1 ? 8 : wire === 5 ? 4 : -1; + requireValue(length >= 0 && p + length <= data.length, 'ONNX字段长度'); + if (wire === 2) visit(field, data.subarray(p, p + length)); + p += length; + } + } + fields(bytes, (field, value) => { + if (field === 14) { + requireValue(value.length < 100_000 && entries.length < 100, 'ONNX metadata大小'); + entries.push(value); + } + }); + let deployment: PolicyDeployment | undefined; + const decoder = new TextDecoder('utf-8', { fatal: true }); + for (const entry of entries) { + let key = '', + value = ''; + fields(entry, (field, bytes) => { + if (field === 1) key = decoder.decode(bytes); + if (field === 2) value = decoder.decode(bytes); + }); + if (key === 'platform_deployment') { + requireValue(!deployment, '重复deployment'); + deployment = validatePolicyDeployment(JSON.parse(value)); + } + } + return deployment; +} + +/** Only the original non-custom Flat job may import older external-trainer ONNX without metadata. + * Runtime.load still validates the real graph, not this JSON declaration. */ +export function resolvePolicyDeployment( + bytes: Uint8Array, + expected?: PolicyDeployment, +): PolicyDeployment | undefined { + const embedded = readPolicyDeployment(bytes); + if (!expected) return embedded; + const validated = validatePolicyDeployment(expected); + if (embedded && policyDeploymentsMatch(validated, embedded)) return embedded; + const legacyFlat = + validated.taskId === 'Unitree-Go2-Flat' && + validated.terrain === undefined && + validated.terrainPreset === undefined && + validated.terrainParams === undefined && + validated.sensorCfg === undefined && + validated.navigation === undefined; + if (!embedded && legacyFlat) return validated; + throw new Error('下载的策略与训练作业部署配置不一致'); +} diff --git a/web_platform/src/rl/fixtures/go2CompiledDynamics.json b/web_platform/src/rl/fixtures/go2CompiledDynamics.json new file mode 100644 index 00000000..e9b56a2f --- /dev/null +++ b/web_platform/src/rl/fixtures/go2CompiledDynamics.json @@ -0,0 +1,279 @@ +{ + "source": "Entity(get_go2_robot_cfg()).compile(); go2_constants.py FULL_COLLISION and armature. Foot friction is overridden by custom terrain startup.", + "joints": { + "floating_base_joint": { + "dof_armature": 0.0, + "dof_damping": 0.0, + "dof_frictionloss": 0.0 + }, + "FL_hip_joint": { + "dof_armature": 0.01, + "dof_damping": 0.0, + "dof_frictionloss": 0.0 + }, + "FL_thigh_joint": { + "dof_armature": 0.01, + "dof_damping": 0.0, + "dof_frictionloss": 0.0 + }, + "FL_calf_joint": { + "dof_armature": 0.02, + "dof_damping": 0.0, + "dof_frictionloss": 0.0 + }, + "FR_hip_joint": { + "dof_armature": 0.01, + "dof_damping": 0.0, + "dof_frictionloss": 0.0 + }, + "FR_thigh_joint": { + "dof_armature": 0.01, + "dof_damping": 0.0, + "dof_frictionloss": 0.0 + }, + "FR_calf_joint": { + "dof_armature": 0.02, + "dof_damping": 0.0, + "dof_frictionloss": 0.0 + }, + "RL_hip_joint": { + "dof_armature": 0.01, + "dof_damping": 0.0, + "dof_frictionloss": 0.0 + }, + "RL_thigh_joint": { + "dof_armature": 0.01, + "dof_damping": 0.0, + "dof_frictionloss": 0.0 + }, + "RL_calf_joint": { + "dof_armature": 0.02, + "dof_damping": 0.0, + "dof_frictionloss": 0.0 + }, + "RR_hip_joint": { + "dof_armature": 0.01, + "dof_damping": 0.0, + "dof_frictionloss": 0.0 + }, + "RR_thigh_joint": { + "dof_armature": 0.01, + "dof_damping": 0.0, + "dof_frictionloss": 0.0 + }, + "RR_calf_joint": { + "dof_armature": 0.02, + "dof_damping": 0.0, + "dof_frictionloss": 0.0 + } + }, + "geoms": { + "base1_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "base2_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "base3_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "FL_hip_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "FL_thigh_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "FL_calf1_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "FL_calf2_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "FL_foot_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 3, + "geom_priority": 1, + "geom_solimp": [0.9, 0.95, 0.023, 0.5, 2.0], + "geom_friction": [0.6, 0.005, 0.0001], + "geom_group": 3 + }, + "FR_hip_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "FR_thigh_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "FR_calf1_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "FR_calf2_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "FR_foot_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 3, + "geom_priority": 1, + "geom_solimp": [0.9, 0.95, 0.023, 0.5, 2.0], + "geom_friction": [0.6, 0.005, 0.0001], + "geom_group": 3 + }, + "RL_hip_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "RL_thigh_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "RL_calf1_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "RL_calf2_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "RL_foot_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 3, + "geom_priority": 1, + "geom_solimp": [0.9, 0.95, 0.023, 0.5, 2.0], + "geom_friction": [0.6, 0.005, 0.0001], + "geom_group": 3 + }, + "RR_hip_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "RR_thigh_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "RR_calf1_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "RR_calf2_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 1, + "geom_priority": 0, + "geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0], + "geom_friction": [1.0, 0.005, 0.0001], + "geom_group": 3 + }, + "RR_foot_collision": { + "geom_contype": 1, + "geom_conaffinity": 0, + "geom_condim": 3, + "geom_priority": 1, + "geom_solimp": [0.9, 0.95, 0.023, 0.5, 2.0], + "geom_friction": [0.6, 0.005, 0.0001], + "geom_group": 3 + } + } +} diff --git a/web_platform/src/rl/fixtures/multiRingDeployment.json b/web_platform/src/rl/fixtures/multiRingDeployment.json new file mode 100644 index 00000000..7c3afbe9 --- /dev/null +++ b/web_platform/src/rl/fixtures/multiRingDeployment.json @@ -0,0 +1,221 @@ +{ + "version": 1, + "taskId": "Unitree-Go2-ObstacleAvoidance", + "browserCompatible": true, + "observationSize": 97, + "actionSize": 12, + "controlHz": 50, + "gaitPeriod": 0.6, + "jointNames": [ + "FL_hip_joint", + "FL_thigh_joint", + "FL_calf_joint", + "FR_hip_joint", + "FR_thigh_joint", + "FR_calf_joint", + "RL_hip_joint", + "RL_thigh_joint", + "RL_calf_joint", + "RR_hip_joint", + "RR_thigh_joint", + "RR_calf_joint" + ], + "defaultJointPosition": [-0.1, 0.9, -1.8, 0.1, 0.9, -1.8, -0.1, 0.9, -1.8, 0.1, 0.9, -1.8], + "actionScale": [0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25], + "stiffness": [20, 20, 40, 20, 20, 40, 20, 20, 40, 20, 20, 40], + "damping": [1, 1, 2, 1, 1, 2, 1, 1, 2, 1, 1, 2], + "effortLimits": [23.5, 23.5, 45, 23.5, 23.5, 45, 23.5, 23.5, 45, 23.5, 23.5, 45], + "observationTerms": [ + "base_ang_vel", + "projected_gravity", + "command", + "phase", + "joint_pos", + "joint_vel", + "actions", + "forward_depth", + "target_error" + ], + "seed": 42, + "terrainPreset": "discrete_obstacles", + "terrainParams": { + "size": 12, + "obstacle_count": 24, + "obstacle_height_min": 0.2, + "obstacle_height_max": 0.6, + "spacing": 1.2, + "friction": 0.8, + "roughness": 0.06, + "step_height": 0.08, + "wave_amplitude": 0.08 + }, + "terrain": { + "representation": "boxes-v1", + "approximation": false, + "size": 12, + "friction": 0.8, + "boxes": [ + { + "pos": [0, 0, -0.1], + "size": [6.0, 6.0, 0.1], + "yaw": 0 + }, + { + "pos": [0.666666666666667, 1.875, 0.261425654654876], + "size": [0.2, 0.2, 0.261425654654876], + "yaw": 0 + }, + { + "pos": [-0.6666666666666665, 4.375, 0.24594635733876358], + "size": [0.2, 0.2, 0.24594635733876358], + "yaw": 0 + }, + { + "pos": [0.666666666666667, -4.375, 0.20724561829094013], + "size": [0.2, 0.2, 0.20724561829094013], + "yaw": 0 + }, + { + "pos": [-2.0, -1.875, 0.29462315279587414], + "size": [0.2, 0.2, 0.29462315279587414], + "yaw": 0 + }, + { + "pos": [-0.6666666666666665, 1.875, 0.17570687544167068], + "size": [0.2, 0.2, 0.17570687544167068], + "yaw": 0 + }, + { + "pos": [-2.0, -3.125, 0.2104081262546454], + "size": [0.2, 0.2, 0.2104081262546454], + "yaw": 0 + }, + { + "pos": [0.666666666666667, 3.125, 0.265880932850599], + "size": [0.2, 0.2, 0.265880932850599], + "yaw": 0 + }, + { + "pos": [2.0, 4.375, 0.2237039504728492], + "size": [0.2, 0.2, 0.2237039504728492], + "yaw": 0 + }, + { + "pos": [-2.0, 0.625, 0.27234138006215547], + "size": [0.2, 0.2, 0.27234138006215547], + "yaw": 0 + }, + { + "pos": [-3.3333333333333335, -0.625, 0.21547042905135239], + "size": [0.2, 0.2, 0.21547042905135239], + "yaw": 0 + }, + { + "pos": [2.0, -0.625, 0.24091436724298468], + "size": [0.2, 0.2, 0.24091436724298468], + "yaw": 0 + }, + { + "pos": [-3.3333333333333335, 0.625, 0.10916487673113245], + "size": [0.2, 0.2, 0.10916487673113245], + "yaw": 0 + }, + { + "pos": [3.333333333333333, -0.625, 0.14557965513030938], + "size": [0.2, 0.2, 0.14557965513030938], + "yaw": 0 + }, + { + "pos": [0.666666666666667, -1.875, 0.15787759272042143], + "size": [0.2, 0.2, 0.15787759272042143], + "yaw": 0 + }, + { + "pos": [-2.0, -0.625, 0.11595839538472551], + "size": [0.2, 0.2, 0.11595839538472551], + "yaw": 0 + }, + { + "pos": [2.0, -3.125, 0.14655817727220605], + "size": [0.2, 0.2, 0.14655817727220605], + "yaw": 0 + }, + { + "pos": [0.666666666666667, 0.625, 0.12020028588194583], + "size": [0.2, 0.2, 0.12020028588194583], + "yaw": 0 + }, + { + "pos": [3.333333333333333, -3.125, 0.15559472062201843], + "size": [0.2, 0.2, 0.15559472062201843], + "yaw": 0 + }, + { + "pos": [-0.6666666666666665, -1.875, 0.22713688885288003], + "size": [0.2, 0.2, 0.22713688885288003], + "yaw": 0 + }, + { + "pos": [0.666666666666667, 4.375, 0.17296643579401685], + "size": [0.2, 0.2, 0.17296643579401685], + "yaw": 0 + }, + { + "pos": [3.333333333333333, 3.125, 0.17403619342337653], + "size": [0.2, 0.2, 0.17403619342337653], + "yaw": 0 + }, + { + "pos": [2.0, -4.375, 0.14190140615429753], + "size": [0.2, 0.2, 0.14190140615429753], + "yaw": 0 + }, + { + "pos": [2.0, 0.625, 0.15339556440982266], + "size": [0.2, 0.2, 0.15339556440982266], + "yaw": 0 + }, + { + "pos": [0.666666666666667, -3.125, 0.2873309175424988], + "size": [0.2, 0.2, 0.2873309175424988], + "yaw": 0 + } + ], + "spawn": [-5.0, 0, 0.32], + "spawnQuaternion": [1, 0, 0, 0], + "target": [5.0, 0], + "actualObstacleCount": 24 + }, + "sensorCfg": { + "fov": 90, + "maxDistance": 4, + "safetyDistance": 0.5, + "avoidanceWeight": 2, + "type": "raycast", + "sensorMode": "multi_ring_raycast", + "rayCount": 48, + "pitchAngles": [0, -20, -45], + "yawCount": 16, + "yawAngles": [ + -45.0, -39.0, -33.0, -27.0, -21.0, -15.0, -9.0, -3.0, 3.0, 9.0, 15.0, 21.0, 27.0, 33.0, 39.0, + 45.0 + ], + "angleUnit": "deg", + "rayOrder": "layer-major", + "offset": [0.3, 0, 0.05], + "alignment": "base", + "terrainOnly": true, + "includeGround": true + }, + "navigation": { + "speed": 0.6, + "arrivalRadius": 0.5, + "distanceScale": 12, + "headingScale": 3.141592653589793, + "yawGain": 1.0, + "maxYawRate": 1.0, + "episodeSeconds": 20, + "onArrival": "stop", + "onReset": "respawn" + } +} diff --git a/web_platform/src/rl/fixtures/multiRingGolden.json b/web_platform/src/rl/fixtures/multiRingGolden.json new file mode 100644 index 00000000..4fc6c467 --- /dev/null +++ b/web_platform/src/rl/fixtures/multiRingGolden.json @@ -0,0 +1,336 @@ +{ + "source": "CPU MuJoCo 3.5.0 mj_ray", + "layouts": { + "default": [ + { + "pos": [0, 0, -0.1], + "size": [6.0, 6.0, 0.1], + "yaw": 0 + }, + { + "pos": [0.666666666666667, 1.875, 0.261425654654876], + "size": [0.2, 0.2, 0.261425654654876], + "yaw": 0 + }, + { + "pos": [-0.6666666666666665, 4.375, 0.24594635733876358], + "size": [0.2, 0.2, 0.24594635733876358], + "yaw": 0 + }, + { + "pos": [0.666666666666667, -4.375, 0.20724561829094013], + "size": [0.2, 0.2, 0.20724561829094013], + "yaw": 0 + }, + { + "pos": [-2.0, -1.875, 0.29462315279587414], + "size": [0.2, 0.2, 0.29462315279587414], + "yaw": 0 + }, + { + "pos": [-0.6666666666666665, 1.875, 0.17570687544167068], + "size": [0.2, 0.2, 0.17570687544167068], + "yaw": 0 + }, + { + "pos": [-2.0, -3.125, 0.2104081262546454], + "size": [0.2, 0.2, 0.2104081262546454], + "yaw": 0 + }, + { + "pos": [0.666666666666667, 3.125, 0.265880932850599], + "size": [0.2, 0.2, 0.265880932850599], + "yaw": 0 + }, + { + "pos": [2.0, 4.375, 0.2237039504728492], + "size": [0.2, 0.2, 0.2237039504728492], + "yaw": 0 + }, + { + "pos": [-2.0, 0.625, 0.27234138006215547], + "size": [0.2, 0.2, 0.27234138006215547], + "yaw": 0 + }, + { + "pos": [-3.3333333333333335, -0.625, 0.21547042905135239], + "size": [0.2, 0.2, 0.21547042905135239], + "yaw": 0 + }, + { + "pos": [2.0, -0.625, 0.24091436724298468], + "size": [0.2, 0.2, 0.24091436724298468], + "yaw": 0 + }, + { + "pos": [-3.3333333333333335, 0.625, 0.10916487673113245], + "size": [0.2, 0.2, 0.10916487673113245], + "yaw": 0 + }, + { + "pos": [3.333333333333333, -0.625, 0.14557965513030938], + "size": [0.2, 0.2, 0.14557965513030938], + "yaw": 0 + }, + { + "pos": [0.666666666666667, -1.875, 0.15787759272042143], + "size": [0.2, 0.2, 0.15787759272042143], + "yaw": 0 + }, + { + "pos": [-2.0, -0.625, 0.11595839538472551], + "size": [0.2, 0.2, 0.11595839538472551], + "yaw": 0 + }, + { + "pos": [2.0, -3.125, 0.14655817727220605], + "size": [0.2, 0.2, 0.14655817727220605], + "yaw": 0 + }, + { + "pos": [0.666666666666667, 0.625, 0.12020028588194583], + "size": [0.2, 0.2, 0.12020028588194583], + "yaw": 0 + }, + { + "pos": [3.333333333333333, -3.125, 0.15559472062201843], + "size": [0.2, 0.2, 0.15559472062201843], + "yaw": 0 + }, + { + "pos": [-0.6666666666666665, -1.875, 0.22713688885288003], + "size": [0.2, 0.2, 0.22713688885288003], + "yaw": 0 + }, + { + "pos": [0.666666666666667, 4.375, 0.17296643579401685], + "size": [0.2, 0.2, 0.17296643579401685], + "yaw": 0 + }, + { + "pos": [3.333333333333333, 3.125, 0.17403619342337653], + "size": [0.2, 0.2, 0.17403619342337653], + "yaw": 0 + }, + { + "pos": [2.0, -4.375, 0.14190140615429753], + "size": [0.2, 0.2, 0.14190140615429753], + "yaw": 0 + }, + { + "pos": [2.0, 0.625, 0.15339556440982266], + "size": [0.2, 0.2, 0.15339556440982266], + "yaw": 0 + }, + { + "pos": [0.666666666666667, -3.125, 0.2873309175424988], + "size": [0.2, 0.2, 0.2873309175424988], + "yaw": 0 + } + ], + "low": [ + { + "pos": [0, 0, -0.1], + "size": [6.0, 6.0, 0.1], + "yaw": 0 + }, + { + "pos": [0.9, 0, 0.025], + "size": [0.2, 0.5, 0.025], + "yaw": 0 + } + ], + "edge": [ + { + "pos": [0, 0, -0.1], + "size": [6.0, 6.0, 0.1], + "yaw": 0 + } + ] + }, + "cases": [ + { + "layout": "default", + "position": [-5, 0, 0.32], + "quaternion": [1.0, 0.0, 0.0, 0.0], + "distances": [ + -1.0, 3.2168989147329183, 1.3910905083086054, 1.309380610573421, 1.2496691592432005, -1.0, + -1.0, -1.0, -1.0, 2.716792619137356, 2.5881904510252074, 5.534249133791317, -1.0, + 7.750361403433658, -1.0, 5.904341622907672, 1.0818076280603424, 1.0818076280603424, + 1.0818076280603424, 1.0818076280603424, 1.0818076280603424, 1.0818076280603424, + 1.0818076280603424, 1.0818076280603424, 1.0818076280603424, 1.0818076280603424, + 1.0818076280603424, 1.0818076280603424, 1.0818076280603424, 1.0818076280603424, + 1.0818076280603424, 1.0818076280603424, 0.5232590180780452, 0.5232590180780452, + 0.5232590180780452, 0.5232590180780452, 0.5232590180780452, 0.5232590180780452, + 0.5232590180780452, 0.5232590180780452, 0.5232590180780452, 0.5232590180780452, + 0.5232590180780452, 0.5232590180780452, 0.5232590180780452, 0.5232590180780452, + 0.5232590180780452, 0.5232590180780452 + ], + "hitIds": [ + -1, 4, 10, 10, 10, -1, -1, -1, -1, 9, 9, 1, -1, 8, -1, 2, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0 + ], + "depth": [ + 1, 0.8042247286832296, 0.34777262707715134, 0.32734515264335523, 0.31241728981080014, 1, 1, + 1, 1, 0.679198154784339, 0.6470476127563018, 1, 1, 1, 1, 1, 0.2704519070150856, + 0.2704519070150856, 0.2704519070150856, 0.2704519070150856, 0.2704519070150856, + 0.2704519070150856, 0.2704519070150856, 0.2704519070150856, 0.2704519070150856, + 0.2704519070150856, 0.2704519070150856, 0.2704519070150856, 0.2704519070150856, + 0.2704519070150856, 0.2704519070150856, 0.2704519070150856, 0.1308147545195113, + 0.1308147545195113, 0.1308147545195113, 0.1308147545195113, 0.1308147545195113, + 0.1308147545195113, 0.1308147545195113, 0.1308147545195113, 0.1308147545195113, + 0.1308147545195113, 0.1308147545195113, 0.1308147545195113, 0.1308147545195113, + 0.1308147545195113, 0.1308147545195113, 0.1308147545195113 + ] + }, + { + "layout": "default", + "position": [-3, 1, 0.42], + "quaternion": [ + 0.9542159228792829, 0.07102055199624979, -0.17699854315088892, 0.23043343819907972 + ], + "distances": [ + -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, + -1.0, 5.4074561027806505, 4.629641840762559, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, + -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, 1.1610973299449852, 1.2130908885760254, + 1.2644878853056216, 1.313884872400026, 1.3596758729124838, 1.400132832042811, + 1.4335269836666904, 1.4582821425665438, 1.4731397038952894, 1.4773064832243956, + 1.4705548683639014, 1.4532525276398263, 1.426314634300804, 1.3910898658423474, + 1.349205630860882, 1.3024034756284684 + ], + "hitIds": [ + -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, 0, 23, -1, -1, -1, -1, -1, + -1, -1, -1, -1, -1, -1, -1, -1, -1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0 + ], + "depth": [ + 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, + 1, 0.2902743324862463, 0.30327272214400636, 0.3161219713264054, 0.3284712181000065, + 0.33991896822812095, 0.35003320801070276, 0.3583817459166726, 0.36457053564163594, + 0.36828492597382234, 0.3693266208060989, 0.36763871709097534, 0.36331313190995657, + 0.356578658575201, 0.34777246646058685, 0.3373014077152205, 0.3256008689071171 + ] + }, + { + "layout": "default", + "position": [-2, -2, 0.5], + "quaternion": [ + 0.920495563973526, -0.20509659844134395, 0.07256269404109637, -0.32458890530384493 + ], + "distances": [ + -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, 4.141121232325999, 2.46859904634721, + 2.3854692510759445, 2.821623442591415, 2.353034357115619, 2.037057192180342, + 1.813411398667459, -1.0, -1.0, -1.0, -1.0, 3.2639344664632044, 2.6365745703702133, + 2.2014351861213233, 1.8851396055702525, 1.647197681451036, 1.463506304912354, + 1.3188672797050087, 1.2032560687123115, 1.1098187807096116, 1.0337317580203136, + 0.9715203831532319, 0.9206357398775848, 0.8370986440649361, 1.207568318572572, + 1.1431256747550447, 1.081361321440165, 1.0230591974383043, 0.9686938988212619, + 0.9185003934090735, 0.8725376029719634, 0.8307419413609812, 0.7929698985660093, + 0.7590304425905215, 0.7287087687688034, 0.7017831189893099, 0.6780362858286884, + 0.657263179268657, 0.6392755644084297 + ], + "hitIds": [ + -1, -1, -1, -1, -1, -1, -1, -1, -1, 22, 24, 24, 0, 0, 0, 0, -1, -1, -1, -1, 0, 0, 0, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 6, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0 + ], + "depth": [ + 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 0.6171497615868025, 0.5963673127689861, 0.7054058606478537, + 0.5882585892789047, 0.5092642980450856, 0.45335284966686473, 1, 1, 1, 1, 0.8159836166158011, + 0.6591436425925533, 0.5503587965303308, 0.4712849013925631, 0.411799420362759, + 0.3658765762280885, 0.3297168199262522, 0.3008140171780779, 0.2774546951774029, + 0.2584329395050784, 0.24288009578830796, 0.2301589349693962, 0.20927466101623401, + 0.301892079643143, 0.28578141868876117, 0.27034033036004124, 0.2557647993595761, + 0.24217347470531547, 0.22962509835226838, 0.21813440074299084, 0.2076854853402453, + 0.19824247464150233, 0.1897576106476304, 0.18217719219220085, 0.17544577974732747, + 0.1695090714571721, 0.16431579481716424, 0.15981889110210742 + ] + }, + { + "layout": "low", + "position": [0, 0, 0.32], + "quaternion": [1.0, 0.0, 0.0, 0.0], + "distances": [ + -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, + -1.0, 1.0818076280603424, 1.0818076280603424, 0.9356174080521878, 0.9356174080521878, + 1.0818076280603424, 1.0818076280603424, 1.0818076280603424, 1.0818076280603424, + 1.0818076280603424, 1.0818076280603424, 1.0818076280603424, 1.0818076280603424, + 0.9356174080521878, 0.9356174080521878, 1.0818076280603424, 1.0818076280603424, + 0.5232590180780452, 0.5232590180780452, 0.5232590180780452, 0.5232590180780452, + 0.5232590180780452, 0.5232590180780452, 0.5232590180780452, 0.5232590180780452, + 0.5232590180780452, 0.5232590180780452, 0.5232590180780452, 0.5232590180780452, + 0.5232590180780452, 0.5232590180780452, 0.5232590180780452, 0.5232590180780452 + ], + "hitIds": [ + -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, 0, 0, 1, 1, 0, 0, 0, 0, 0, + 0, 0, 0, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0 + ], + "depth": [ + 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 0.2704519070150856, 0.2704519070150856, + 0.23390435201304696, 0.23390435201304696, 0.2704519070150856, 0.2704519070150856, + 0.2704519070150856, 0.2704519070150856, 0.2704519070150856, 0.2704519070150856, + 0.2704519070150856, 0.2704519070150856, 0.23390435201304696, 0.23390435201304696, + 0.2704519070150856, 0.2704519070150856, 0.1308147545195113, 0.1308147545195113, + 0.1308147545195113, 0.1308147545195113, 0.1308147545195113, 0.1308147545195113, + 0.1308147545195113, 0.1308147545195113, 0.1308147545195113, 0.1308147545195113, + 0.1308147545195113, 0.1308147545195113, 0.1308147545195113, 0.1308147545195113, + 0.1308147545195113, 0.1308147545195113 + ] + }, + { + "layout": "low", + "position": [0, 0, 0.32], + "quaternion": [ + 0.9927681027239722, 0.09709860284347692, 0.054650688238454904, -0.04468397715907572 + ], + "distances": [ + 1.5463415120082626, 1.593685768444766, 1.6628142035077924, 1.7583500416114888, + 1.8874730476566093, 2.06143867460523, 2.298467554987033, 2.6296404516571896, + 3.112149423056716, 3.863359519926606, 5.167240524912454, -1.0, -1.0, -1.0, -1.0, -1.0, + 0.62209731944295, 0.629163205894759, 0.5724502740761562, 0.5541151601654627, + 0.567642029183309, 0.5840261942036087, 0.6035165215332615, 0.626413352751139, + 0.6530744139420617, 0.6839211881333935, 0.7194452948914949, 0.7602140509721573, + 0.8068737942201454, 0.8601486324595742, 1.0831765236086466, 1.1642551456388435, + 0.3961557091780092, 0.39829918966870376, 0.4012471161917819, 0.40500178050690006, + 0.4095651162997735, 0.4149379498546247, 0.4211190461681518, 0.4281039291887044, + 0.4358834560016674, 0.44444212949443834, 0.4537561437057185, 0.4637911723119031, + 0.4744999352392512, 0.48581961273489255, 0.4976692212346268, 0.5099471205068027 + ], + "hitIds": [ + 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, -1, -1, -1, -1, -1, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, + 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0 + ], + "depth": [ + 0.38658537800206566, 0.3984214421111915, 0.4157035508769481, 0.4395875104028722, + 0.4718682619141523, 0.5153596686513074, 0.5746168887467582, 0.6574101129142974, + 0.778037355764179, 0.9658398799816516, 1, 1, 1, 1, 1, 1, 0.1555243298607375, + 0.15729080147368976, 0.14311256851903906, 0.13852879004136567, 0.14191050729582724, + 0.14600654855090217, 0.15087913038331538, 0.15660333818778474, 0.16326860348551542, + 0.17098029703334838, 0.17986132372287372, 0.19005351274303933, 0.20171844855503634, + 0.21503715811489355, 0.27079413090216164, 0.29106378640971087, 0.0990389272945023, + 0.09957479741717594, 0.10031177904794547, 0.10125044512672501, 0.10239127907494337, + 0.10373448746365617, 0.10527976154203796, 0.1070259822971761, 0.10897086400041685, + 0.11111053237360959, 0.11343903592642962, 0.11594779307797577, 0.1186249838098128, + 0.12145490318372314, 0.1244173053086567, 0.12748678012670067 + ] + }, + { + "layout": "edge", + "position": [5.45, 0, 0.32], + "quaternion": [1.0, 0.0, 0.0, 0.0], + "distances": [ + -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, + -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, + -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, + -1.0, -1.0, -1.0 + ], + "hitIds": [ + -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, + -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, + -1, -1 + ], + "depth": [ + 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, + 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1 + ] + } + ] +} diff --git a/web_platform/src/rl/fixtures/obstacleDeployment.json b/web_platform/src/rl/fixtures/obstacleDeployment.json new file mode 100644 index 00000000..9b9068a0 --- /dev/null +++ b/web_platform/src/rl/fixtures/obstacleDeployment.json @@ -0,0 +1,212 @@ +{ + "version": 1, + "taskId": "Unitree-Go2-ObstacleAvoidance", + "browserCompatible": true, + "observationSize": 81, + "actionSize": 12, + "controlHz": 50, + "gaitPeriod": 0.6, + "jointNames": [ + "FL_hip_joint", + "FL_thigh_joint", + "FL_calf_joint", + "FR_hip_joint", + "FR_thigh_joint", + "FR_calf_joint", + "RL_hip_joint", + "RL_thigh_joint", + "RL_calf_joint", + "RR_hip_joint", + "RR_thigh_joint", + "RR_calf_joint" + ], + "defaultJointPosition": [-0.1, 0.9, -1.8, 0.1, 0.9, -1.8, -0.1, 0.9, -1.8, 0.1, 0.9, -1.8], + "actionScale": [0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25], + "stiffness": [20, 20, 40, 20, 20, 40, 20, 20, 40, 20, 20, 40], + "damping": [1, 1, 2, 1, 1, 2, 1, 1, 2, 1, 1, 2], + "effortLimits": [23.5, 23.5, 45, 23.5, 23.5, 45, 23.5, 23.5, 45, 23.5, 23.5, 45], + "observationTerms": [ + "base_ang_vel", + "projected_gravity", + "command", + "phase", + "joint_pos", + "joint_vel", + "actions", + "forward_depth", + "target_error" + ], + "seed": 42, + "terrainPreset": "discrete_obstacles", + "terrainParams": { + "size": 12, + "obstacle_count": 24, + "obstacle_height_min": 0.2, + "obstacle_height_max": 0.6, + "spacing": 1.2, + "friction": 0.8, + "roughness": 0.06, + "step_height": 0.08, + "wave_amplitude": 0.08 + }, + "terrain": { + "representation": "boxes-v1", + "approximation": false, + "size": 12, + "friction": 0.8, + "boxes": [ + { + "pos": [0, 0, -0.1], + "size": [6.0, 6.0, 0.1], + "yaw": 0 + }, + { + "pos": [0.666666666666667, 1.875, 0.261425654654876], + "size": [0.2, 0.2, 0.261425654654876], + "yaw": 0 + }, + { + "pos": [-0.6666666666666665, 4.375, 0.24594635733876358], + "size": [0.2, 0.2, 0.24594635733876358], + "yaw": 0 + }, + { + "pos": [0.666666666666667, -4.375, 0.20724561829094013], + "size": [0.2, 0.2, 0.20724561829094013], + "yaw": 0 + }, + { + "pos": [-2.0, -1.875, 0.29462315279587414], + "size": [0.2, 0.2, 0.29462315279587414], + "yaw": 0 + }, + { + "pos": [-0.6666666666666665, 1.875, 0.17570687544167068], + "size": [0.2, 0.2, 0.17570687544167068], + "yaw": 0 + }, + { + "pos": [-2.0, -3.125, 0.2104081262546454], + "size": [0.2, 0.2, 0.2104081262546454], + "yaw": 0 + }, + { + "pos": [0.666666666666667, 3.125, 0.265880932850599], + "size": [0.2, 0.2, 0.265880932850599], + "yaw": 0 + }, + { + "pos": [2.0, 4.375, 0.2237039504728492], + "size": [0.2, 0.2, 0.2237039504728492], + "yaw": 0 + }, + { + "pos": [-2.0, 0.625, 0.27234138006215547], + "size": [0.2, 0.2, 0.27234138006215547], + "yaw": 0 + }, + { + "pos": [-3.3333333333333335, -0.625, 0.21547042905135239], + "size": [0.2, 0.2, 0.21547042905135239], + "yaw": 0 + }, + { + "pos": [2.0, -0.625, 0.24091436724298468], + "size": [0.2, 0.2, 0.24091436724298468], + "yaw": 0 + }, + { + "pos": [-3.3333333333333335, 0.625, 0.10916487673113245], + "size": [0.2, 0.2, 0.10916487673113245], + "yaw": 0 + }, + { + "pos": [3.333333333333333, -0.625, 0.14557965513030938], + "size": [0.2, 0.2, 0.14557965513030938], + "yaw": 0 + }, + { + "pos": [0.666666666666667, -1.875, 0.15787759272042143], + "size": [0.2, 0.2, 0.15787759272042143], + "yaw": 0 + }, + { + "pos": [-2.0, -0.625, 0.11595839538472551], + "size": [0.2, 0.2, 0.11595839538472551], + "yaw": 0 + }, + { + "pos": [2.0, -3.125, 0.14655817727220605], + "size": [0.2, 0.2, 0.14655817727220605], + "yaw": 0 + }, + { + "pos": [0.666666666666667, 0.625, 0.12020028588194583], + "size": [0.2, 0.2, 0.12020028588194583], + "yaw": 0 + }, + { + "pos": [3.333333333333333, -3.125, 0.15559472062201843], + "size": [0.2, 0.2, 0.15559472062201843], + "yaw": 0 + }, + { + "pos": [-0.6666666666666665, -1.875, 0.22713688885288003], + "size": [0.2, 0.2, 0.22713688885288003], + "yaw": 0 + }, + { + "pos": [0.666666666666667, 4.375, 0.17296643579401685], + "size": [0.2, 0.2, 0.17296643579401685], + "yaw": 0 + }, + { + "pos": [3.333333333333333, 3.125, 0.17403619342337653], + "size": [0.2, 0.2, 0.17403619342337653], + "yaw": 0 + }, + { + "pos": [2.0, -4.375, 0.14190140615429753], + "size": [0.2, 0.2, 0.14190140615429753], + "yaw": 0 + }, + { + "pos": [2.0, 0.625, 0.15339556440982266], + "size": [0.2, 0.2, 0.15339556440982266], + "yaw": 0 + }, + { + "pos": [0.666666666666667, -3.125, 0.2873309175424988], + "size": [0.2, 0.2, 0.2873309175424988], + "yaw": 0 + } + ], + "spawn": [-5.0, 0, 0.32], + "spawnQuaternion": [1, 0, 0, 0], + "target": [5.0, 0], + "actualObstacleCount": 24 + }, + "sensorCfg": { + "fov": 90, + "maxDistance": 4, + "safetyDistance": 0.5, + "avoidanceWeight": 2, + "type": "raycast", + "rayCount": 32, + "offset": [0.3, 0, 0.05], + "alignment": "base", + "terrainOnly": true, + "includeGround": true + }, + "navigation": { + "speed": 0.6, + "arrivalRadius": 0.5, + "distanceScale": 12, + "headingScale": 3.141592653589793, + "yawGain": 1.0, + "maxYawRate": 1.0, + "episodeSeconds": 20, + "onArrival": "stop", + "onReset": "respawn" + } +} diff --git a/web_platform/src/rl/fixtures/obstacleRayGolden.json b/web_platform/src/rl/fixtures/obstacleRayGolden.json new file mode 100644 index 00000000..91ee54aa --- /dev/null +++ b/web_platform/src/rl/fixtures/obstacleRayGolden.json @@ -0,0 +1,90 @@ +{ + "source": "CPU MuJoCo 3.5.0 mj_ray; task_config seed42 discrete_obstacles; distances normalized by4", + "cases": [ + { + "position": [-5, 0, 0.32], + "quaternion": [1, 0, 0, 0], + "depth": [ + 1, 1, 0.8064353265688073, 0.7754070525042633, 0.3493131891686196, 0.33844992227063814, + 0.32906116691883763, 0.3209809620250947, 0.31407497396875733, 0.3285019206818616, + 0.3862286490933719, 1, 1, 1, 1, 1, 1, 1, 1, 0.634959321852795, 0.6416072675585701, + 0.650082300931107, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1 + ] + }, + { + "position": [-2, 1, 0.32], + "quaternion": [0.9689124217106447, 0, 0.24740395925452294, 0], + "depth": [ + 0.16227742543146714, 0.154643351828059, 0.1480582588705743, 0.1423615934562481, + 0.13742675650547637, 0.1331529312359637, 0.12945920813695258, 0.12628029481538672, + 0.12356334175298786, 0.12126556878831947, 0.11935247679185605, 0.11779649499930897, + 0.11657595910035341, 0.11567434601043744, 0.1150797130480744, 0.1147843050751654, + 0.1147843050751654, 0.1150797130480744, 0.11567434601043744, 0.11657595910035341, + 0.11779649499930897, 0.11935247679185605, 0.12126556878831947, 0.12356334175298787, + 0.12628029481538672, 0.12945920813695258, 0.1331529312359637, 0.13742675650547637, + 0.14236159345624813, 0.14805825887057428, 0.154643351828059, 0.16227742543146714 + ] + }, + { + "position": [0, -1, 0.32], + "quaternion": [0.955336489125606, 0.29552020666133955, 0, 0], + "depth": [ + 0.2262087980617112, 0.23859992736340305, 0.2531146297321152, 0.27024831583146797, + 0.2906703480610164, 0.31530672670613974, 0.34547502041327477, 0.3831145628697801, + 0.431200828086753, 0.49454232783018703, 0.5814468749554897, 0.707609556942636, + 0.9066658372041383, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1 + ] + }, + { + "position": [2, 2, 0.32], + "quaternion": [0.8253356149096783, 0, 0, 0.5646424733950354], + "depth": [ + 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 0.47425413995403354, + 0.4738672755338355, 0.4746985930390542, 0.47675881629202255, 0.48007471719272266, 1, 1, 1, + 1, 1 + ] + } + ], + "boxEdges": [ + { + "origin": [0, 0, 0], + "direction": [1, 0, 0], + "distance": 1.0 + }, + { + "origin": [1, 0, 0], + "direction": [1, 0, 0], + "distance": 0.0 + }, + { + "origin": [1, 0, 0], + "direction": [-1, 0, 0], + "distance": -0.0 + }, + { + "origin": [1, 0, 0], + "direction": [0, 1, 0], + "distance": 1.0 + }, + { + "origin": [2, 0, 0], + "direction": [0, 1, 0], + "distance": -1.0 + }, + { + "origin": [2, 0, 0], + "direction": [-1, 0, 0], + "distance": 1.0 + }, + { + "origin": [1, 1, 0], + "direction": [0, 0, 1], + "distance": 1.0 + }, + { + "origin": [2, 2, 0], + "direction": [1, 0, 0], + "distance": -1.0 + } + ] +} diff --git a/web_platform/src/rl/runtime/Go2ObstacleAvoidanceBindings.test.ts b/web_platform/src/rl/runtime/Go2ObstacleAvoidanceBindings.test.ts new file mode 100644 index 00000000..0c86674b --- /dev/null +++ b/web_platform/src/rl/runtime/Go2ObstacleAvoidanceBindings.test.ts @@ -0,0 +1,55 @@ +import { expect, it, vi } from 'vitest'; +import type { MjData, MjModel } from '@mujoco/mujoco'; +import fixture from '../fixtures/obstacleDeployment.json'; +import { validatePolicyDeployment } from '../deployment'; +vi.mock('./Go2wPolicyBindings', () => ({ + Go2wPolicyBindings: class { + baseBodyId = 0; + constructor( + public model: MjModel, + public data: MjData, + ) {} + clear() {} + observe( + _time: number, + _last: Float32Array, + command: { linearX: number; linearY: number; angularZ: number }, + ) { + const obs = new Float32Array(47); + obs.set([command.linearX, command.linearY, command.angularZ], 6); + return obs; + } + }, +})); +import { Go2ObstacleAvoidanceBindings } from './Go2ObstacleAvoidanceBindings'; + +it('浏览器连续导航超过20秒仍观察;换点改变command/目标误差,跌倒越界仍明确停止', () => { + const original = validatePolicyDeployment(fixture); + const deployment = { ...original, terrain: { ...original.terrain!, boxes: [] } }; + const data = { + time: 0, + xpos: new Float64Array([0, 0, 0.32]), + xquat: new Float64Array([1, 0, 0, 0]), + } as unknown as MjData; + const bindings = new Go2ObstacleAvoidanceBindings( + { ngeom: 0 } as MjModel, + data, + vi.fn(), + deployment, + ); + bindings.setNavigationTarget([3, 0]); + const before = bindings.observe(21, new Float32Array(12)); + bindings.setNavigationTarget([0, 3]); + const after = bindings.observe(22, new Float32Array(12)); + expect(before[6]).toBeGreaterThan(0); + expect(after[6]).toBeCloseTo(0); + expect(after[8]).toBe(1); + expect(after[79]).not.toBe(before[79]); + data.xpos[2] = 0.1; + expect(() => bindings.observe(23, new Float32Array(12))).toThrow(/跌倒\/越界/); + bindings.setNavigationTarget([2, 0]); + expect(() => bindings.observe(24, new Float32Array(12))).toThrow(/安全停止/); + data.xpos[2] = 0.32; + data.xpos[0] = 100; + expect(() => bindings.observe(25, new Float32Array(12))).toThrow(/安全停止/); +}); diff --git a/web_platform/src/rl/runtime/Go2ObstacleAvoidanceBindings.ts b/web_platform/src/rl/runtime/Go2ObstacleAvoidanceBindings.ts new file mode 100644 index 00000000..31e86424 --- /dev/null +++ b/web_platform/src/rl/runtime/Go2ObstacleAvoidanceBindings.ts @@ -0,0 +1,165 @@ +import type { MjModel, MjData } from '@mujoco/mujoco'; +import type { NavigationStatus } from '../types'; +import type { PolicyDeployment } from '../deployment'; +import { + buildGo2ObstacleAvoidanceObservation, + obstacleNavigation, + rayBoxDistance, + sampleForwardRays, + forwardRayDirections, + type PerceptionRay, +} from '../tasks/go2ObstacleAvoidance'; +import { Go2wPolicyBindings } from './Go2wPolicyBindings'; + +/** Only named boxes from the compiled training scene are sensed: never robot visual/collision geoms. */ +export class Go2ObstacleAvoidanceBindings extends Go2wPolicyBindings { + private readonly terrainIds: number[] = []; + private readonly boxes: { + center: number[]; + half: number[]; + rotation: number[]; + axisAligned: boolean; + }[] = []; + private readonly directions: number[][]; + private startTime: number; + private readonly defaultTarget: [number, number]; + private currentTarget: [number, number]; + rays: PerceptionRay[] = []; + constructor( + model: MjModel, + data: MjData, + setActuator: (id: number, value: number) => void, + private readonly deployment: PolicyDeployment, + ) { + super(model, data, setActuator, deployment.effortLimits); + if (!deployment.terrain || !deployment.sensorCfg) + throw new Error('避障策略缺少地图/传感器配置'); + this.directions = forwardRayDirections(deployment.sensorCfg); + this.defaultTarget = [deployment.terrain.target[0], deployment.terrain.target[1]]; + this.currentTarget = [...this.defaultTarget]; + this.startTime = Number(data.time); + for (let id = 0; id < model.ngeom; id++) { + const geom = model.geom(id); + try { + if (!geom.name.startsWith('__training_terrain_')) continue; + if ( + Number(model.geom_type[id]) !== 6 || + Number(model.geom_bodyid[id]) !== 0 || + Number(model.geom_group[id]) !== 2 + ) + throw new Error('训练地图必须是静态group2 box'); + this.terrainIds.push(id); + this.boxes.push({ + center: Array.from(data.geom_xpos.subarray(id * 3, id * 3 + 3), Number), + half: Array.from(model.geom_size.subarray(id * 3, id * 3 + 3), Number), + rotation: Array.from(data.geom_xmat.subarray(id * 9, id * 9 + 9), Number), + axisAligned: Array.from(data.geom_xmat.subarray(id * 9, id * 9 + 9), Number).every( + (v, i) => v === (i % 4 === 0 ? 1 : 0), + ), + }); + } finally { + geom.delete(); + } + } + if (this.terrainIds.length !== deployment.terrain.boxes.length) + throw new Error('请先加载策略配套训练地图'); + } + setNavigationTarget(target: [number, number]): void { + if (target.length !== 2 || !target.every(Number.isFinite)) + throw new Error('导航目标必须是两个有限坐标'); + const limit = this.deployment.terrain!.size / 2 - 0.5; + this.currentTarget = target.map((v) => Math.max(-limit, Math.min(limit, v))) as [ + number, + number, + ]; + } + resetNavigationTarget(): void { + this.currentTarget = [...this.defaultTarget]; + } + navigationStatus(): NavigationStatus { + const [x, y] = this.currentTarget; + let height = 0; + for (const id of this.terrainIds) { + const distance = rayBoxDistance( + [x, y, 10000], + [0, 0, -1], + this.data.geom_xpos.subarray(id * 3, id * 3 + 3), + this.model.geom_size.subarray(id * 3, id * 3 + 3), + this.data.geom_xmat.subarray(id * 9, id * 9 + 9), + ); + if (distance >= 0) height = Math.max(height, 10000 - distance); + } + return { + target: [...this.currentTarget], + defaultTarget: [...this.defaultTarget], + targetHeight: height, + distance: Math.hypot( + x - Number(this.data.xpos[this.baseBodyId * 3]), + y - Number(this.data.xpos[this.baseBodyId * 3 + 1]), + ), + }; + } + reset(time: number) { + this.resetNavigationTarget(); + this.startTime = time; + this.rays = []; + } + override clear() { + super.clear(); + this.rays = []; + } + override observe(time: number, lastAction: Float32Array): Float32Array { + const terrain = this.deployment.terrain!, + sensor = this.deployment.sensorCfg!; + const position = Array.from( + this.data.xpos.subarray(this.baseBodyId * 3, this.baseBodyId * 3 + 3), + Number, + ); + const quaternion = Array.from( + this.data.xquat.subarray(this.baseBodyId * 4, this.baseBodyId * 4 + 4), + Number, + ); + // Interactive navigation has no episode deadline. Training/evaluation retain 20s. + // Safety failures still require an explicit reset/re-enable, never a new target. + if ( + position[2] < 0.12 || + Math.abs(position[0]) > terrain.size / 2 - 0.3 || + Math.abs(position[1]) > terrain.size / 2 - 0.3 + ) + throw new Error('导航安全停止(跌倒/越界),请重置仿真后重新启用'); + const navigation = obstacleNavigation( + position, + quaternion, + this.currentTarget, + terrain.size, + this.deployment.navigation?.speed, + ); + const sample = sampleForwardRays( + position, + quaternion, + sensor, + (origin, direction) => { + let nearest = -1; + for (const box of this.boxes) { + const distance = rayBoxDistance( + origin, + direction, + box.center, + box.half, + box.rotation, + box.axisAligned, + ); + if (distance >= 0 && (nearest < 0 || distance < nearest)) nearest = distance; + } + return nearest; + }, + this.directions, + ); + this.rays = sample.rays; + return buildGo2ObstacleAvoidanceObservation( + super.observe(time - this.startTime, lastAction, navigation.command), + sample.depth, + navigation.targetError, + ); + } +} diff --git a/web_platform/src/rl/runtime/Go2wPolicyBindings.ts b/web_platform/src/rl/runtime/Go2wPolicyBindings.ts index 6152c5ec..c3380307 100644 --- a/web_platform/src/rl/runtime/Go2wPolicyBindings.ts +++ b/web_platform/src/rl/runtime/Go2wPolicyBindings.ts @@ -27,15 +27,16 @@ function rotateInverse( /** 将 mjlab Go2 velocity 的 47 维 actor 观测和 12 维关节位置动作映射到 MuJoCo。 */ export class Go2wPolicyBindings implements PolicyRuntimeBindings { private readonly joints: BoundJoint[]; - private readonly baseBodyId: number; + protected readonly baseBodyId: number; private readonly baseFreeJointId: number; private readonly gyroSensorId?: number; private readonly wheelActuatorIds: number[]; constructor( - private readonly model: MjModel, - private readonly data: MjData, + protected readonly model: MjModel, + protected readonly data: MjData, private readonly setActuator: (id: number, value: number) => void, + private readonly effortLimits?: readonly number[], ) { const jointIds = new Map(), actuatorIds = new Map(), @@ -199,7 +200,11 @@ export class Go2wPolicyBindings implements PolicyRuntimeBindings { GO2W_VELOCITY_TASK.damping[index] * Number(this.data.qvel[item.qvelAddress]); this.setActuator( item.actuatorId, - item.positionActuator ? target : torque / item.controlScale, + item.positionActuator + ? target + : (this.effortLimits + ? Math.max(-this.effortLimits[index], Math.min(this.effortLimits[index], torque)) + : torque) / item.controlScale, ); } for (const id of this.wheelActuatorIds) this.setActuator(id, 0); diff --git a/web_platform/src/rl/runtime/OnnxPolicyRuntime.test.ts b/web_platform/src/rl/runtime/OnnxPolicyRuntime.test.ts new file mode 100644 index 00000000..471a9fb0 --- /dev/null +++ b/web_platform/src/rl/runtime/OnnxPolicyRuntime.test.ts @@ -0,0 +1,217 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import fixture from '../fixtures/obstacleDeployment.json'; +import { validatePolicyDeployment } from '../deployment'; +const mock = vi.hoisted(() => ({ + create: vi.fn(), + tensors: [] as { dispose: ReturnType }[], +})); +vi.mock('onnxruntime-web/wasm', () => ({ + env: { wasm: {} }, + InferenceSession: { create: mock.create }, + Tensor: class { + dispose = vi.fn(); + constructor( + public type: string, + public data: Float32Array, + public dims: number[], + ) { + mock.tensors.push(this); + } + }, +})); +import { OnnxPolicyRuntime } from './OnnxPolicyRuntime'; +const deployment = validatePolicyDeployment(fixture); +const flush = async () => { + for (let i = 0; i < 10; i++) await Promise.resolve(); +}; +function setup(size = 81) { + const session = { + inputNames: ['obs'], + outputNames: ['actions'], + inputMetadata: [{ isTensor: true, type: 'float32', shape: [1, size] }], + outputMetadata: [{ isTensor: true, type: 'float32', shape: [1, 12] }], + run: vi.fn(), + release: vi.fn().mockResolvedValue(undefined), + }; + const bindings = { + observe: vi.fn(() => new Float32Array(size)), + apply: vi.fn(), + clear: vi.fn(), + reset: vi.fn(), + }; + mock.create.mockResolvedValue(session); + return { session, bindings }; +} +beforeEach(() => { + vi.clearAllMocks(); + mock.tensors.length = 0; +}); +describe('OnnxPolicyRuntime held-action', () => { + it('动态目标转发不重建ORT、不启用策略或打断in-flight,status实时读取', async () => { + const { session, bindings } = setup(); + let finish!: (outputs: Record) => void; + session.run.mockImplementation( + () => + new Promise((resolve) => { + finish = resolve; + }), + ); + const navigation = { + target: [1, 2] as [number, number], + defaultTarget: [5, 0] as [number, number], + distance: 3, + targetHeight: 0, + }; + const targetBindings = { + ...bindings, + navigationStatus: () => navigation, + setNavigationTarget: vi.fn(), + resetNavigationTarget: vi.fn(), + }; + const runtime = await OnnxPolicyRuntime.load( + new Uint8Array(), + 'goal.onnx', + targetBindings, + deployment, + ); + runtime.setNavigationTarget([2, 3]); + expect(targetBindings.setNavigationTarget).toHaveBeenCalledWith([2, 3]); + expect(runtime.status().enabled).toBe(false); + expect(runtime.status().navigation).toBe(navigation); + runtime.setEnabled(true, 0); + runtime.step(0); + runtime.setNavigationTarget([3, 4]); + runtime.step(0.02); + expect(session.run).toHaveBeenCalledOnce(); + expect(mock.create).toHaveBeenCalledOnce(); + runtime.resetNavigationTarget(); + expect(targetBindings.resetNavigationTarget).toHaveBeenCalledOnce(); + finish({ actions: { type: 'float32', data: new Float32Array(12), dispose: vi.fn() } }); + await flush(); + runtime.dispose(); + expect(runtime.status().navigation).toBeUndefined(); + }); + it('81维契约单in-flight,持有最新动作,不阻塞物理步;reset拒绝旧promise', async () => { + const { session, bindings } = setup(); + let finish!: (outputs: Record) => void; + session.run.mockImplementation( + () => + new Promise((resolve) => { + finish = resolve; + }), + ); + const runtime = await OnnxPolicyRuntime.load( + new Uint8Array([1]), + 'policy.onnx', + bindings, + deployment, + ); + runtime.setEnabled(true, 0); + runtime.step(0); + runtime.step(0.02); + runtime.step(0.04); + expect(session.run).toHaveBeenCalledOnce(); + expect(bindings.apply).toHaveBeenCalledTimes(3); + const output = { type: 'float32', data: new Float32Array(12).fill(0.3), dispose: vi.fn() }; + finish({ actions: output }); + await flush(); + runtime.step(0.06); + expect(bindings.apply.mock.lastCall?.[0][0]).toBeCloseTo(0.3); + expect(session.run).toHaveBeenCalledTimes(2); + expect(output.dispose).toHaveBeenCalledOnce(); + runtime.reset(0); + finish({ actions: { ...output, data: new Float32Array(12).fill(0.9) } }); + await flush(); + runtime.step(0); + expect(bindings.apply.mock.lastCall?.[0][0]).toBe(0); + expect(bindings.reset).toHaveBeenCalledWith(0); + runtime.dispose(); + expect(session.release).not.toHaveBeenCalled(); + finish({ actions: output }); + await flush(); + expect(session.release).toHaveBeenCalledOnce(); + expect(mock.tensors.every((t) => t.dispose.mock.calls.length === 1)).toBe(true); + }); + it('81策略不能默认为47、Rough维度不能加载,失败释放session', async () => { + const { session, bindings } = setup(); + await expect(OnnxPolicyRuntime.load(new Uint8Array(), 'wrong.onnx', bindings)).rejects.toThrow( + /维度/, + ); + expect(session.release).toHaveBeenCalledOnce(); + session.inputMetadata[0].shape = [1, 234]; + await expect( + OnnxPolicyRuntime.load(new Uint8Array(), 'rough.onnx', bindings, deployment), + ).rejects.toThrow(/维度/); + }); + it('原47维Flat保持向后兼容,非法观测/输出失败关闭并清动作', async () => { + const { session, bindings } = setup(47); + session.run.mockResolvedValue({ + actions: { type: 'float32', data: new Float32Array(12).fill(NaN), dispose: vi.fn() }, + }); + const runtime = await OnnxPolicyRuntime.load(new Uint8Array(), 'flat.onnx', bindings); + expect(runtime.status().observationSize).toBe(47); + runtime.setEnabled(true, 0); + runtime.step(0); + await flush(); + expect(runtime.status().enabled).toBe(false); + expect(runtime.status().error).toMatch(/非有限/); + runtime.reset(0); + bindings.observe.mockReturnValue(new Float32Array(81)); + runtime.setEnabled(true, 0); + runtime.step(0); + expect(runtime.status().error).toMatch(/维度/); + runtime.dispose(); + await flush(); + }); +}); + +it('默认Flat作业的兼容契约仍检查真实graph,拒绝81维及动态特征', async () => { + const flat = validatePolicyDeployment({ + ...fixture, + taskId: 'Unitree-Go2-Flat', + observationSize: 47, + observationTerms: fixture.observationTerms.slice(0, 7), + terrain: undefined, + terrainPreset: undefined, + terrainParams: undefined, + sensorCfg: undefined, + navigation: undefined, + }); + const { session, bindings } = setup(47); + const runtime = await OnnxPolicyRuntime.load(new Uint8Array(), 'legacy.onnx', bindings, flat); + expect(runtime.status().observationSize).toBe(47); + runtime.dispose(); + await flush(); + session.inputMetadata[0].shape = [1, 81]; + await expect( + OnnxPolicyRuntime.load(new Uint8Array(), 'legacy-wrong.onnx', bindings, flat), + ).rejects.toThrow(/维度/); + session.inputMetadata[0].shape = [1, -1]; + await expect( + OnnxPolicyRuntime.load(new Uint8Array(), 'legacy-dynamic.onnx', bindings, flat), + ).rejects.toThrow(/固定/); +}); + +it('97 metadata不能冒充81 graph;真实97 shape沿用single-flight生命周期', async () => { + const multi = validatePolicyDeployment({ + ...fixture, + observationSize: 97, + sensorCfg: { ...fixture.sensorCfg, sensorMode: 'multi_ring_raycast', rayCount: 48 }, + }); + const { session, bindings } = setup(81); + await expect( + OnnxPolicyRuntime.load(new Uint8Array([1]), 'multi.onnx', bindings, multi), + ).rejects.toThrow(/维度/); + expect(session.release).toHaveBeenCalledOnce(); + const next = setup(97); + const runtime = await OnnxPolicyRuntime.load( + new Uint8Array([1]), + 'multi.onnx', + next.bindings, + multi, + ); + expect(runtime.status().observationSize).toBe(97); + runtime.dispose(); + await flush(); + expect(next.session.release).toHaveBeenCalledOnce(); +}); diff --git a/web_platform/src/rl/runtime/OnnxPolicyRuntime.ts b/web_platform/src/rl/runtime/OnnxPolicyRuntime.ts index 84d1cde7..da6af466 100644 --- a/web_platform/src/rl/runtime/OnnxPolicyRuntime.ts +++ b/web_platform/src/rl/runtime/OnnxPolicyRuntime.ts @@ -1,6 +1,8 @@ import * as ort from 'onnxruntime-web/wasm'; import { GO2W_VELOCITY_TASK, clampGo2wCommand } from '../tasks/go2wVelocity'; -import type { RLCommand, RLPolicyStatus } from '../types'; +import type { PolicyDeployment } from '../deployment'; +import { GO2_OBSTACLE_AVOIDANCE_TASK } from '../tasks/go2ObstacleAvoidance'; +import type { NavigationStatus, RLCommand, RLPolicyStatus } from '../types'; ort.env.wasm.numThreads = 1; ort.env.wasm.proxy = false; @@ -9,6 +11,10 @@ export interface PolicyRuntimeBindings { observe(time: number, lastAction: Float32Array, command: RLCommand): Float32Array; apply(action: Float32Array): void; clear(): void; + reset?(time: number): void; + navigationStatus?(): NavigationStatus; + setNavigationTarget?(target: [number, number]): void; + resetNavigationTarget?(): void; } function message(error: unknown): string { @@ -38,13 +44,25 @@ export class OnnxPolicyRuntime { private readonly path: string, private readonly inputName: string, private readonly outputName: string, + private readonly task: { + id: string; + name: string; + observationSize: number; + actionSize: number; + controlHz: number; + }, ) {} static async load( model: Uint8Array, path: string, bindings: PolicyRuntimeBindings, + deployment?: PolicyDeployment, ): Promise { + const task = + deployment?.taskId === GO2_OBSTACLE_AVOIDANCE_TASK.id + ? { ...GO2_OBSTACLE_AVOIDANCE_TASK, observationSize: deployment.observationSize } + : GO2W_VELOCITY_TASK; const session = await ort.InferenceSession.create(model.slice(), { executionProviders: ['wasm'], graphOptimizationLevel: 'all', @@ -71,28 +89,19 @@ export class OnnxPolicyRuntime { throw new Error(`策略输入 batch 必须为 1 或动态维度,实际为 ${inputBatch}`); if (typeof outputBatch === 'number' && outputBatch !== -1 && outputBatch !== 1) throw new Error(`策略输出 batch 必须为 1 或动态维度,实际为 ${outputBatch}`); - if ( - typeof fixedInput === 'number' && - fixedInput > 0 && - fixedInput !== GO2W_VELOCITY_TASK.observationSize - ) - throw new Error( - `策略观测维度不匹配:模型 ${fixedInput},任务 ${GO2W_VELOCITY_TASK.observationSize}`, - ); - if ( - typeof fixedOutput === 'number' && - fixedOutput > 0 && - fixedOutput !== GO2W_VELOCITY_TASK.actionSize - ) - throw new Error( - `策略动作维度不匹配:模型 ${fixedOutput},任务 ${GO2W_VELOCITY_TASK.actionSize}`, - ); + if (typeof fixedInput === 'number' && fixedInput > 0 && fixedInput !== task.observationSize) + throw new Error(`策略观测维度不匹配:模型 ${fixedInput},任务 ${task.observationSize}`); + if (typeof fixedOutput === 'number' && fixedOutput > 0 && fixedOutput !== task.actionSize) + throw new Error(`策略动作维度不匹配:模型 ${fixedOutput},任务 ${task.actionSize}`); + if (deployment && (fixedInput !== task.observationSize || fixedOutput !== task.actionSize)) + throw new Error('部署策略必须声明固定的观测/动作特征维度'); return new OnnxPolicyRuntime( session, bindings, path, session.inputNames[0], session.outputNames[0], + task, ); } catch (error) { await session.release(); @@ -102,22 +111,32 @@ export class OnnxPolicyRuntime { status(): RLPolicyStatus { return { - taskId: GO2W_VELOCITY_TASK.id, - taskName: GO2W_VELOCITY_TASK.name, + taskId: this.task.id, + taskName: this.task.name, path: this.path, loaded: !this.disposed, enabled: this.enabled, - controlHz: GO2W_VELOCITY_TASK.controlHz, - observationSize: GO2W_VELOCITY_TASK.observationSize, - actionSize: GO2W_VELOCITY_TASK.actionSize, + controlHz: this.task.controlHz, + observationSize: this.task.observationSize, + actionSize: this.task.actionSize, inputName: this.inputName, outputName: this.outputName, command: { ...this.commandValue }, inferenceCount: this.inferenceCount, lastInferenceMs: this.lastInferenceMs, + navigation: this.navigationStatus(), error: this.error, }; } + navigationStatus(): NavigationStatus | undefined { + return this.disposed ? undefined : this.bindings.navigationStatus?.(); + } + setNavigationTarget(target: [number, number]): void { + if (!this.disposed) this.bindings.setNavigationTarget?.(target); + } + resetNavigationTarget(): void { + if (!this.disposed) this.bindings.resetNavigationTarget?.(); + } setCommand(command: RLCommand): void { this.commandValue = clampGo2wCommand(command); } @@ -138,6 +157,7 @@ export class OnnxPolicyRuntime { this.nextInferenceTime = time; this.error = undefined; this.bindings.clear(); + this.bindings.reset?.(time); } step(time: number): void { @@ -147,12 +167,14 @@ export class OnnxPolicyRuntime { let observation: Float32Array; try { observation = this.bindings.observe(time, this.action, this.commandValue); + if (observation.length !== this.task.observationSize || !observation.every(Number.isFinite)) + throw new Error('策略观测维度错误或包含非有限数'); } catch (error) { this.fail(error); return; } this.inFlight = true; - this.nextInferenceTime = time + 1 / GO2W_VELOCITY_TASK.controlHz; + this.nextInferenceTime = time + 1 / this.task.controlHz; const started = performance.now(), epoch = this.epoch; const input = new ort.Tensor('float32', observation, [1, observation.length]); @@ -163,9 +185,9 @@ export class OnnxPolicyRuntime { const output = outputs[this.outputName]; if (!output || output.type !== 'float32') throw new Error(`找不到 float32 输出:${this.outputName}`); - if (output.data.length !== GO2W_VELOCITY_TASK.actionSize) + if (output.data.length !== this.task.actionSize) throw new Error( - `策略动作维度错误:期望 ${GO2W_VELOCITY_TASK.actionSize},实际 ${output.data.length}`, + `策略动作维度错误:期望 ${this.task.actionSize},实际 ${output.data.length}`, ); const next = Float32Array.from(output.data as Float32Array, Number); for (const value of next) diff --git a/web_platform/src/rl/tasks/go2ObstacleAvoidance.test.ts b/web_platform/src/rl/tasks/go2ObstacleAvoidance.test.ts new file mode 100644 index 00000000..7ef69c5b --- /dev/null +++ b/web_platform/src/rl/tasks/go2ObstacleAvoidance.test.ts @@ -0,0 +1,114 @@ +import { describe, it, expect } from 'vitest'; +import deploymentJSON from '../fixtures/obstacleDeployment.json'; +import golden from '../fixtures/obstacleRayGolden.json'; +import { validatePolicyDeployment } from '../deployment'; +import { + buildGo2ObstacleAvoidanceObservation, + obstacleNavigation, + rayBoxDistance, + sampleForwardRays, +} from './go2ObstacleAvoidance'; +const deployment = validatePolicyDeployment(deploymentJSON); +const identity = [1, 0, 0, 0, 1, 0, 0, 0, 1]; +describe('go2ObstacleAvoidance', () => { + it('严格47+32+2顺序及有限值、归一化范围校验', () => { + const base = Array.from({ length: 47 }, (_, i) => i), + depth = Array(32).fill(0.75); + const result = buildGo2ObstacleAvoidanceObservation(base, depth, [0.25, 0.5]); + expect(result).toBeInstanceOf(Float32Array); + expect(result).toHaveLength(81); + expect(Array.from(result.slice(0, 47))).toEqual(base); + expect(Array.from(result.slice(47, 79))).toEqual(depth); + expect(Array.from(result.slice(79))).toEqual([0.25, 0.5]); + expect(() => buildGo2ObstacleAvoidanceObservation([], depth, [0, 0])).toThrow(/维度/); + expect(() => buildGo2ObstacleAvoidanceObservation(base, depth, [NaN, 0])).toThrow(/非有限/); + expect(() => + buildGo2ObstacleAvoidanceObservation(base, Array(32).fill(Infinity), [0, 0]), + ).toThrow(/非有限/); + expect(() => buildGo2ObstacleAvoidanceObservation(base, Array(32).fill(1.1), [0, 0])).toThrow( + /0~1/, + ); + }); + it('目标前向、身后、yaw旋转、到达停止及离开恢复与后端公式一致', () => { + expect(obstacleNavigation([-5, 0, 0.32], [1, 0, 0, 0], [5, 0], 12)).toEqual({ + command: { linearX: 0.6, linearY: 0, angularZ: 0 }, + targetError: [0, 10 / 12], + }); + expect(obstacleNavigation([5, 0, 0.32], [1, 0, 0, 0], [-5, 0], 12).command).toEqual({ + linearX: 0, + linearY: 0, + angularZ: 1, + }); + expect(obstacleNavigation([4.8, 0, 0.32], [0, 0, 0, 1], [5, 0], 12).command.linearX).toBe(0); + expect(obstacleNavigation([4.4, 0, 0.32], [1, 0, 0, 0], [5, 0], 12).command.linearX).toBe(0.6); + expect( + obstacleNavigation([0, 0, 0], [Math.SQRT1_2, 0, 0, Math.SQRT1_2], [0, 2], 12).targetError[0], + ).toBeCloseTo(0); + }); + it('动态换目标维持81维末尾归一化误差,到达后恢复转向', () => { + const position = [0, 0, 0.32], + quaternion = [1, 0, 0, 0]; + expect(obstacleNavigation(position, quaternion, [0, 0], 12).command.angularZ).toBe(0); + for (const [target, heading] of [ + [[0, 3], 0.5], + [[0, -3], -0.5], + ] as const) { + const navigation = obstacleNavigation(position, quaternion, target, 12); + const observation = buildGo2ObstacleAvoidanceObservation( + Array(47).fill(0), + Array(32).fill(1), + navigation.targetError, + ); + expect(observation).toHaveLength(81); + expect(Array.from(observation.slice(79))).toEqual([heading, 0.25]); + expect(navigation.command.angularZ).toBe(Math.sign(heading)); + } + expect(obstacleNavigation(position, quaternion, [24, 0], 12).targetError[1]).toBe(1); + }); + it('全部32射线与CPU mj_ray golden对照:完整yaw/pitch/roll且包含地面命中', () => { + for (const fixture of golden.cases) { + const result = sampleForwardRays( + fixture.position, + fixture.quaternion, + deployment.sensorCfg!, + (origin, direction) => { + const hits = deployment + .terrain!.boxes.map((box) => + rayBoxDistance(origin, direction, box.pos, box.size, identity), + ) + .filter((d) => d >= 0); + return hits.length ? Math.min(...hits) : -1; + }, + ); + result.depth.forEach((d, i) => expect(d).toBeCloseTo(fixture.depth[i], 6)); + expect(result.rays).toHaveLength(32); + } + }); + it('box内部/表面/平行/边缘/miss与CPU mj_ray一致', () => { + golden.boxEdges.forEach((item) => + expect( + rayBoxDistance(item.origin, item.direction, [0, 0, 0], [1, 1, 1], identity), + ).toBeCloseTo(item.distance, 8), + ); + }); + it('按右到左含端点采样,maxRange等号视为命中,越界/miss为1', () => { + let count = 0; + const sample = sampleForwardRays( + [0, 0, 0], + [1, 0, 0, 0], + deployment.sensorCfg!, + () => [4, 4.01, -1, 0][count++ % 4], + ); + expect(sample.depth.slice(0, 4)).toEqual([1, 1, 1, 0]); + expect(sample.rays.slice(0, 4).map((r) => r.hit)).toEqual([true, false, false, true]); + expect(sample.rays[0].end[1]).toBeLessThan(0); + expect(sample.rays[30].end[1]).toBeGreaterThan(0); + }); +}); + +it('使用部署中的target_velocity命令幅值,不改变81维目标误差或到达停机', () => { + const result = obstacleNavigation([0, 0, 0.32], [1, 0, 0, 0], [5, 0], 12, 0.9); + expect(result.command.linearX).toBe(0.9); + expect(result.targetError).toEqual([0, 5 / 12]); + expect(obstacleNavigation([4.8, 0, 0.32], [1, 0, 0, 0], [5, 0], 12, 1.2).command.linearX).toBe(0); +}); diff --git a/web_platform/src/rl/tasks/go2ObstacleAvoidance.ts b/web_platform/src/rl/tasks/go2ObstacleAvoidance.ts new file mode 100644 index 00000000..b6c8e7a5 --- /dev/null +++ b/web_platform/src/rl/tasks/go2ObstacleAvoidance.ts @@ -0,0 +1,160 @@ +import { GO2W_VELOCITY_TASK } from './go2wVelocity'; +import { OBSTACLE_TASK_ID, obstacleSensorPattern, type ObstacleSensorConfig } from '../deployment'; +import type { RLCommand } from '../types'; + +export const GO2_OBSTACLE_AVOIDANCE_TASK = { + ...GO2W_VELOCITY_TASK, + id: OBSTACLE_TASK_ID, + name: 'Go2 前视射线避障导航', + observationSize: 81, +}; +export interface PerceptionRay { + origin: number[]; + end: number[]; + hit: boolean; +} +export function rotateVector(q: readonly number[], v: readonly number[]): number[] { + const [w, x, y, z] = q, + [vx, vy, vz] = v; + const tx = 2 * (y * vz - z * vy), + ty = 2 * (z * vx - x * vz), + tz = 2 * (x * vy - y * vx); + return [ + vx + w * tx + y * tz - z * ty, + vy + w * ty + z * tx - x * tz, + vz + w * tz + x * ty - y * tx, + ]; +} +export function obstacleNavigation( + position: readonly number[], + quaternion: readonly number[], + target: readonly number[], + size: number, + speed = 0.6, +): { command: RLCommand; targetError: number[] } { + const [w, x, y, z] = quaternion; + const yaw = Math.atan2(2 * (w * z + x * y), 1 - 2 * (y * y + z * z)); + const dx = target[0] - position[0], + dy = target[1] - position[1], + distance = Math.hypot(dx, dy); + const delta = Math.atan2(dy, dx) - yaw; + const heading = distance < 0.5 ? 0 : Math.atan2(Math.sin(delta), Math.cos(delta)); + return { + command: + distance < 0.5 + ? { linearX: 0, linearY: 0, angularZ: 0 } + : { + linearX: speed * Math.max(0, Math.cos(heading)), + linearY: 0, + angularZ: Math.max(-1, Math.min(1, heading)), + }, + targetError: [heading / Math.PI, Math.min(1, distance / size)], + }; +} +/** Slab intersection with a physical box in its local frame. Inside hits the exit, as mj_ray does. */ +export function rayBoxDistance( + origin: readonly number[], + direction: readonly number[], + center: ArrayLike, + halfSize: ArrayLike, + rotation: ArrayLike, + axisAligned = false, +): number { + if (axisAligned) { + let near = -Infinity, + far = Infinity; + for (let i = 0; i < 3; i++) { + const o = origin[i] - center[i], + d = direction[i], + h = halfSize[i]; + if (Math.abs(d) < 1e-12) { + if (Math.abs(o) > h) return -1; + continue; + } + let a = (-h - o) / d, + b = (h - o) / d; + if (a > b) { + const t = a; + a = b; + b = t; + } + if (a > near) near = a; + if (b < far) far = b; + if (near > far || far < 0) return -1; + } + return near >= 0 ? near : far; + } + const x = origin[0] - center[0], + y = origin[1] - center[1], + z = origin[2] - center[2]; + let near = -Infinity, + far = Infinity; + for (let i = 0; i < 3; i++) { + const o = rotation[i] * x + rotation[3 + i] * y + rotation[6 + i] * z; + const d = + rotation[i] * direction[0] + rotation[3 + i] * direction[1] + rotation[6 + i] * direction[2]; + if (Math.abs(d) < 1e-12) { + if (Math.abs(o) > halfSize[i]) return -1; + continue; + } + const a = (-halfSize[i] - o) / d, + b = (halfSize[i] - o) / d; + near = Math.max(near, Math.min(a, b)); + far = Math.min(far, Math.max(a, b)); + if (near > far || far < 0) return -1; + } + return near >= 0 ? near : far; +} +/** Per-binding directions are precomputed, avoiding trig and temporary arrays per ray/box. */ +export function forwardRayDirections(sensor: ObstacleSensorConfig): number[][] { + const pattern = obstacleSensorPattern(sensor.sensorMode, sensor.fov); + return pattern.pitchAngles.flatMap((pitch) => + pattern.yawAngles.map((yaw) => { + const p = (pitch * Math.PI) / 180, + y = (yaw * Math.PI) / 180; + return [Math.cos(p) * Math.cos(y), Math.cos(p) * Math.sin(y), Math.sin(p)]; + }), + ); +} +export function sampleForwardRays( + position: readonly number[], + quaternion: readonly number[], + sensor: ObstacleSensorConfig, + cast: (origin: number[], direction: number[]) => number, + localDirections = forwardRayDirections(sensor), +): { depth: number[]; rays: PerceptionRay[] } { + const offset = rotateVector(quaternion, sensor.offset), + origin = position.map((v, i) => v + offset[i]); + const rays: PerceptionRay[] = [], + depth: number[] = []; + for (const local of localDirections) { + const direction = rotateVector(quaternion, local); + const distance = cast(origin, direction); + if (!Number.isFinite(distance)) throw new Error('物理射线返回非有限距离'); + const hit = distance >= 0 && distance <= sensor.maxDistance; + const length = hit ? distance : sensor.maxDistance; + depth.push(Math.max(0, Math.min(1, length / sensor.maxDistance))); + rays.push({ origin, end: direction.map((v, axis) => origin[axis] + v * length), hit }); + } + return { depth, rays }; +} +export function buildGo2ObstacleAvoidanceObservation( + base: ArrayLike, + depth: ArrayLike, + targetError: ArrayLike, +): Float32Array { + if ( + base.length !== 47 || + (depth.length !== 32 && depth.length !== 48) || + targetError.length !== 2 + ) + throw new Error('避障观测维度必须为47+(32或48)+2'); + const result = Float32Array.from([ + ...Array.from(base), + ...Array.from(depth), + ...Array.from(targetError), + ]); + if (!result.every(Number.isFinite)) throw new Error('避障观测包含非有限数'); + if (Array.from(depth).some((v) => v < 0 || v > 1)) throw new Error('归一化深度必须在0~1之间'); + return result; +} diff --git a/web_platform/src/rl/tasks/multiRing.test.ts b/web_platform/src/rl/tasks/multiRing.test.ts new file mode 100644 index 00000000..b68e0801 --- /dev/null +++ b/web_platform/src/rl/tasks/multiRing.test.ts @@ -0,0 +1,90 @@ +import { expect, it } from 'vitest'; +import legacy from '../fixtures/obstacleDeployment.json'; +import fixture from '../fixtures/multiRingDeployment.json'; +import golden from '../fixtures/multiRingGolden.json'; +import { policyDeploymentsMatch, validatePolicyDeployment } from '../deployment'; +import { + buildGo2ObstacleAvoidanceObservation, + forwardRayDirections, + rayBoxDistance, + sampleForwardRays, +} from './go2ObstacleAvoidance'; +const deployment = validatePolicyDeployment(fixture); +const rotation = [1, 0, 0, 0, 1, 0, 0, 0, 1]; +it('legacy缺字段与规范化metadata语义相等,multi拒绝矛盾及未知组合', () => { + expect( + policyDeploymentsMatch( + legacy as unknown as typeof deployment, + validatePolicyDeployment(legacy), + ), + ).toBe(true); + expect(deployment.observationSize).toBe(97); + for (const patch of [ + { rayCount: 32 }, + { sensorMode: 'unknown' }, + { pitchAngles: [0, -45, -20] }, + { yawCount: 32 }, + { yawAngles: Array(16).fill(0) }, + { angleUnit: 'rad' }, + { rayOrder: 'yaw-major' }, + { garbage: 1 }, + ]) + expect(() => + validatePolicyDeployment({ ...fixture, sensorCfg: { ...fixture.sensorCfg, ...patch } }), + ).toThrow(); + expect(() => validatePolicyDeployment({ ...fixture, observationSize: 81 })).toThrow(/观测维数/); + expect(() => + validatePolicyDeployment({ + ...fixture, + sensorCfg: { ...fixture.sensorCfg, sensorMode: undefined }, + }), + ).toThrow(); +}); +it('所有48ray CPU mj_ray golden:完整姿态、5cm障碍、坑边/地板终点、层序', () => { + for (const c of golden.cases) { + const boxes = golden.layouts[c.layout as keyof typeof golden.layouts]; + const sample = sampleForwardRays(c.position, c.quaternion, deployment.sensorCfg!, (o, d) => { + let nearest = -1; + for (const b of boxes) { + const t = rayBoxDistance(o, d, b.pos, b.size, rotation); + if (t >= 0 && (nearest < 0 || t < nearest)) nearest = t; + } + return nearest; + }); + expect(sample.rays).toHaveLength(48); + c.depth.forEach((d, i) => expect(sample.depth[i]).toBeCloseTo(d, 6)); + const obs = buildGo2ObstacleAvoidanceObservation( + Array(47).fill(0.2), + sample.depth, + [-0.5, 0.3], + ); + expect(obs).toHaveLength(97); + expect(obs[95]).toBe(-0.5); + expect(obs[96]).toBeCloseTo(0.3); + } + const low = golden.cases.find((c) => c.layout === 'low')!; + expect(low.hitIds.slice(0, 16)).not.toContain(1); + expect(low.hitIds.slice(16)).toContain(1); + const rays = forwardRayDirections(deployment.sensorCfg!); + for (const [layer, pitch] of [0, -20, -45].entries()) { + expect(rays[layer * 16][2]).toBeCloseTo(Math.sin((pitch * Math.PI) / 180)); + expect(rays[layer * 16][1]).toBeLessThan(0); + expect(rays[layer * 16 + 15][1]).toBeGreaterThan(0); + } +}); +it('48ray完整range边界、miss、有限值,floor实际距离不隐藏', () => { + let i = 0; + const result = sampleForwardRays( + [0, 0, 0.32], + [1, 0, 0, 0], + deployment.sensorCfg!, + () => [-1, 0, 4, 4 + 1e-9][i++ % 4], + ); + expect(result.depth).toEqual(Array.from({ length: 48 }, (_, j) => [1, 0, 1, 1][j % 4])); + expect(result.rays.map((r) => r.hit)).toEqual( + Array.from({ length: 48 }, (_, j) => [false, true, true, false][j % 4]), + ); + expect(() => + sampleForwardRays([0, 0, 0.32], [1, 0, 0, 0], deployment.sensorCfg!, () => NaN), + ).toThrow(); +}); diff --git a/web_platform/src/rl/tasks/raycastBenchmark.ts b/web_platform/src/rl/tasks/raycastBenchmark.ts new file mode 100644 index 00000000..5d764e04 --- /dev/null +++ b/web_platform/src/rl/tasks/raycastBenchmark.ts @@ -0,0 +1,60 @@ +/** Reproducible full-frame benchmark, not a single-ray microbenchmark. See TRAINING.md. */ +import { forwardRayDirections, rayBoxDistance, sampleForwardRays } from './go2ObstacleAvoidance'; +import { validatePolicyDeployment } from '../deployment'; +import fixture from '../fixtures/multiRingDeployment.json'; +export function benchmarkForwardRays(warmup = 3000, samples = 10000) { + const deployment = validatePolicyDeployment(fixture), + sensor = deployment.sensorCfg!; + const directions = forwardRayDirections(sensor); + const rotation = [1, 0, 0, 0, 1, 0, 0, 0, 1]; + const defaultBoxes = deployment.terrain!.boxes; + // 16x16 grid + floor, same maximal box count as rough/wave deployment. + const worstBoxes = [ + defaultBoxes[0], + ...Array.from({ length: 256 }, (_, i) => ({ + pos: [-3.75 + Math.floor(i / 16) * 0.5, -5.625 + (i % 16) * 0.75, 0.1], + size: [0.25, 0.375, 0.1], + yaw: 0, + })), + ]; + return [defaultBoxes, worstBoxes].map((boxes) => { + const cast = (o: number[], d: number[]) => { + let nearest = -1; + for (const box of boxes) { + const t = rayBoxDistance(o, d, box.pos, box.size, rotation, true); + if (t >= 0 && (nearest < 0 || t < nearest)) nearest = t; + } + return nearest; + }; + let checksum = 0; + const frame = (i: number) => { + // Change pose each frame: no memoized results. Full attitude, origin, all intersections + output. + const yaw = (i % 100) * 0.004, + q = [ + Math.cos(yaw / 2) * Math.cos(0.1), + Math.sin(0.1) * Math.cos(yaw / 2), + Math.sin(0.1) * Math.sin(yaw / 2), + Math.sin(yaw / 2) * Math.cos(0.1), + ]; + const r = sampleForwardRays([-5 + (i % 10) * 0.05, 0, 0.32], q, sensor, cast, directions); + checksum += r.depth[i % 48]; + }; + for (let i = 0; i < warmup; i++) frame(i); + const times = []; + for (let i = 0; i < samples; i++) { + const start = performance.now(); + frame(i); + times.push(performance.now() - start); + } + times.sort((a, b) => a - b); + return { + boxes: boxes.length, + rays: 48, + warmup, + samples, + p50Ms: times[Math.floor(samples * 0.5)], + p95Ms: times[Math.floor(samples * 0.95)], + checksum, + }; + }); +} diff --git a/web_platform/src/rl/types.ts b/web_platform/src/rl/types.ts index 66be3f46..3ed2fb4b 100644 --- a/web_platform/src/rl/types.ts +++ b/web_platform/src/rl/types.ts @@ -4,8 +4,15 @@ export interface RLCommand { angularZ: number; } +export interface NavigationStatus { + target: [number, number]; + defaultTarget: [number, number]; + distance: number; + targetHeight: number; +} + export interface RLPolicyStatus { - taskId: 'unitree-go2w-velocity'; + taskId: string; taskName: string; path: string; loaded: boolean; @@ -18,6 +25,7 @@ export interface RLPolicyStatus { command: RLCommand; inferenceCount: number; lastInferenceMs: number; + navigation?: NavigationStatus; error?: string; } diff --git a/web_platform/src/simulation/PhysicsAdapter.trainingTransaction.test.ts b/web_platform/src/simulation/PhysicsAdapter.trainingTransaction.test.ts new file mode 100644 index 00000000..bd5b3777 --- /dev/null +++ b/web_platform/src/simulation/PhysicsAdapter.trainingTransaction.test.ts @@ -0,0 +1,259 @@ +import { readFileSync } from 'node:fs'; +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import type { ProjectManifest } from '../project/types'; +import fixture from '../rl/fixtures/obstacleDeployment.json'; +import { validatePolicyDeployment } from '../rl/deployment'; +const ort = vi.hoisted(() => ({ create: vi.fn() })); +vi.mock('onnxruntime-web/wasm', () => ({ + env: { wasm: {} }, + InferenceSession: { create: ort.create }, +})); +vi.mock('@mujoco/mujoco', async (original) => { + const actual = await original(); + return { + ...actual, + default: () => + actual.default({ wasmBinary: readFileSync('node_modules/@mujoco/mujoco/mujoco.wasm') }), + }; +}); +import { MainThreadPhysicsAdapter } from './PhysicsAdapter'; +const deployment = validatePolicyDeployment(fixture); +function graph(size = 47) { + return { + inputNames: ['obs'], + outputNames: ['actions'], + inputMetadata: [{ isTensor: true, type: 'float32', shape: [1, size] }], + outputMetadata: [{ isTensor: true, type: 'float32', shape: [1, 12] }], + release: vi.fn().mockResolvedValue(undefined), + }; +} +function project(): ProjectManifest { + const doc = new DOMParser().parseFromString( + readFileSync('training_server/rl/src/assets/robots/unitree_go2/xmls/go2.xml', 'utf8'), + 'application/xml', + ); + doc.querySelectorAll('mesh, geom[mesh]').forEach((e) => e.remove()); + const actuators = doc.createElement('actuator'); + for (const name of deployment.jointNames) { + const motor = doc.createElement('motor'); + motor.setAttribute('joint', name); + motor.setAttribute('name', `${name}_motor`); + actuators.append(motor); + } + doc.documentElement.append(actuators); + const data = new TextEncoder().encode(new XMLSerializer().serializeToString(doc)); + return { + id: 'transaction', + name: 'go2', + files: [{ path: 'go2.xml', data, size: data.length, source: 'file', mimeType: 'text/xml' }], + entries: [{ path: 'go2.xml', format: 'mjcf', label: 'go2' }], + maps: [], + selectedEntry: 'go2.xml', + totalBytes: data.length, + }; +} +async function existing() { + const adapter = new MainThreadPhysicsAdapter(), + source = project(), + oldGraph = graph(); + await adapter.load(source, 'go2.xml'); + ort.create.mockResolvedValueOnce(oldGraph); + await adapter.loadRLPolicy(new Uint8Array([1]), 'old-flat.onnx'); + adapter.setRLPolicyEnabled(true); + const old = adapter.session!; + old.data.time = 7; + old.data.qvel[0] = 0.123; + old.data.qpos[0] = 1.25; + old.data.ctrl[0] = 0.3; + adapter.setPaused(false); + return { adapter, source, old, oldGraph, before: adapter.snapshot()! }; +} +beforeEach(() => { + vi.clearAllMocks(); + ort.create.mockReset(); +}); +describe('PhysicsAdapter training transaction', () => { + it.each(['wrong graph', 'ORT initialization', 'binding'] as const)( + '%s失败保留原session/workspace/策略及物理状态', + async (failure) => { + const { adapter, source, old, oldGraph, before } = await existing(); + const workspace = adapter.workspace, + dispose = vi.spyOn(old, 'dispose'); + const candidate = graph(47); + if (failure === 'ORT initialization') + ort.create.mockRejectedValueOnce(new Error('ORT init failed')); + else ort.create.mockResolvedValueOnce(candidate); + if (failure === 'binding') { + const text = new TextDecoder() + .decode(source.files[0].data) + .replace('name="FL_hip_joint_motor"', 'name="incompatible_name"'); + source.files[0].data = new TextEncoder().encode(text); + } + try { + await expect( + adapter.load(source, 'go2.xml', { + trainingDeployment: deployment, + trainingPolicy: { data: new Uint8Array([2]), path: 'candidate.onnx' }, + }), + ).rejects.toThrow( + failure === 'wrong graph' ? /维度/ : failure === 'binding' ? /驱动器/ : /ORT init/, + ); + expect(adapter.session).toBe(old); + expect(adapter.workspace).toBe(workspace); + expect(adapter.snapshot()).toEqual(before); + expect(dispose).not.toHaveBeenCalled(); + expect(oldGraph.release).not.toHaveBeenCalled(); + if (failure === 'wrong graph') expect(candidate.release).toHaveBeenCalledOnce(); + } finally { + adapter.dispose(); + } + }, + ); + it('完成ORT初始化前不发布候选,成功后仍保留旧资源以供viewer失败回滚', async () => { + const { adapter, source, old, oldGraph, before } = await existing(); + let finish!: (value: ReturnType) => void; + let started!: () => void; + const initialized = new Promise((resolve) => { + started = resolve; + }); + ort.create.mockImplementationOnce(() => { + started(); + return new Promise((resolve) => { + finish = resolve; + }); + }); + const candidate = graph(81); + try { + const loading = adapter.load(source, 'go2.xml', { + trainingDeployment: deployment, + trainingPolicy: { data: new Uint8Array([2]), path: 'new.onnx' }, + }); + await initialized; + expect(adapter.session).toBe(old); + expect(oldGraph.release).not.toHaveBeenCalled(); + finish(candidate); + await loading; + expect(adapter.session).not.toBe(old); + expect(adapter.snapshot()?.rlPolicy).toMatchObject({ + path: 'new.onnx', + observationSize: 81, + enabled: true, + }); + expect(oldGraph.release).not.toHaveBeenCalled(); + adapter.rollbackRetired(); + expect(adapter.session).toBe(old); + expect(adapter.snapshot()).toEqual(before); + for (let i = 0; i < 10; i++) await Promise.resolve(); + expect(candidate.release).toHaveBeenCalledOnce(); + expect(oldGraph.release).not.toHaveBeenCalled(); + } finally { + adapter.dispose(); + } + }); +}); + +it('候选ORT等待期间另一场景完成加载,迟到候选不得覆盖最新场景', async () => { + const { adapter, source } = await existing(); + let finish!: (value: ReturnType) => void; + let started!: () => void; + const initialized = new Promise((resolve) => { + started = resolve; + }); + ort.create.mockImplementationOnce(() => { + started(); + return new Promise((resolve) => { + finish = resolve; + }); + }); + const candidate = graph(81); + try { + const stale = adapter.load(source, 'go2.xml', { + trainingDeployment: deployment, + trainingPolicy: { data: new Uint8Array([2]), path: 'stale.onnx' }, + }); + await initialized; + await adapter.load(source, 'go2.xml'); + const latest = adapter.session, + before = adapter.snapshot(); + const rejected = expect(stale).rejects.toThrow(/取消/); + finish(candidate); + await rejected; + expect(adapter.session).toBe(latest); + expect(adapter.snapshot()).toEqual(before); + for (let i = 0; i < 10; i++) await Promise.resolve(); + expect(candidate.release).toHaveBeenCalledOnce(); + } finally { + adapter.dispose(); + } +}); + +describe('PhysicsAdapter custom_boxes applied scene compiler', () => { + it('真实WASM读取工程路径/嵌套body变换/旋转多实例;与共享CPU布局halfsize和原点一致', async () => { + const { default: shared } = + await import('../../../training_server/tests/fixtures/custom-boxes.json'); + const source = project(); + const descriptor = { + schemaVersion: 1, + id: 'warehouse', + name: '仓库', + coordinateSystem: { units: 'm', up: 'Z', forward: '+X' }, + physics: { source: 'physics/world.xml' }, + spawnPoints: [], + }; + for (const [path, text] of [ + ['maps/warehouse/map.json', JSON.stringify(descriptor)], + [ + 'maps/warehouse/physics/world.xml', + '', + ], + ]) { + const data = new TextEncoder().encode(text); + source.files.push({ path, data, size: data.length, source: 'file', mimeType: 'text/xml' }); + } + const assets = [ + { + id: 'one', + name: 'one', + selection: { + kind: 'project' as const, + descriptorPath: 'maps/warehouse/map.json', + positionX: 1, + positionY: 1, + }, + }, + ]; + const adapter = new MainThreadPhysicsAdapter(); + try { + await adapter.load(source, 'go2.xml', { mapAssets: assets }); + const initial = adapter.exportTrainingTerrain(assets, { spawn: [-2, -1], target: [2, -1] }); + // Geometry from actual WASM matches the JSON consumed by CPU TerrainGenerator tests. + expect(initial.boxes[1]).toEqual(shared.boxes[1]); + expect(initial.actualObstacleCount).toBe(1); + adapter.session!.data.qpos[0] = 8; + expect(adapter.exportTrainingTerrain(assets).spawn[0]).not.toBe(8); + expect(() => adapter.exportTrainingTerrain([])).toThrow(/过时/); + const multi = [ + ...assets, + { + id: 'two', + name: 'two', + selection: { ...assets[0].selection, positionX: -2, positionY: 1, yawDeg: 45 }, + }, + ]; + await adapter.load(source, 'go2.xml', { mapAssets: multi }); + const layout = adapter.exportTrainingTerrain(multi, { spawn: [-2, -1], target: [2, -1] }); + expect(layout.boxes).toHaveLength(3); + expect(layout.boxes[2].pos[0]).toBeCloseTo(-2 - Math.SQRT1_2, 12); + expect(layout.boxes[2].pos[1]).toBeCloseTo(1 + Math.SQRT1_2, 12); + expect(layout.boxes[2].size[0]).toBeCloseTo(0.7 * Math.SQRT1_2, 12); + expect(layout.boxes[2].size[1]).toBeCloseTo(0.7 * Math.SQRT1_2, 12); + adapter.rollbackRetired(); + expect(adapter.exportTrainingTerrain(assets).actualObstacleCount).toBe(1); + // A deployment replacement cannot pretend to be the original applied editor map. + await adapter.load(source, 'go2.xml', { mapAssets: assets, trainingDeployment: deployment }); + expect(() => adapter.exportTrainingTerrain(assets)).toThrow(/已应用/); + } finally { + adapter.dispose(); + } + }); +}); diff --git a/web_platform/src/simulation/PhysicsAdapter.ts b/web_platform/src/simulation/PhysicsAdapter.ts index d706d3bb..7ca5f702 100644 --- a/web_platform/src/simulation/PhysicsAdapter.ts +++ b/web_platform/src/simulation/PhysicsAdapter.ts @@ -1,3 +1,10 @@ +import { + composeTrainingMap, + trainingTerrainFromCompiledScene, + type CompiledMapGeometry, + type TrainingSceneCoordinates, +} from '../map/trainingMap'; +import type { PolicyDeployment, TrainingTerrain } from '../rl/deployment'; import type { MainModule } from '@mujoco/mujoco'; import type { ProjectFile, ProjectManifest } from '../project/types'; import { prepareProjectForMujoco } from '../project/importer'; @@ -36,6 +43,9 @@ export interface PhysicsLoadProgress { } export interface PhysicsLoadOptions { + trainingDeployment?: PolicyDeployment; + /** Candidate policy is initialized/validated before replacing the active session. */ + trainingPolicy?: { data: Uint8Array; path: string }; urdfMode?: UrdfLoadMode; baseMode?: UrdfBaseMode; enhancements?: UrdfEnhancementOptions; @@ -68,9 +78,15 @@ export interface PhysicsAdapter { setControllerEnabled(enabled: boolean): void; sendControllerCommand(command: ControllerCommand): void; removeController(): void; - loadRLPolicy(model: Uint8Array, path: string): Promise; + loadRLPolicy( + model: Uint8Array, + path: string, + deployment?: PolicyDeployment, + ): Promise; setRLPolicyEnabled(enabled: boolean): void; setRLCommand(command: RLCommand): void; + setNavigationTarget(target: [number, number]): void; + resetNavigationTarget(): void; removeRLPolicy(): void; configureDataRecorder(config: Partial): DataRecorderStatus | undefined; startDataRecording(): DataRecorderStatus | undefined; @@ -82,6 +98,10 @@ export interface PhysicsAdapter { releaseRetired(): void; rollbackRetired(): void; exportMjcf(): Uint8Array; + exportTrainingTerrain( + assets: readonly PlacedMapAsset[], + coordinates?: TrainingSceneCoordinates, + ): TrainingTerrain; dispose(): void; } @@ -108,6 +128,15 @@ export class MainThreadPhysicsAdapter implements PhysicsAdapter { session: SimulationSession | null = null; workspace: MemfsWorkspace | null = null; private supportFiles: ProjectFile[] = []; + private trainingScenes = new WeakMap< + SimulationSession, + { + assets: string; + geometries: CompiledMapGeometry[]; + initialPose: number[]; + extent: number; + } + >(); private loadGeneration = 0; private disposed = false; private retiredSession: SimulationSession | null = null; @@ -196,7 +225,7 @@ export class MainThreadPhysicsAdapter implements PhysicsAdapter { intermediate.dispose(); } } - if (hasMaps) { + if (hasMaps && !options.trainingDeployment?.terrain) { report(0.72, '组合机器人与物理地图'); if (entry?.format === 'urdf' && urdfMode === 'native') throw new Error('原生 URDF 模式暂不支持地图,请切换为转换模式'); @@ -232,13 +261,102 @@ export class MainThreadPhysicsAdapter implements PhysicsAdapter { modelRelativePath = mapPath; modelPath = workspace.path(modelRelativePath); } + if (options.trainingDeployment?.terrain) { + if (entry?.format === 'urdf' && urdfMode === 'native') + throw new Error('训练地图需要URDF转换模式'); + const intermediate = new SimulationSession(module, modelPath); + try { + const flatPath = `${modelRelativePath}.training.xml`; + if (!module.mj_saveLastXML(workspace.path(flatPath), intermediate.model)) + throw new Error('无法展开训练场景'); + workspace.writeGenerated( + flatPath, + composeTrainingMap( + new TextEncoder().encode(workspace.readText(flatPath)), + options.trainingDeployment, + ), + ); + modelPath = workspace.path(flatPath); + } finally { + intermediate.dispose(); + } + warnings.push('已用策略配套训练布局替换场景地形;编辑器地图未更改。Go2-W动力学不等同Go2。'); + } report(0.84, '编译模型与物理数据'); console.info('[MuJoCo] 编译模型', modelPath); nextSession = new SimulationSession(module, modelPath, warnings); + if (options.trainingDeployment) nextSession.configureDeployment(options.trainingDeployment); if (entry?.format === 'urdf' && urdfMode === 'native') { const offset = nextSession.alignLowestPointToGround(); warnings.push(`原生 URDF 已整体平移 ${offset.toFixed(4)} m,使最低点位于 z=0`); } + if (placedMapAssets?.length && !options.trainingDeployment) { + const sanitize = (id: string) => id.replace(/[^a-zA-Z0-9_-]/g, '_'); + const prefixes = placedMapAssets.map((asset) => + asset.selection.kind === 'builtin' + ? `__platform_map_${sanitize(asset.id)}__` + : `__platform_map_${sanitize(`${resolveProjectMap(prepared.manifest, asset.selection.descriptorPath).definition.id}_${asset.id}`)}_`, + ); + const model = nextSession.model, + data = nextSession.data; + const geometries: CompiledMapGeometry[] = []; + for (let id = 0; id < model.ngeom; id++) { + const geom = model.geom(id); + try { + const name = geom.name; + geometries.push({ + name, + type: Number(model.geom_type[id]), + position: Array.from(data.geom_xpos.slice(id * 3, id * 3 + 3)), + rotation: Array.from(data.geom_xmat.slice(id * 9, id * 9 + 9)), + size: Array.from(model.geom_size.slice(id * 3, id * 3 + 3)), + friction: Array.from(model.geom_friction.slice(id * 3, id * 3 + 3)), + collision: Boolean(model.geom_contype[id] || model.geom_conaffinity[id]), + static: Number(model.body_weldid[Number(model.geom_bodyid[id])]) === 0, + map: + prefixes.some((prefix) => name.startsWith(prefix)) || + name === '__platform_map_ground__', + }); + } finally { + geom.delete(); + } + } + const roots = Array.from({ length: model.njnt }, (_, id) => id).filter( + (id) => Number(model.jnt_type[id]) === 0, + ); + if (roots.length === 1) { + const address = Number(model.jnt_qposadr[roots[0]]); + this.trainingScenes.set(nextSession, { + assets: JSON.stringify(placedMapAssets), + geometries, + initialPose: Array.from(data.qpos.slice(address, address + 7)), + extent: Math.max( + 4, + ...placedMapAssets.map((asset) => + asset.selection.kind === 'builtin' + ? Math.max( + Math.abs(asset.selection.config.positionX), + Math.abs(asset.selection.config.positionY), + ) + + asset.selection.config.size / 2 + : 0, + ), + ), + }); + } + } + if (options.trainingPolicy) { + if (!options.trainingDeployment?.terrain) throw new Error('事务策略加载需要配套训练地图'); + report(0.9, '校验候选场景的 ONNX 策略'); + await nextSession.loadRLPolicy( + options.trainingPolicy.data, + options.trainingPolicy.path, + options.trainingDeployment, + ); + // Configure/reset already supplied the correct spawn. Enable while still paused; + // the candidate cannot advance until the viewer transaction commits. + nextSession.setRLPolicyEnabled(true); + } if (warnings.length) console.info('[MuJoCo] 兼容与地图处理', warnings); console.info('[MuJoCo] 模型编译完成'); report(0.94, '生成初始仿真状态'); @@ -263,6 +381,20 @@ export class MainThreadPhysicsAdapter implements PhysicsAdapter { ); } } + exportTrainingTerrain( + assets: readonly PlacedMapAsset[], + coordinates?: TrainingSceneCoordinates, + ): TrainingTerrain { + const scene = this.session && this.trainingScenes.get(this.session); + if (!scene || scene.assets !== JSON.stringify(assets)) + throw new Error('没有匹配的已应用碰撞场景/唯一浮动机器人,或场景已过时;请先应用地图'); + return trainingTerrainFromCompiledScene( + scene.geometries, + scene.initialPose, + scene.extent, + coordinates, + ); + } advance(now: number): FrameResult { return this.session?.advance(now) ?? { steps: 0, stepMs: 0, overBudget: false }; } @@ -315,13 +447,23 @@ export class MainThreadPhysicsAdapter implements PhysicsAdapter { removeController(): void { this.session?.removeController(); } - async loadRLPolicy(model: Uint8Array, path: string): Promise { + async loadRLPolicy( + model: Uint8Array, + path: string, + deployment?: PolicyDeployment, + ): Promise { if (!this.session) throw new Error('请先加载模型'); - return this.session.loadRLPolicy(model, path); + return this.session.loadRLPolicy(model, path, deployment); } setRLPolicyEnabled(enabled: boolean): void { this.session?.setRLPolicyEnabled(enabled); } + setNavigationTarget(target: [number, number]): void { + this.session?.setNavigationTarget(target); + } + resetNavigationTarget(): void { + this.session?.resetNavigationTarget(); + } setRLCommand(command: RLCommand): void { this.session?.setRLCommand(command); } diff --git a/web_platform/src/simulation/SimulationSession.policyLifecycle.test.ts b/web_platform/src/simulation/SimulationSession.policyLifecycle.test.ts new file mode 100644 index 00000000..ce78d360 --- /dev/null +++ b/web_platform/src/simulation/SimulationSession.policyLifecycle.test.ts @@ -0,0 +1,71 @@ +import { describe, it, expect, vi } from 'vitest'; +const pending = vi.hoisted(() => ({ + resolvers: [] as (() => void)[], + runtimes: [] as { dispose: ReturnType; status: ReturnType }[], +})); +vi.mock('../rl/runtime/Go2wPolicyBindings', () => ({ + Go2wPolicyBindings: class { + constructor( + _model: unknown, + _data: unknown, + private write: (id: number, value: number) => void, + ) {} + clear() { + this.write(0, 0); + } + }, +})); +vi.mock('../rl/runtime/OnnxPolicyRuntime', () => ({ + OnnxPolicyRuntime: { + load: (_model: unknown, _path: string, bindings: { clear(): void }) => + new Promise((resolve) => { + const runtime = { + dispose: vi.fn(() => bindings.clear()), + status: vi.fn(() => ({ loaded: true })), + }; + pending.runtimes.push(runtime); + pending.resolvers.push(() => resolve(runtime)); + }), + }, +})); +import { SimulationSession } from './SimulationSession'; +async function flush() { + for (let i = 0; i < 20; i++) await Promise.resolve(); +} +function session(): SimulationSession { + return Object.assign(Object.create(SimulationSession.prototype) as object, { + model: {}, + data: { ctrl: new Float64Array(1) }, + rlPolicyLoadGeneration: 0, + disposed: false, + setActuator: vi.fn(), + }) as unknown as SimulationSession; +} +describe('SimulationSession policy loading race', () => { + it('模型释放后迟到ORT session必须release,不能写已释放MjData', async () => { + pending.resolvers.length = pending.runtimes.length = 0; + const s = session(); + const load = s.loadRLPolicy(new Uint8Array(), 'late.onnx'); + await flush(); + Object.assign(s, { disposed: true }); + pending.resolvers[0](); + await expect(load).rejects.toThrow(/取消/); + expect(pending.runtimes[0].dispose).toHaveBeenCalledOnce(); + expect(s.setActuator).not.toHaveBeenCalled(); + }); + it('同一模型两次加载只接受最新一份,旧clear不能清新策略动作', async () => { + pending.resolvers.length = pending.runtimes.length = 0; + const s = session(); + const first = s.loadRLPolicy(new Uint8Array(), 'first.onnx'); + await flush(); + const second = s.loadRLPolicy(new Uint8Array(), 'second.onnx'); + await flush(); + pending.resolvers[1](); + await expect(second).resolves.toMatchObject({ loaded: true }); + pending.resolvers[0](); + await expect(first).rejects.toThrow(/取消/); + expect(pending.runtimes[0].dispose).toHaveBeenCalledOnce(); + expect(pending.runtimes[1].dispose).not.toHaveBeenCalled(); + expect(s.setActuator).not.toHaveBeenCalled(); + }); +}); diff --git a/web_platform/src/simulation/SimulationSession.ts b/web_platform/src/simulation/SimulationSession.ts index 746f4104..5e71a4dd 100644 --- a/web_platform/src/simulation/SimulationSession.ts +++ b/web_platform/src/simulation/SimulationSession.ts @@ -1,3 +1,10 @@ +import { Go2ObstacleAvoidanceBindings } from '../rl/runtime/Go2ObstacleAvoidanceBindings'; +import { + OBSTACLE_TASK_ID, + policyDeploymentsMatch, + validatePolicyDeployment, + type PolicyDeployment, +} from '../rl/deployment'; import type { MainModule, MjData, MjModel, MjvPerturb, MjvScene } from '@mujoco/mujoco'; import { meshIdFromSceneDataId } from './geometry'; import { PythonControllerRuntime } from '../controller/PythonControllerRuntime'; @@ -106,6 +113,9 @@ export class SimulationSession { private pythonController?: PythonControllerRuntime; private controllerLoadGeneration = 0; private rlPolicy?: OnnxPolicyRuntime; + private obstacleBindings?: Go2ObstacleAvoidanceBindings; + private deployment?: PolicyDeployment; + private deploymentInitialQpos?: Float64Array; private rlPolicyLoadGeneration = 0; private dataRecorder!: DataRecorder; @@ -189,6 +199,7 @@ export class SimulationSession { reset(): void { this.setPaused(true); this.module.mj_resetData(this.model, this.data); + if (this.deploymentInitialQpos) this.data.qpos.set(this.deploymentInitialQpos); this.module.mj_forward(this.model, this.data); this.clearExternalForce(); this.data.ctrl.fill(0); @@ -253,20 +264,132 @@ export class SimulationSession { if (!enabled) this.data.ctrl.fill(0); } - async loadRLPolicy(model: Uint8Array, path: string): Promise { + navigationStatus() { + return this.rlPolicy?.navigationStatus(); + } + setNavigationTarget(target: [number, number]): void { + this.rlPolicy?.setNavigationTarget(target); + } + resetNavigationTarget(): void { + this.rlPolicy?.resetNavigationTarget(); + } + + perceptionRays() { + return this.obstacleBindings?.rays ?? []; + } + + configureDeployment(input: PolicyDeployment): void { + const d = validatePolicyDeployment(input); + if (!d.terrain) return; + const seen = new Set(); + for (let id = 0; id < this.model.ngeom; id++) { + const geom = this.model.geom(id); + try { + if (Number(this.model.geom_bodyid[id]) !== 0) { + if (geom.name.startsWith('__training_terrain_')) + throw new Error('机器人几何不能使用训练地图保留名称'); + continue; + } + const match = /^__training_terrain_(\d+)$/.exec(geom.name); + const index = match ? Number(match[1]) : -1; + const box = d.terrain.boxes[index]; + if ( + !box || + seen.has(index) || + Number(this.model.geom_type[id]) !== 6 || + Number(this.model.geom_group[id]) !== 2 + ) + throw new Error('配套地图只支持声明的静态box,不允许额外plane或其他形状'); + seen.add(index); + for (let axis = 0; axis < 3; axis++) { + if ( + Math.abs(Number(this.data.geom_xpos[id * 3 + axis]) - box.pos[axis]) > 1e-6 || + Math.abs(Number(this.model.geom_size[id * 3 + axis]) - box.size[axis]) > 1e-6 + ) + throw new Error('编译后的训练地图坐标/半尺寸不匹配'); + } + } finally { + geom.delete(); + } + } + if (seen.size !== d.terrain.boxes.length) throw new Error('配套训练地图缺少box'); + // Construction validates robot topology before changing the initial pose. + new Go2wPolicyBindings(this.model, this.data, (id, value) => this.setActuator(id, value)); + const initialQpos = Float64Array.from(this.model.qpos0); + let freeCount = 0; + for (let id = 0; id < this.model.njnt; id++) { + const joint = this.model.jnt(id); + try { + const address = Number(joint.qposadr); + if (Number(this.model.jnt_type[id]) === 0) { + freeCount++; + initialQpos.set([...d.terrain.spawn, ...d.terrain.spawnQuaternion], address); + } else { + const index = d.jointNames.indexOf(joint.name); + if (index >= 0) initialQpos[address] = d.defaultJointPosition[index]; + } + } finally { + joint.delete(); + } + } + if (freeCount !== 1) throw new Error('训练评测需要唯一浮动基座'); + for (let id = 0; id < this.model.nactuator; id++) { + const actuator = this.model.actuator(id); + try { + const jointId = Number(actuator.trnid[0]); + if (jointId < 0 || jointId >= this.model.njnt) continue; + const joint = this.model.jnt(jointId); + try { + const index = d.jointNames.indexOf(joint.name); + if (index >= 0) { + actuator.forcelimited = 1; + actuator.forcerange[0] = -d.effortLimits[index]; + actuator.forcerange[1] = d.effortLimits[index]; + } + } finally { + joint.delete(); + } + } finally { + actuator.delete(); + } + } + this.deployment = d; + // qpos0 contains hinge reference angles, not the desired joint pose. Do not mutate it. + this.deploymentInitialQpos = initialQpos; + this.reset(); + } + + async loadRLPolicy( + model: Uint8Array, + path: string, + deployment?: PolicyDeployment, + ): Promise { const generation = ++this.rlPolicyLoadGeneration; - const bindings = new Go2wPolicyBindings(this.model, this.data, (id, value) => - this.setActuator(id, value), - ); + if (deployment) deployment = validatePolicyDeployment(deployment); + if ( + deployment?.terrain && + (!this.deployment || !policyDeploymentsMatch(deployment, this.deployment)) + ) + throw new Error('策略配套地图尚未加载或配置不匹配'); + let cancelled = false; + const writeActuator = (id: number, value: number) => { + if (!this.disposed && !cancelled) this.setActuator(id, value); + }; + const bindings = + deployment?.taskId === OBSTACLE_TASK_ID + ? new Go2ObstacleAvoidanceBindings(this.model, this.data, writeActuator, deployment) + : new Go2wPolicyBindings(this.model, this.data, writeActuator, deployment?.effortLimits); const { OnnxPolicyRuntime: Runtime } = await import('../rl/runtime/OnnxPolicyRuntime'); - const runtime = await Runtime.load(model, path, bindings); + const runtime = await Runtime.load(model, path, bindings, deployment); if (this.disposed || generation !== this.rlPolicyLoadGeneration) { + cancelled = true; runtime.dispose(); throw new Error('模型已切换,ONNX 策略加载已取消'); } this.data.ctrl.fill(0); this.rlPolicy?.dispose(); this.rlPolicy = runtime; + this.obstacleBindings = bindings instanceof Go2ObstacleAvoidanceBindings ? bindings : undefined; return runtime.status(); } @@ -306,6 +429,7 @@ export class SimulationSession { this.rlPolicyLoadGeneration += 1; this.rlPolicy?.dispose(); this.rlPolicy = undefined; + this.obstacleBindings = undefined; this.data.ctrl.fill(0); } diff --git a/web_platform/src/training/LocalTrainingClient.test.ts b/web_platform/src/training/LocalTrainingClient.test.ts index b2b87da1..36701230 100644 --- a/web_platform/src/training/LocalTrainingClient.test.ts +++ b/web_platform/src/training/LocalTrainingClient.test.ts @@ -148,3 +148,66 @@ describe('LocalTrainingClient', () => { expect(() => new LocalTrainingClient('http://localhost:8765', '')).toThrow('访问令牌'); }); }); + +it('LocalTrainingClient完整透传自定义任务与嵌套地形/传感器参数', async () => { + const fetchMock = vi.fn().mockResolvedValue(Response.json({ id: 'job' })); + vi.stubGlobal('fetch', fetchMock); + const request = { + taskId: 'Unitree-Go2-ObstacleAvoidance', + numEnvs: 16, + maxIterations: 1, + seed: 42, + runName: 'avoid', + device: 'gpu' as const, + gpuIds: [0], + wandbMode: 'disabled' as const, + terrainPreset: 'discrete_obstacles', + terrainParams: { friction: 0.9, obstacle_count: 12 }, + sensorType: 'raycast' as const, + sensorCfg: { fov: 60, maxDistance: 3 }, + }; + await new LocalTrainingClient('http://127.0.0.1:8765', 'test').start(request); + expect(JSON.parse(String(fetchMock.mock.calls[0][1].body))).toEqual(request); +}); + +it('上传使用File原始body、显式模板、Bearer及AbortSignal,不发送JSON/base64/服务器路径', async () => { + const fetchMock = vi + .fn() + .mockResolvedValue(new Response(JSON.stringify({ id: 'source' }), { status: 201 })); + vi.stubGlobal('fetch', fetchMock); + const file = new File(['weights'], '单个 actor.ONNX'); + const controller = new AbortController(); + await new LocalTrainingClient('http://localhost:8765', 'secret-token').uploadPretrained( + file, + 'go2-legacy47-v1', + controller.signal, + ); + const [url, init] = fetchMock.mock.calls[0] as [string, RequestInit]; + const query = new URL(url).searchParams; + expect(query.get('format')).toBe('onnx'); + expect(query.get('template')).toBe('go2-legacy47-v1'); + expect(query.get('name')).toBe(file.name); + expect(url).not.toContain('secret-token'); + expect(init.body).toBe(file); + expect(init.signal).toBe(controller.signal); + expect(new Headers(init.headers).get('Content-Type')).toBe('application/octet-stream'); + expect(new Headers(init.headers).get('Authorization')).toBe('Bearer secret-token'); + expect(new Headers(init.headers).has('Content-Length')).toBe(false); // Browser owns this forbidden header. +}); + +it('上传前拒绝ZIP、空文件和超过各格式上限,仍由服务验证实际模型', () => { + const fetchMock = vi.fn(); + vi.stubGlobal('fetch', fetchMock); + const client = new LocalTrainingClient('http://localhost:8765', 'token'); + for (const [name, size] of [ + ['model.zip', 1], + ['model.pt', 0], + ['model.pt', 256 * 1024 ** 2 + 1], + ['model.onnx', 64 * 1024 ** 2 + 1], + ] as const) { + const file = new File(['x'], name); + Object.defineProperty(file, 'size', { value: size }); + expect(() => client.uploadPretrained(file, 'go2-legacy47-v1')).toThrow(); + } + expect(fetchMock).not.toHaveBeenCalled(); +}); diff --git a/web_platform/src/training/LocalTrainingClient.ts b/web_platform/src/training/LocalTrainingClient.ts index 336950bd..7d8c8469 100644 --- a/web_platform/src/training/LocalTrainingClient.ts +++ b/web_platform/src/training/LocalTrainingClient.ts @@ -1,5 +1,6 @@ import type { ParameterConstraint, + PretrainedSource, RewardPreset, TuningCapability, TuningCreateRequest, @@ -57,6 +58,26 @@ export class LocalTrainingClient { health(): Promise { return this.json('/api/training/health'); } + uploadPretrained( + file: File, + template: 'go2-legacy47-v1', + signal?: AbortSignal, + ): Promise { + const format = file.name.split('.').pop()?.toLowerCase(); + if (format !== 'pt' && format !== 'onnx') + throw new Error('请选择单个.pt或.onnx文件,不支持ZIP/目录'); + const limit = (format === 'pt' ? 256 : 64) * 1024 ** 2; + if (!file.size || file.size > limit) + throw new Error(`文件不能为空,${format}上限为${limit / 1024 ** 2}MiB`); + if (template !== 'go2-legacy47-v1') throw new Error('请先确认Go2 legacy47模板'); + const query = new URLSearchParams({ format, template, name: file.name }); + return this.json(`/api/training/pretrained-sources/upload?${query}`, { + method: 'POST', + headers: { 'Content-Type': 'application/octet-stream' }, + body: file, + signal, + }); + } start(request: TrainingRequest): Promise { return this.json('/api/training/jobs', { method: 'POST', diff --git a/web_platform/src/training/LocalTrainingPanel.test.tsx b/web_platform/src/training/LocalTrainingPanel.test.tsx index ec0e4542..5ee27c8d 100644 --- a/web_platform/src/training/LocalTrainingPanel.test.tsx +++ b/web_platform/src/training/LocalTrainingPanel.test.tsx @@ -1,3 +1,8 @@ +import customLayout from '../../../training_server/tests/fixtures/custom-boxes.json'; +import { validateCustomTerrain, validatePolicyDeployment } from '../rl/deployment'; +import fixture from '../rl/fixtures/obstacleDeployment.json'; +import { DEFAULT_PHYSICAL_MAP_CONFIG } from '../map/types'; +import { TRAINING_JOB_KEY } from './storage'; import { fireEvent, render, screen, waitFor } from '@testing-library/react'; import { beforeEach, describe, expect, it, vi } from 'vitest'; import { LocalTrainingPanel } from './LocalTrainingPanel'; @@ -161,3 +166,308 @@ describe('LocalTrainingPanel', () => { ); }); }); + +const customTasks = ['Unitree-Go2-Flat', 'Unitree-Go2-ObstacleAvoidance', 'Unitree-Go2-Rough']; +const customMetadata = customTasks.map((id) => ({ + id, + name: id, + browserCompatible: id !== 'Unitree-Go2-Rough', + terrainPresets: [ + 'plane', + 'discrete_obstacles', + 'rough', + 'pyramid_stairs', + 'wave', + 'custom_boxes', + ], + terrainParameters: { + size: { min: 8, max: 24, default: 12 }, + friction: { min: 0.2, max: 2, default: 0.8 }, + obstacle_count: { min: 1, max: 100, default: 24, integer: true }, + }, + sensorTypes: id.includes('Obstacle') ? ['raycast'] : [], + sensorModes: ['single_ring_raycast', 'multi_ring_raycast'], + sensorParameters: id.includes('Obstacle') + ? { fov: { min: 30, max: 120, default: 90 }, maxDistance: { min: 1, max: 5, default: 4 } } + : {}, + mapSyncScope: '单块预设参数', +})); +function customServer(job?: unknown) { + const fetchMock = vi.fn(async (url: string, init?: RequestInit) => { + if (url.endsWith('/health')) + return Response.json({ + ready: true, + tasks: customTasks, + taskMetadata: customMetadata, + trainerRoot: '/local', + }); + if (url.includes('/presets')) return Response.json({ presets: [] }); + if (url.endsWith('/policy.onnx')) return new Response(new Uint8Array([8, 9])); + if (init?.method === 'POST' || job) + return Response.json( + job ?? { + id: 'c'.repeat(32), + taskId: 'Unitree-Go2-ObstacleAvoidance', + state: 'queued', + progress: 0, + iteration: 0, + maxIterations: 2, + message: '', + logs: [], + artifactReady: false, + }, + ); + return Response.json({ error: 'not found' }, { status: 404 }); + }); + vi.stubGlobal('fetch', fetchMock); + return fetchMock; +} +async function connectCustom() { + fireEvent.change(screen.getByLabelText('训练服务访问令牌'), { target: { value: 'token' } }); + fireEvent.click(screen.getByRole('button', { name: '连接' })); + await screen.findByText('/local'); +} +describe('LocalTrainingPanel 自定义任务', () => { + it('动态任务/地图/传感器配置传入payload,切换任务清掉sensor与reward状态', async () => { + const fetchMock = customServer(); + render(); + await connectCustom(); + fireEvent.change(screen.getByLabelText('训练任务'), { + target: { value: 'Unitree-Go2-ObstacleAvoidance' }, + }); + expect(screen.getByLabelText('训练地形')).toHaveValue('discrete_obstacles'); + expect(screen.getByLabelText('奖励配置 preset')).toBeDisabled(); + fireEvent.change(screen.getByLabelText('感知角 FOV'), { target: { value: '60' } }); + fireEvent.change(screen.getByLabelText('训练地形'), { target: { value: 'rough' } }); + expect(screen.getByText(/训练专用 box/)).toBeInTheDocument(); + fireEvent.change(screen.getByLabelText('训练任务'), { target: { value: 'Unitree-Go2-Flat' } }); + expect(screen.queryByLabelText('感知角 FOV')).not.toBeInTheDocument(); + expect(screen.getByLabelText('训练地形')).toHaveValue(''); + fireEvent.change(screen.getByLabelText('训练任务'), { + target: { value: 'Unitree-Go2-ObstacleAvoidance' }, + }); + expect(screen.getByLabelText('感知角 FOV')).toHaveValue(90); + fireEvent.change(screen.getByLabelText('感知角 FOV'), { target: { value: '60' } }); + fireEvent.click(screen.getByRole('button', { name: '发起本地训练' })); + await waitFor(() => + expect(fetchMock.mock.calls.some((call) => call[1]?.method === 'POST')).toBe(true), + ); + const payload = JSON.parse( + String(fetchMock.mock.calls.find((call) => call[1]?.method === 'POST')?.[1]?.body), + ); + expect(payload).toMatchObject({ + taskId: 'Unitree-Go2-ObstacleAvoidance', + terrainPreset: 'discrete_obstacles', + sensorType: 'raycast', + sensorCfg: { fov: 60 }, + }); + expect(payload.rewardPresetId).toBeUndefined(); + }); + it('无地图/多实例明确报错而不提交;超限参数阻止请求', async () => { + const fetchMock = customServer(); + render(); + await connectCustom(); + fireEvent.click(screen.getByRole('button', { name: '同步当前场景地图' })); + expect(screen.getByRole('alert')).toHaveTextContent('没有已应用'); + fireEvent.change(screen.getByLabelText('训练地形'), { target: { value: 'plane' } }); + fireEvent.change(screen.getByLabelText('地图尺寸 m'), { target: { value: '100' } }); + fireEvent.click(screen.getByRole('button', { name: '发起本地训练' })); + await waitFor(() => expect(screen.getByRole('alert')).toHaveTextContent('超出允许范围')); + expect(fetchMock.mock.calls.some((call) => call[1]?.method === 'POST')).toBe(false); + }); + it('旧Rough作业禁用导入,不误当Flat', async () => { + customServer({ + id: 'r'.repeat(32), + taskId: 'Unitree-Go2-Rough', + state: 'succeeded', + artifactReady: true, + progress: 1, + iteration: 1, + maxIterations: 1, + logs: [], + message: '', + }); + localStorage.setItem(TRAINING_JOB_KEY, 'r'.repeat(32)); + render(); + await connectCustom(); + expect(await screen.findByRole('button', { name: '导入策略' })).toBeDisabled(); + expect(screen.getByText(/234维 Rough/)).toBeInTheDocument(); + }); +}); + +describe('LocalTrainingPanel 配套策略交接', () => { + it('同步权威布局并上传完整boxes,不再重跑预设种子', async () => { + const fetchMock = customServer(); + render( + validateCustomTerrain(customLayout)} + onPolicyReady={vi.fn()} + sceneMaps={[ + { + id: 'one', + name: 'one', + selection: { + kind: 'builtin', + config: { + ...DEFAULT_PHYSICAL_MAP_CONFIG, + preset: 'discrete_obstacles', + size: 16, + friction: 1.2, + seed: 123, + obstacleCount: 17, + }, + }, + }, + ]} + />, + ); + await connectCustom(); + fireEvent.click(screen.getByRole('button', { name: '同步当前场景地图' })); + expect(screen.getByLabelText('训练地形')).toHaveValue('custom_boxes'); + expect(screen.getByText('已将视口中 1 个自定义障碍物编译为训练地图布局')).toBeInTheDocument(); + expect(screen.getByLabelText('出生 X')).toHaveValue(-2); + expect(screen.getByLabelText('目标 Y')).toHaveValue(-1); + expect(screen.getByText(/旋转障碍会膨胀/)).toBeInTheDocument(); + fireEvent.click(screen.getByRole('button', { name: '发起本地训练' })); + await waitFor(() => + expect(fetchMock.mock.calls.some((call) => call[1]?.method === 'POST')).toBe(true), + ); + const request = JSON.parse( + String(fetchMock.mock.calls.find((call) => call[1]?.method === 'POST')?.[1]?.body), + ); + expect(request.customTerrainBoxes).toEqual(customLayout); + expect(request.terrainPreset).toBe('custom_boxes'); + expect(request.terrainParams).toEqual({ size: 12, friction: 0.8 }); + }); + it('等待地图/策略异步回调完成,失败显示错误,不提前解除busy', async () => { + localStorage.setItem(TRAINING_JOB_KEY, 'd'.repeat(32)); + customServer({ + id: 'd'.repeat(32), + taskId: fixture.taskId, + state: 'succeeded', + artifactReady: true, + progress: 1, + iteration: 1, + maxIterations: 1, + logs: ['Mean value loss: 0.0141'], + message: '', + deployment: fixture, + }); + let reject!: (error: Error) => void; + const ready = vi.fn( + () => + new Promise((_resolve, failure) => { + reject = failure; + }), + ); + render(); + await connectCustom(); + const button = await screen.findByRole('button', { name: '导入策略' }); + fireEvent.click(button); + await waitFor(() => expect(ready).toHaveBeenCalledOnce()); + expect(ready.mock.calls[0]).toEqual([expect.any(File), validatePolicyDeployment(fixture)]); + expect(button).toBeDisabled(); + expect(screen.getByText('价值损失')).toBeInTheDocument(); + reject(new Error('配套地图失败')); + await waitFor(() => expect(screen.getByRole('alert')).toHaveTextContent('配套地图失败')); + expect(button).toBeEnabled(); + }); +}); + +it('custom_boxes拒绝草稿、过时编译场景和未重新同步的坐标;安全区报错不删障碍', async () => { + const fetchMock = customServer(); + const assets = [ + { + id: 'one', + name: 'one', + selection: { + kind: 'builtin' as const, + config: { ...DEFAULT_PHYSICAL_MAP_CONFIG, preset: 'flat' as const }, + }, + }, + ]; + let current = validateCustomTerrain(customLayout); + const compileScene = vi.fn( + (coordinates?: import('../map/trainingMap').TrainingSceneCoordinates) => ({ + ...structuredClone(current), + ...(coordinates?.spawn ? { spawn: [...coordinates.spawn, 0.32] } : {}), + ...(coordinates?.target ? { target: [...coordinates.target] } : {}), + }), + ); + const { rerender } = render( + , + ); + await connectCustom(); + fireEvent.click(screen.getByRole('button', { name: '同步当前场景地图' })); + expect(screen.getByRole('alert')).toHaveTextContent('草稿'); + expect(compileScene).not.toHaveBeenCalled(); + rerender( + , + ); + fireEvent.click(screen.getByRole('button', { name: '同步当前场景地图' })); + fireEvent.change(screen.getByLabelText('目标 X'), { target: { value: 1 } }); + fireEvent.change(screen.getByLabelText('目标 Y'), { target: { value: 2 } }); + fireEvent.click(screen.getByRole('button', { name: '发起本地训练' })); + await waitFor(() => expect(screen.getByRole('alert')).toHaveTextContent('重新同步')); + fireEvent.click(screen.getByRole('button', { name: '同步当前场景地图' })); + expect(screen.getByRole('alert')).toHaveTextContent('安全区'); + expect(current.boxes).toHaveLength(2); + fireEvent.change(screen.getByLabelText('目标 X'), { target: { value: 2 } }); + fireEvent.change(screen.getByLabelText('目标 Y'), { target: { value: -1 } }); + fireEvent.click(screen.getByRole('button', { name: '同步当前场景地图' })); + current = { ...current, boxes: [current.boxes[0], { ...current.boxes[1], pos: [1, 3, 0.5] }] }; + fireEvent.click(screen.getByRole('button', { name: '发起本地训练' })); + await waitFor(() => expect(screen.getByRole('alert')).toHaveTextContent('过时')); + expect(fetchMock.mock.calls.some((call) => call[1]?.method === 'POST')).toBe(false); +}); + +it('显式选择multi48传入训练请求,默认与切换任务仍single32', async () => { + const fetchMock = customServer(); + render(); + await connectCustom(); + fireEvent.change(screen.getByLabelText('训练任务'), { + target: { value: 'Unitree-Go2-ObstacleAvoidance' }, + }); + expect(screen.getByLabelText('传感器模式')).toHaveValue('single_ring_raycast'); + fireEvent.change(screen.getByLabelText('传感器模式'), { + target: { value: 'multi_ring_raycast' }, + }); + fireEvent.click(screen.getByRole('button', { name: '发起本地训练' })); + await waitFor(() => expect(fetchMock.mock.calls.some((c) => c[1]?.method === 'POST')).toBe(true)); + const payload = JSON.parse( + String(fetchMock.mock.calls.find((c) => c[1]?.method === 'POST')?.[1]?.body), + ); + expect(payload.sensorCfg.sensorMode).toBe('multi_ring_raycast'); +}); + +it('Flat奖励菜单只展示服务确认归属Flat的preset,Obstacle和无身份项不展示', async () => { + const entries = [ + { id: 'a'.repeat(32), name: '合法历史Flat', taskId: 'Unitree-Go2-Flat' }, + { id: 'b'.repeat(32), name: 'Obstacle不能混入', taskId: 'Unitree-Go2-ObstacleAvoidance' }, + { id: 'c'.repeat(32), name: '缺少权威身份' }, + ]; + vi.stubGlobal( + 'fetch', + vi.fn(async (url: string) => { + if (url.endsWith('/health')) + return Response.json({ + ready: true, + tasks: customTasks, + taskMetadata: customMetadata, + trainerRoot: '/local', + }); + if (url.endsWith('/presets')) return Response.json({ presets: entries }); + return Response.json({ error: 'not found' }, { status: 404 }); + }), + ); + render(); + await connectCustom(); + expect(await screen.findByRole('option', { name: '合法历史Flat' })).toBeInTheDocument(); + expect(screen.queryByRole('option', { name: 'Obstacle不能混入' })).not.toBeInTheDocument(); + expect(screen.queryByRole('option', { name: '缺少权威身份' })).not.toBeInTheDocument(); +}); diff --git a/web_platform/src/training/LocalTrainingPanel.tsx b/web_platform/src/training/LocalTrainingPanel.tsx index 344bb8a6..ced159ba 100644 --- a/web_platform/src/training/LocalTrainingPanel.tsx +++ b/web_platform/src/training/LocalTrainingPanel.tsx @@ -1,4 +1,17 @@ -import { useEffect, useState, type ReactNode } from 'react'; +import { PretrainedIdentity, PretrainedSourceSelect } from './PretrainedSourceSelect'; +import { pretrainedSelectionError } from './pretrainedSelection'; +import { TrainingMetricsPanel } from './TrainingMetricsPanel'; +import { trainingLosses } from './trainingLosses'; +import type { PolicyDeployment, TrainingTerrain } from '../rl/deployment'; +import { + OBSTACLE_TASK_ID, + validatePolicyDeployment, + validateCustomTerrain, + readPolicyDeployment, +} from '../rl/deployment'; +import type { TrainingSceneCompiler } from '../map/trainingMap'; +import type { PlacedMapAsset } from '../map/types'; +import { useEffect, useRef, useState, type ReactNode } from 'react'; import { Download, ExternalLink, Link, Play, Server, Square } from 'lucide-react'; import { Badge, Button, ProgressBar, PropertyRow, Select } from '../components/ui'; import { LocalTrainingClient } from './LocalTrainingClient'; @@ -33,7 +46,17 @@ function stateLabel(state: TrainingJob['state']): string { }[state]; } -export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File): void }) { +export function LocalTrainingPanel({ + onPolicyReady, + compileScene, + sceneMaps = [], + sceneDirty = false, +}: { + onPolicyReady(file: File, deployment?: PolicyDeployment): void | Promise; + compileScene?: TrainingSceneCompiler; + sceneMaps?: readonly PlacedMapAsset[]; + sceneDirty?: boolean; +}) { const [endpoint, setEndpoint] = useState(() => localStored(TRAINING_ENDPOINT_KEY, DEFAULT_TRAINING_ENDPOINT), ); @@ -42,6 +65,9 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File const [job, setJob] = useState(); const [presets, setPresets] = useState([]); const [rewardPresetId, setRewardPresetId] = useState(''); + const [pretrainedSourceId, setPretrainedSourceId] = useState(''); + const [uploading, setUploading] = useState(false); + const [connectionRevision, setConnectionRevision] = useState(0); const [busy, setBusy] = useState(false), [error, setError] = useState(); const [taskId, setTaskId] = useState('Unitree-Go2-Flat'), @@ -53,15 +79,85 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File [gpuIds, setGpuIds] = useState('0'), [wandbMode, setWandbMode] = useState('offline'); + const [terrainPreset, setTerrainPreset] = useState(''); + const [customTerrainBoxes, setCustomTerrainBoxes] = useState(); + const [syncedScene, setSyncedScene] = useState(); + const [terrainParams, setTerrainParams] = useState>({}); + const [sensorMode, setSensorMode] = useState<'single_ring_raycast' | 'multi_ring_raycast'>( + 'single_ring_raycast', + ); + const [sensorCfg, setSensorCfg] = useState>({}); + const metadata = server?.taskMetadata?.find((item) => item.id === taskId); + const sourceSelectionError = pretrainedSelectionError( + server?.pretrainedSources, + taskId, + pretrainedSourceId, + ); + const selectTask = (id: string) => { + setTaskId(id); + setCustomTerrainBoxes(undefined); + setSyncedScene(undefined); + setRewardPresetId(''); + setTerrainParams({}); + setSensorCfg({}); + setSensorMode('single_ring_raycast'); + setTerrainPreset(id === OBSTACLE_TASK_ID ? 'discrete_obstacles' : ''); + }; + const syncMap = () => { + try { + if (sceneDirty) throw new Error('请先应用地图草稿,再同步训练地图'); + if (!compileScene || !sceneMaps.length) throw new Error('没有已应用的碰撞地图'); + if (!metadata?.terrainPresets.includes('custom_boxes')) + throw new Error('当前服务/任务不支持custom_boxes,请升级训练服务'); + setSyncedScene(undefined); + const layout = compileScene( + customTerrainBoxes + ? { + spawn: [customTerrainBoxes.spawn[0], customTerrainBoxes.spawn[1]], + target: [customTerrainBoxes.target[0], customTerrainBoxes.target[1]], + } + : undefined, + ); + setCustomTerrainBoxes(layout); + setTerrainPreset('custom_boxes'); + setTerrainParams({ size: layout.size, friction: layout.friction }); + validateCustomTerrain(layout); + setSyncedScene(JSON.stringify(sceneMaps)); + setError(undefined); + } catch (value) { + setError(errorText(value)); + } + }; + const connectionEpoch = useRef(0); + const connected = () => { + connectionEpoch.current += 1; + setConnectionRevision(connectionEpoch.current); + setServer(undefined); + setJob(undefined); + setPresets([]); + setRewardPresetId(''); + setError(undefined); + }; + useEffect( + () => () => { + connectionEpoch.current += 1; + }, + [], + ); const connect = async () => { setBusy(true); setError(undefined); + const epoch = ++connectionEpoch.current; + setConnectionRevision(epoch); try { const client = new LocalTrainingClient(endpoint, token), info = await client.health(); + if (epoch !== connectionEpoch.current) return; setServer(info); try { - setPresets(await client.presets()); + const nextPresets = await client.presets(); + if (epoch !== connectionEpoch.current) return; + setPresets(nextPresets); } catch { setPresets([]); } @@ -70,11 +166,12 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File } catch { /* 当前会话仍可连接 */ } - if (info.tasks.length && !info.tasks.includes(taskId)) setTaskId(info.tasks[0]); + if (info.tasks.length && !info.tasks.includes(taskId)) selectTask(info.tasks[0]); const remembered = info.activeJobId ?? localStored(TRAINING_JOB_KEY); if (remembered) { try { const recovered = await client.job(remembered); + if (epoch !== connectionEpoch.current) return; setJob(recovered); try { localStorage.setItem(TRAINING_JOB_KEY, recovered.id); @@ -129,10 +226,57 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File typeof event.data !== 'object' ) return; - const data = event.data as { type?: string; sessionId?: string; policy?: unknown }; + const data = event.data as { + type?: string; + sessionId?: string; + policy?: unknown; + taskId?: string; + }; const source = event.source as Window; if (data.type === 'mujoco-tuning-ready') { - source.postMessage({ type: 'mujoco-tuning-credentials', endpoint, token }, event.origin); + try { + if (taskId === OBSTACLE_TASK_ID && terrainPreset === 'custom_boxes') { + if ( + sceneDirty || + syncedScene !== JSON.stringify(sceneMaps) || + !compileScene || + !customTerrainBoxes + ) + throw new Error('自定义地图已过时,请重新同步后打开调参'); + const current = compileScene({ + spawn: [customTerrainBoxes.spawn[0], customTerrainBoxes.spawn[1]], + target: [customTerrainBoxes.target[0], customTerrainBoxes.target[1]], + }); + if (JSON.stringify(current) !== JSON.stringify(customTerrainBoxes)) + throw new Error('碰撞场景已过时,请重新同步'); + } + source.postMessage( + { + type: 'mujoco-tuning-credentials', + endpoint, + token, + trainingContext: { + taskId: taskId === OBSTACLE_TASK_ID ? taskId : 'Unitree-Go2-Flat', + seed, + ...(pretrainedSourceId ? { pretrainedSourceId } : {}), + ...(taskId === OBSTACLE_TASK_ID + ? { + taskConfig: { + terrainPreset, + terrainParams, + sensorType: 'raycast', + sensorCfg: { ...sensorCfg, sensorMode }, + ...(terrainPreset === 'custom_boxes' ? { customTerrainBoxes } : {}), + }, + } + : {}), + }, + }, + event.origin, + ); + } catch (value) { + setError(errorText(value)); + } } if (data.type === 'mujoco-tuning-import-policy' && data.sessionId) { const reply = (ok: boolean, message?: string) => { @@ -159,7 +303,12 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File if (!(policy instanceof File) || !/\.onnx$/i.test(policy.name)) throw new Error('调参工作台返回的 ONNX 策略无效'); if (policy.size > 64 * 1024 * 1024) throw new Error('ONNX 策略不能超过 64 MiB'); - onPolicyReady(policy); + if (data.taskId === OBSTACLE_TASK_ID) { + const deployment = readPolicyDeployment(new Uint8Array(await policy.arrayBuffer())); + if (deployment?.taskId !== OBSTACLE_TASK_ID) + throw new Error('避障最佳策略缺少匹配部署契约'); + await onPolicyReady(policy, deployment); + } else await onPolicyReady(policy); reply(true); } catch (value) { const message = errorText(value); @@ -171,7 +320,23 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File }; window.addEventListener('message', receive); return () => window.removeEventListener('message', receive); - }, [endpoint, onPolicyReady, token]); + }, [ + endpoint, + pretrainedSourceId, + onPolicyReady, + token, + taskId, + seed, + terrainPreset, + terrainParams, + sensorCfg, + sensorMode, + customTerrainBoxes, + syncedScene, + sceneMaps, + sceneDirty, + compileScene, + ]); const openTuningDashboard = () => { rememberTrainingConnection(endpoint, token); @@ -179,9 +344,20 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File }; const start = async () => { + if (uploading) return; + if (sourceSelectionError) { + setError(sourceSelectionError); + return; + } setBusy(true); setError(undefined); try { + if ( + taskId === OBSTACLE_TASK_ID && + sensorMode === 'multi_ring_raycast' && + !metadata?.sensorModes?.includes(sensorMode) + ) + throw new Error('训练服务不支持multi_ring_raycast,请升级服务'); const ids = device === 'gpu' ? gpuIds @@ -191,6 +367,43 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File : []; if (ids.some((id) => !Number.isInteger(id) || id < 0)) throw new Error('GPU 编号必须是非负整数'); + for (const [values, schema] of [ + [terrainParams, metadata?.terrainParameters], + [sensorCfg, metadata?.sensorParameters], + ] as const) { + for (const [key, value] of Object.entries(values)) { + const bounds = schema?.[key]; + if ( + !bounds || + !Number.isFinite(value) || + value < bounds.min || + value > bounds.max || + (bounds.integer && !Number.isInteger(value)) + ) + throw new Error(`参数 ${key} 超出允许范围`); + } + } + if ((terrainParams.obstacle_height_min ?? 0.2) > (terrainParams.obstacle_height_max ?? 0.6)) + throw new Error('障碍物最小高度不能超过最大高度'); + if ((sensorCfg.safetyDistance ?? 0.5) >= (sensorCfg.maxDistance ?? 4)) + throw new Error('安全距离必须小于探测距离'); + if (terrainPreset === 'custom_boxes') { + if ( + sceneDirty || + !syncedScene || + syncedScene !== JSON.stringify(sceneMaps) || + !compileScene || + !customTerrainBoxes + ) + throw new Error('自定义地图未同步或场景/坐标已更改,请重新同步'); + validateCustomTerrain(customTerrainBoxes); + const current = compileScene({ + spawn: [customTerrainBoxes.spawn[0], customTerrainBoxes.spawn[1]], + target: [customTerrainBoxes.target[0], customTerrainBoxes.target[1]], + }); + if (JSON.stringify(current) !== JSON.stringify(customTerrainBoxes)) + throw new Error('已编译碰撞场景已过时,请重新同步'); + } const next = await new LocalTrainingClient(endpoint, token).start({ taskId, numEnvs, @@ -200,7 +413,13 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File device, gpuIds: ids, wandbMode, - rewardPresetId: rewardPresetId || undefined, + rewardPresetId: taskId === 'Unitree-Go2-Flat' ? rewardPresetId || undefined : undefined, + ...(pretrainedSourceId ? { pretrainedSourceId } : {}), + ...(terrainPreset ? { terrainPreset, terrainParams } : {}), + ...(terrainPreset === 'custom_boxes' ? { customTerrainBoxes } : {}), + ...(taskId === OBSTACLE_TASK_ID + ? { sensorType: 'raycast' as const, sensorCfg: { ...sensorCfg, sensorMode } } + : {}), }); setJob(next); try { @@ -231,7 +450,11 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File setBusy(true); setError(undefined); try { - onPolicyReady(await new LocalTrainingClient(endpoint, token).downloadPolicy(job.id)); + if (job.taskId !== 'Unitree-Go2-Flat' && !job.deployment) + throw new Error('该任务缺少浏览器部署契约'); + const deployment = job.deployment ? validatePolicyDeployment(job.deployment) : undefined; + const file = await new LocalTrainingClient(endpoint, token).downloadPolicy(job.id); + await onPolicyReady(file, deployment); } catch (value) { setError(errorText(value)); } finally { @@ -249,7 +472,12 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File aria-label="本地训练服务地址" className="field h-7 min-w-0 flex-1 px-2 text-xs text-text-primary" value={endpoint} - onChange={(event) => setEndpoint(event.target.value)} + disabled={busy} + onChange={(event) => { + if (busy) return; + connected(); + setEndpoint(event.target.value); + }} /> {server?.ready && !job && ( -

+
+ { + setServer( + (current) => + current && { + ...current, + pretrainedSources: [ + ...(current.pretrainedSources ?? []).filter((s) => s.id !== source.id), + source, + ], + }, + ); + setPretrainedSourceId(source.id); + }, + }} + /> + {metadata && ( + <> + + + + +

+ 从全部已应用实例的实际碰撞几何编译世界AABB;旋转障碍会膨胀,底板标准化为z=[-0.2,0],出生高度标准化为0.32m。仅保证训练与浏览器使用相同boxes,不等于原OBB。mesh/hfield、地下结构、混合摩擦明确拒绝。 +

+ {terrainPreset === 'custom_boxes' && customTerrainBoxes && ( + <> +
+ {(['spawn', 'target'] as const).flatMap((key) => + [0, 1].map((i) => ( + { + setSyncedScene(undefined); + setCustomTerrainBoxes( + (old) => + old && { + ...old, + [key]: old[key].map((v, j) => (i === j ? value : v)), + }, + ); + }} + /> + )), + )} +
+

+ 此处起终点仅用于固定评估和部署初始演示,训练会在同一连通自由区域内逐episode随机采样。参考点须保留0.5m圆形安全区;修改后请重新同步。 +

+ {syncedScene === JSON.stringify(sceneMaps) && !sceneDirty && ( +

+ 已将视口中 {customTerrainBoxes.actualObstacleCount}{' '} + 个自定义障碍物编译为训练地图布局 +

+ )} + + )} + {terrainPreset && terrainPreset !== 'custom_boxes' && ( +
+ {Object.entries(metadata.terrainParameters).map(([key, bounds]) => ( + setTerrainParams((old) => ({ ...old, [key]: value }))} + /> + ))} +
+ )} + {['rough', 'wave', 'pyramid_stairs'].includes(terrainPreset) && ( +

训练专用 box 离散近似布局,不等于编辑器高度场。

+ )} + {taskId === OBSTACLE_TASK_ID && ( +
+ 避障传感器高级设置 + + + + {Object.entries(metadata.sensorParameters).map(([key, bounds]) => ( + setSensorCfg((old) => ({ ...old, [key]: value }))} + /> + ))} +
+ )} + {!metadata.browserCompatible && ( +

此任务可训练,但浏览器不支持其观测契约,不能一键部署。

+ )} + + )}
@@ -387,7 +768,7 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File variant="primary" className="w-full" icon={} - disabled={busy} + disabled={busy || uploading || Boolean(sourceSelectionError)} onClick={() => void start()} > 发起本地训练 @@ -396,7 +777,7 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File 训练使用本地 mjlab 任务资产,不会把浏览器中的模型上传到网络。服务一次只运行一个训练任务。

-
+
)} {job && (
@@ -416,11 +797,26 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File {stateLabel(job.state)}
+ + {job.taskId === 'Unitree-Go2-Rough' && ( +

234维 Rough 策略仅支持后端评测,浏览器不可加载。

+ )} + {job.deployment?.terrain && ( +

+ 导入将替换当前物理地图并启动配套策略; + {job.deployment.terrain.approximation ? '训练专用近似布局' : '配套碰撞布局'} + 。请先保存场景。 +

+ )}
+ {trainingLosses(job.logs).map(({ label, value }) => ( + + ))}
+ {job.logs.length > 0 && (
最近日志 @@ -443,7 +839,12 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File ) : ( <> } + {message &&

{message}

} + {error &&

{error}

} +
+ ); +} diff --git a/web_platform/src/training/TrainingMetricHistory.test.ts b/web_platform/src/training/TrainingMetricHistory.test.ts new file mode 100644 index 00000000..239574f7 --- /dev/null +++ b/web_platform/src/training/TrainingMetricHistory.test.ts @@ -0,0 +1,63 @@ +import { TrainingMetricHistory } from './TrainingMetricHistory'; +const header = (i: number) => `Learning iteration ${i} / 1000`; + +it('解析多个迭代全部指标、重复poll去重、同迭代分批补全', () => { + const h = new TrainingMetricHistory(); + const logs = [ + header(1), + 'Mean value loss: 1e-3', + 'Mean surrogate loss: -0.1', + header(2), + 'Mean value loss: 2', + ]; + expect(h.update('a', logs)).toBe(true); + expect(h.update('a', logs)).toBe(false); + expect(h.series()[0].points.map((p) => [p.step, p.value])).toEqual([ + [1, 0.001], + [2, 2], + ]); + expect( + h.update('a', [...logs, 'Mean entropy loss: -3', 'Mean reward: 4', 'Mean episode length: 50']), + ).toBe(true); + expect(h.series()).toHaveLength(5); + expect(h.series().find((s) => s.tag === '平均奖励')!.points[0]).toEqual({ + step: 2, + value: 4, + wallTime: 0, + }); +}); +it('滚动截断按重叠上下文补全,无header/overlap不猜迭代,忽略非有限/缺失/其它数字', () => { + const h = new TrainingMetricHistory(); + h.update('a', [header(3), 'Mean value loss: 3']); + h.update('a', ['Mean value loss: 3', 'Mean reward: 7']); + expect(h.series()[1].points[0].step).toBe(3); + h.update('a', ['Mean entropy loss: 5']); + expect(h.series()).toHaveLength(2); + h.update('a', [ + header(4), + 'Mean value loss: NaN', + 'Mean reward: Infinity', + 'Mean reward: 1e999', + 'Mean reward: 3 ms', + 'Total timesteps: 40', + 'Iteration time: 2', + ]); + expect(h.series()[0].points).toHaveLength(1); + h.update('b', ['Mean reward: 1']); + expect(h.series()).toEqual([]); + h.update('b', [header(0), 'Mean reward: .5']); + expect(h.series()[0].points[0].step).toBe(0); +}); +it('容量和去重索引有界,旧重发不挤掉最新迭代,job切换清空', () => { + const h = new TrainingMetricHistory(); + for (let i = 0; i < 650; i++) h.update('a', [header(i), `Mean reward: ${i}`]); + const points = h.series()[0].points; + expect(points).toHaveLength(500); + expect(points[0].step).toBe(150); + expect(points.at(-1)!.step).toBe(649); + h.update('a', [header(1), 'Mean reward: 999']); + expect(h.series()[0].points).toEqual(points); + h.update('new-job', [header(0), 'Mean value loss: 1']); + expect(h.series()[0].points).toHaveLength(1); + expect(() => new TrainingMetricHistory(1)).toThrow(); +}); diff --git a/web_platform/src/training/TrainingMetricHistory.ts b/web_platform/src/training/TrainingMetricHistory.ts new file mode 100644 index 00000000..fa7fb36f --- /dev/null +++ b/web_platform/src/training/TrainingMetricHistory.ts @@ -0,0 +1,93 @@ +import type { ScalarSeries } from './types'; + +export const TRAINING_METRICS = { + value: '价值损失', + surrogate: '策略损失', + entropy: '熵损失', + reward: '平均奖励', + episodeLength: '平均回合长度', +} as const; +type Metric = keyof typeof TRAINING_METRICS; +const METRIC_PATTERN = + /^\s*Mean (value loss|surrogate loss|entropy loss|reward|episode length):\s*([-+]?(?:\d+(?:\.\d*)?|\.\d+)(?:e[-+]?\d+)?)\s*$/i; +const KEYS: Record = { + 'value loss': 'value', + 'surrogate loss': 'surrogate', + 'entropy loss': 'entropy', + reward: 'reward', + 'episode length': 'episodeLength', +}; + +/** Bounded iteration merge. Headerless lines are accepted only with proven suffix/prefix overlap. + * On reconnect without overlap, orphan scalars are skipped rather than guessed onto a new step. */ +export class TrainingMetricHistory { + private readonly rows = new Map>>(); + private previous: { line: string; iteration?: number }[] = []; + private jobId?: string; + constructor(private readonly capacity = 500) { + if (!Number.isInteger(capacity) || capacity < 2) throw new Error('指标容量必须至少为2'); + } + update(jobId: string, logs: readonly string[]): boolean { + let changed = false; + if (this.jobId !== jobId) { + this.rows.clear(); + this.previous = []; + this.jobId = jobId; + changed = true; + } + const lines = logs + .flatMap((line) => line.split('\n')) + .slice(-1000) + // rsl_rl uses ANSI SGR colors around iteration headers. + // eslint-disable-next-line no-control-regex + .map((line) => line.replace(/\u001b\[[0-9;]*m/g, '').trim()); + let iteration: number | undefined; + for (let overlap = Math.min(lines.length, this.previous.length); overlap > 0; overlap--) { + const start = this.previous.length - overlap; + if (lines.slice(0, overlap).every((line, i) => this.previous[start + i].line === line)) { + iteration = this.previous[start].iteration; + break; + } + } + const contexts: typeof this.previous = []; + for (const line of lines) { + const header = /\bLearning iteration\s+(\d+)\s*\/\s*\d+/i.exec(line); + if (header) { + const step = Number(header[1]); + iteration = Number.isSafeInteger(step) ? step : undefined; + } + contexts.push({ line, iteration }); + const scalar = METRIC_PATTERN.exec(line); + if (iteration === undefined || !scalar) continue; + const value = Number(scalar[2]); + if (!Number.isFinite(value)) continue; + const key = KEYS[scalar[1].toLowerCase()]; + if (!this.rows.has(iteration) && this.rows.size >= this.capacity) { + const oldest = Math.min(...this.rows.keys()); + if (iteration < oldest) continue; + this.rows.delete(oldest); + } + const row = this.rows.get(iteration) ?? {}; + if (row[key] !== value) { + row[key] = value; + changed = true; + } + this.rows.set(iteration, row); + } + this.previous = contexts; + return changed; + } + series(): ScalarSeries[] { + const rows = [...this.rows].sort(([a], [b]) => a - b); + return Object.entries(TRAINING_METRICS) + .map(([key, tag]) => ({ + tag, + points: rows.flatMap(([step, row]) => + row[key as Metric] === undefined + ? [] + : [{ step, wallTime: 0, value: row[key as Metric]! }], + ), + })) + .filter((series) => series.points.length > 0); + } +} diff --git a/web_platform/src/training/TrainingMetricsPanel.test.tsx b/web_platform/src/training/TrainingMetricsPanel.test.tsx new file mode 100644 index 00000000..e6a1001c --- /dev/null +++ b/web_platform/src/training/TrainingMetricsPanel.test.tsx @@ -0,0 +1,36 @@ +import { fireEvent, render, screen } from '@testing-library/react'; +import { TrainingMetricsPanel } from './TrainingMetricsPanel'; +import { ScalarChart } from '../components/charts/ScalarChart'; +vi.mock('../components/charts/ScalarChart', () => ({ + ScalarChart: vi.fn(({ title }: { title: string }) => ( +
{title}曲线
+ )), +})); + +it('折叠不挂载图表,价值/策略/综合独立缩放,job切换清空历史', () => { + const logs = [ + 'Learning iteration 1 / 10', + 'Mean value loss: 1', + 'Mean surrogate loss: -2', + 'Mean entropy loss: -3', + 'Mean reward: 4', + 'Mean episode length: 50', + ]; + const view = render(); + expect(screen.queryByTestId('metric-chart')).not.toBeInTheDocument(); + fireEvent.click(screen.getByRole('button', { name: /训练指标趋势/ })); + expect(screen.getByText('价值损失曲线')).toBeInTheDocument(); + vi.mocked(ScalarChart).mockClear(); + view.rerender(); + expect(ScalarChart).not.toHaveBeenCalled(); + fireEvent.click(screen.getByRole('tab', { name: '策略损失' })); + expect(screen.getByText('策略损失曲线')).toBeInTheDocument(); + fireEvent.click(screen.getByRole('tab', { name: '综合' })); + expect(screen.getAllByTestId('metric-chart')).toHaveLength(5); + expect(screen.getByText('平均回合长度曲线')).toBeInTheDocument(); + view.rerender(); + expect(screen.queryByTestId('metric-chart')).not.toBeInTheDocument(); + expect(screen.getByText(/尚无带迭代编号/)).toBeInTheDocument(); + fireEvent.click(screen.getByRole('button', { name: /训练指标趋势/ })); + expect(screen.queryByRole('tabpanel')).not.toBeInTheDocument(); +}); diff --git a/web_platform/src/training/TrainingMetricsPanel.tsx b/web_platform/src/training/TrainingMetricsPanel.tsx new file mode 100644 index 00000000..4c3e7708 --- /dev/null +++ b/web_platform/src/training/TrainingMetricsPanel.tsx @@ -0,0 +1,81 @@ +import { memo, useEffect, useState } from 'react'; +import { ScalarChart } from '../components/charts/ScalarChart'; +import { TrainingMetricHistory, TRAINING_METRICS } from './TrainingMetricHistory'; +import type { ScalarSeries } from './types'; + +const MetricChart = memo(function MetricChart({ series }: { series: ScalarSeries }) { + return ( +
+

+ {series.tag} · 最新原值 {series.points.at(-1)!.value.toPrecision(4)} +

+ +
+ ); +}); + +export const TrainingMetricsPanel = memo(function TrainingMetricsPanel({ + jobId, + logs, +}: { + jobId: string; + logs: readonly string[]; +}) { + const [history] = useState(() => new TrainingMetricHistory()); + const [series, setSeries] = useState([]); + const [open, setOpen] = useState(false); + const [tab, setTab] = useState('value'); + useEffect(() => { + if (history.update(jobId, logs)) setSeries(history.series()); + }, [history, jobId, logs]); + const selected = series.filter( + (item) => + tab === 'all' || + item.tag === (tab === 'value' ? TRAINING_METRICS.value : TRAINING_METRICS.surrogate), + ); + return ( +
+ + {open && ( +
+
+ {[ + ['value', '价值损失'], + ['surrogate', '策略损失'], + ['all', '综合'], + ].map(([id, label]) => ( + + ))} +
+
+ {selected.length ? ( + selected.map((item) => ) + ) : ( +

尚无带迭代编号的指标日志

+ )} +
+

+ 各指标独立纵轴;EMA 0.4 + 仅用于曲线,悬停显示原值。最多保留最近500个有指标的迭代,重连仅恢复服务端日志尾部。 +

+
+ )} +
+ ); +}); diff --git a/web_platform/src/training/pretrainedSelection.ts b/web_platform/src/training/pretrainedSelection.ts new file mode 100644 index 00000000..9209babe --- /dev/null +++ b/web_platform/src/training/pretrainedSelection.ts @@ -0,0 +1,21 @@ +import type { PretrainedSource } from './types'; + +/** A nonempty selection is an initialization intent, never a random-training fallback. */ +export function pretrainedSelectionError( + sources: readonly PretrainedSource[] | undefined, + taskId: string, + selectedId: string, +): string | undefined { + if (!selectedId) return undefined; + const source = sources?.find((item) => item.id === selectedId); + const reason = !source + ? '目录缺少该内容ID' + : !source.ready + ? '来源未通过验证' + : !source.compatibleTasks.includes(taskId) + ? '来源与当前任务不兼容' + : undefined; + return reason + ? `所选基础策略已失效:${reason}。请明确选择其他有效基础策略,或选择“不选择(随机初始化)”;不会自动退回随机初始化。` + : undefined; +} diff --git a/web_platform/src/training/trainingLosses.test.ts b/web_platform/src/training/trainingLosses.test.ts new file mode 100644 index 00000000..1bbac25a --- /dev/null +++ b/web_platform/src/training/trainingLosses.test.ts @@ -0,0 +1,18 @@ +import { expect, it } from 'vitest'; +import { trainingLosses } from './trainingLosses'; +it('提取最新有限PPO损失,兼容无指标旧日志', () => { + expect( + trainingLosses([ + 'Mean value loss: 1.2', + 'Mean surrogate loss: -0.03', + 'Mean entropy loss: 1e-3', + 'Mean value loss: 0.1', + 'Mean value loss: Infinity', + ]), + ).toEqual([ + { label: '价值损失', value: 0.1 }, + { label: '策略损失', value: -0.03 }, + { label: '熵损失', value: 0.001 }, + ]); + expect(trainingLosses(['Starting training'])).toEqual([]); +}); diff --git a/web_platform/src/training/trainingLosses.ts b/web_platform/src/training/trainingLosses.ts new file mode 100644 index 00000000..b560f2d3 --- /dev/null +++ b/web_platform/src/training/trainingLosses.ts @@ -0,0 +1,14 @@ +/** rsl_rl console scalars, newest finite value per loss; old servers need no new endpoint. */ +export function trainingLosses(logs: readonly string[]): { label: string; value: number }[] { + const values = new Map(); + for (const line of logs) { + const match = /Mean (value|surrogate|entropy) loss:\s*([-+\d.eE]+)/.exec(line); + if (match && Number.isFinite(Number(match[2]))) values.set(match[1], Number(match[2])); + } + const labels: Record = { + value: '价值损失', + surrogate: '策略损失', + entropy: '熵损失', + }; + return Array.from(values, ([key, value]) => ({ label: labels[key], value })); +} diff --git a/web_platform/src/training/types.ts b/web_platform/src/training/types.ts index 15247b61..b8a862ea 100644 --- a/web_platform/src/training/types.ts +++ b/web_platform/src/training/types.ts @@ -1,13 +1,75 @@ +import type { PolicyDeployment, TrainingTerrain } from '../rl/deployment'; + +export interface TrainingParameter { + min: number; + max: number; + default: number; + integer?: boolean; +} +export interface TrainingTaskMetadata { + id: string; + name: string; + browserCompatible: boolean; + terrainPresets: string[]; + terrainParameters: Record; + sensorTypes: string[]; + sensorModes?: string[]; + sensorParameters: Record; + mapSyncScope: string; +} export type TrainingJobState = 'queued' | 'running' | 'succeeded' | 'failed' | 'cancelled'; export type TrainingDevice = 'cpu' | 'gpu'; export type WandbMode = 'offline' | 'online' | 'disabled'; +export interface PretrainedInitialization { + sourceId: string; + registeredId: string; + label: string; + manifest: { + source_iteration: number | null; + sourceFormat?: 'pt' | 'onnx'; + contract?: string; + derived_fields?: { + normalizer_count?: { policy: string; value: number }; + exploration_std?: { policy: string; value: number }; + }; + source_actor_dim: number; + normalization: string; + artifacts: { checkpoint: PretrainedArtifact } & Partial< + Record<'upload' | 'onnx' | 'env' | 'agent', PretrainedArtifact> + >; + }; +} +export interface PretrainedArtifact { + name: string; + sha256: string; + bytes: number; +} +export interface PretrainedUploadCapability { + enabled: boolean; + templateId: 'go2-legacy47-v1'; + formats: { pt: number; onnx: number }; + endpoint: string; +} +export interface PretrainedSource { + id: string; + label: string; + ready: boolean; + compatibleTasks: string[]; + observationSizes?: number[]; + initialization?: PretrainedInitialization; + error?: string; +} + export interface TrainingServerInfo { + pretrainedSources?: PretrainedSource[]; + pretrainedUpload?: PretrainedUploadCapability; version: string; ready: boolean; trainerRoot: string; python: string; tasks: string[]; + taskMetadata?: TrainingTaskMetadata[]; activeJobId?: string; resourceOwner?: string; tuning?: TuningCapability; @@ -15,6 +77,11 @@ export interface TrainingServerInfo { } export interface TrainingRequest { + customTerrainBoxes?: TrainingTerrain; + terrainPreset?: string; + terrainParams?: Record; + sensorType?: 'raycast'; + sensorCfg?: Partial; taskId: string; numEnvs: number; maxIterations: number; @@ -24,9 +91,12 @@ export interface TrainingRequest { gpuIds: number[]; wandbMode: WandbMode; rewardPresetId?: string; + pretrainedSourceId?: string; } export interface TrainingJob { + pretrained?: PretrainedInitialization; + deployment?: PolicyDeployment; id: string; state: TrainingJobState; taskId: string; @@ -60,16 +130,11 @@ export interface RewardConfiguration { params: Record; } -export interface ObjectiveWeights { - velocity_tracking: number; - action_smoothness: number; - posture_stability: number; - fall_avoidance: number; - foot_slip: number; - energy: number; -} +export type ObjectiveWeights = Record; export interface TuningCapability { + pretrainedSources?: PretrainedSource[]; + pretrainedUpload?: PretrainedUploadCapability; ready: boolean; configured: boolean; apiKeyConfigured: boolean; @@ -79,7 +144,12 @@ export interface TuningCapability { } export interface TuningCreateRequest { - taskId: 'Unitree-Go2-Flat'; + pretrainedSourceId?: string; + taskId: 'Unitree-Go2-Flat' | 'Unitree-Go2-ObstacleAvoidance'; + taskConfig?: Pick< + TrainingRequest, + 'terrainPreset' | 'terrainParams' | 'sensorType' | 'sensorCfg' | 'customTerrainBoxes' + > & { seed?: number }; mode: TuningMode; runName: string; numEnvs: number; @@ -160,7 +230,11 @@ export interface TuningSession { mode: TuningMode; createdAt: string; updatedAt: string; - config: TuningCreateRequest & { rungs: number[]; promote: number[] }; + config: TuningCreateRequest & { + rungs: number[]; + promote: number[]; + pretrained?: PretrainedInitialization; + }; objectiveWeights: ObjectiveWeights; message: string; currentTrialId?: string; @@ -194,6 +268,7 @@ export interface TuningMetricsResponse { export interface RewardPreset { id: string; + taskId: string; name: string; sessionId: string; trialId: string; diff --git a/web_platform/src/tuning/AgentDecisionTimeline.tsx b/web_platform/src/tuning/AgentDecisionTimeline.tsx index be7e56d3..46d0e8d5 100644 --- a/web_platform/src/tuning/AgentDecisionTimeline.tsx +++ b/web_platform/src/tuning/AgentDecisionTimeline.tsx @@ -25,6 +25,7 @@ import { formatMetric, mergeRewardPatch, OBJECTIVE_META, + OBSTACLE_OBJECTIVE_META, PARAMETER_BY_PATH, rewardConfigurationDiff, } from './domain'; @@ -223,7 +224,9 @@ function ExpectedImpact({ proposal }: { proposal: TuningProposal }) { return (
{values.map(([key, value]) => { - const label = OBJECTIVE_META.find((item) => item.key === key)?.label ?? key; + const label = + [...OBJECTIVE_META, ...OBSTACLE_OBJECTIVE_META].find((item) => item.key === key)?.label ?? + key; return (

{label}

diff --git a/web_platform/src/tuning/MetricsComparisonBoard.tsx b/web_platform/src/tuning/MetricsComparisonBoard.tsx index fa2f086d..80e41484 100644 --- a/web_platform/src/tuning/MetricsComparisonBoard.tsx +++ b/web_platform/src/tuning/MetricsComparisonBoard.tsx @@ -4,7 +4,7 @@ import uPlot from 'uplot'; import { useShallow } from 'zustand/react/shallow'; import { Badge, Button, Select } from '../components/ui'; import type { ScalarPoint, TuningTrial } from '../training/types'; -import { formatMetric, OBJECTIVE_META } from './domain'; +import { formatMetric, objectiveMeta } from './domain'; import { useTuningStore } from './tuningStore'; const LINE_COLORS = ['#60a5fa', '#a78bfa', '#f472b6', '#2dd4bf', '#fb7185', '#94a3b8']; @@ -275,6 +275,7 @@ function MultiTrialPlot({ } function ScoreBreakdown({ current, best }: { current?: TuningTrial; best?: TuningTrial }) { + const taskId = useTuningStore((state) => state.sessionConfig?.taskId); const currentComponents = current?.evaluation?.score?.components ?? {}; const bestComponents = best?.evaluation?.score?.components ?? {}; return ( @@ -282,7 +283,11 @@ function ScoreBreakdown({ current, best }: { current?: TuningTrial; best?: Tunin

Score Breakdown

-

相对基线改善 · 中线为 0

+

+ {taskId === 'Unitree-Go2-ObstacleAvoidance' + ? '固定客观归一化 · 0–1(越高越好)' + : '相对基线改善 · 中线为 0'} +

{best && ( @@ -291,7 +296,7 @@ function ScoreBreakdown({ current, best }: { current?: TuningTrial; best?: Tunin )}
- {OBJECTIVE_META.map(({ key, label }) => { + {objectiveMeta(taskId).map(({ key, label }) => { const currentValue = currentComponents[key] ?? 0; const bestValue = bestComponents[key] ?? 0; const currentWidth = Math.min(50, Math.abs(currentValue) * 50); @@ -332,15 +337,20 @@ function ScoreBreakdown({ current, best }: { current?: TuningTrial; best?: Tunin } function MetricKpis({ current, best }: { current?: TuningTrial; best?: TuningTrial }) { + const taskId = useTuningStore((state) => state.sessionConfig?.taskId); const currentMetrics = current?.evaluation?.metrics ?? {}; const bestMetrics = best?.evaluation?.metrics ?? {}; return (
- {OBJECTIVE_META.map(({ key, shortLabel, metric }) => { + {objectiveMeta(taskId).map(({ key, shortLabel, metric }) => { const currentValue = currentMetrics[metric]; const bestValue = bestMetrics[metric]; const worse = - currentValue !== undefined && bestValue !== undefined && currentValue > bestValue; + currentValue !== undefined && + bestValue !== undefined && + (taskId === 'Unitree-Go2-ObstacleAvoidance' + ? currentValue < bestValue + : currentValue > bestValue); return (

diff --git a/web_platform/src/tuning/ScalarChart.tsx b/web_platform/src/tuning/ScalarChart.tsx index c32996ae..07ad6c8a 100644 --- a/web_platform/src/tuning/ScalarChart.tsx +++ b/web_platform/src/tuning/ScalarChart.tsx @@ -1,205 +1 @@ -import { useEffect, useMemo, useRef } from 'react'; -import { RotateCcw, ZoomIn, ZoomOut } from 'lucide-react'; -import uPlot from 'uplot'; -import type { ScalarSeries } from '../training/types'; - -const COLORS = ['#38d39f', '#60a5fa', '#f59e0b', '#f472b6', '#a78bfa', '#fb7185']; - -function smooth(values: (number | null)[], factor: number): (number | null)[] { - if (factor <= 0) return values; - let previous: number | undefined; - return values.map((value) => { - if (value === null) return null; - previous = previous === undefined ? value : factor * previous + (1 - factor) * value; - return previous; - }); -} - -export function ScalarChart({ - series, - smoothing, - title = '训练与评估 Scalars', -}: { - series: ScalarSeries[]; - smoothing: number; - title?: string; -}) { - const host = useRef(null); - const chartRef = useRef(null); - const trackZoom = useRef(false); - const zoomRanges = useRef>>({}); - const prepared = useMemo(() => { - const steps = Array.from( - new Set(series.flatMap((item) => item.points.map((point) => point.step))), - ).sort((a, b) => a - b); - const columns: uPlot.AlignedData = [steps]; - for (const item of series) { - const byStep = new Map(item.points.map((point) => [point.step, point.value])); - columns.push( - smooth( - steps.map((step) => byStep.get(step) ?? null), - smoothing, - ), - ); - } - return columns; - }, [series, smoothing]); - - useEffect(() => { - if (!host.current || series.length === 0 || prepared[0].length === 0) return; - const element = host.current; - const width = Math.max(280, Math.floor(element.getBoundingClientRect().width)); - trackZoom.current = false; - const chart = new uPlot( - { - width, - height: 280, - cursor: { drag: { x: true, y: true, setScale: true } }, - scales: { x: { time: false } }, - hooks: { - setScale: [ - (instance, key) => { - if (!trackZoom.current || (key !== 'x' && key !== 'y')) return; - const scale = instance.scales[key]; - if (typeof scale.min === 'number' && typeof scale.max === 'number') - zoomRanges.current[key] = { min: scale.min, max: scale.max }; - }, - ], - }, - axes: [ - { stroke: '#8fa0b5', grid: { stroke: '#213044' } }, - { stroke: '#8fa0b5', grid: { stroke: '#213044' } }, - ], - series: [ - { label: 'Step' }, - ...series.map((item, index) => ({ - label: item.tag, - stroke: COLORS[index % COLORS.length], - width: 2, - spanGaps: true, - })), - ], - }, - prepared, - element, - ); - chartRef.current = chart; - trackZoom.current = true; - for (const key of ['x', 'y'] as const) { - const range = zoomRanges.current[key]; - if (range) chart.setScale(key, range); - } - const wheelZoom = (event: WheelEvent) => { - if (!event.deltaY) return; - event.preventDefault(); - event.stopPropagation(); - const bounds = chart.over.getBoundingClientRect(); - if (!bounds.width || !bounds.height) return; - const xRatio = Math.min(1, Math.max(0, (event.clientX - bounds.left) / bounds.width)); - const yRatio = Math.min(1, Math.max(0, (event.clientY - bounds.top) / bounds.height)); - const unit = event.deltaMode === 1 ? 16 : event.deltaMode === 2 ? window.innerHeight : 1; - const factor = Math.min(2, Math.max(0.5, Math.exp(event.deltaY * unit * 0.002))); - for (const key of ['x', 'y'] as const) { - const scale = chart.scales[key]; - if (typeof scale.min !== 'number' || typeof scale.max !== 'number') continue; - const ratio = key === 'x' ? xRatio : 1 - yRatio; - const anchor = scale.min + (scale.max - scale.min) * ratio; - chart.setScale(key, { - min: anchor - (anchor - scale.min) * factor, - max: anchor + (scale.max - anchor) * factor, - }); - } - }; - chart.over.addEventListener('wheel', wheelZoom, { passive: false }); - let frame = 0; - let lastWidth = width; - const observer = new ResizeObserver(() => { - window.cancelAnimationFrame(frame); - frame = window.requestAnimationFrame(() => { - const nextWidth = Math.max(280, Math.floor(element.getBoundingClientRect().width)); - if (nextWidth !== lastWidth) { - lastWidth = nextWidth; - chart.setSize({ width: nextWidth, height: 280 }); - } - }); - }); - observer.observe(element); - return () => { - observer.disconnect(); - chart.over.removeEventListener('wheel', wheelZoom); - window.cancelAnimationFrame(frame); - trackZoom.current = false; - chartRef.current = null; - chart.destroy(); - }; - }, [prepared, series]); - - const zoom = (factor: number) => { - const chart = chartRef.current; - if (!chart) return; - for (const key of ['x', 'y']) { - const scale = chart.scales[key]; - if (typeof scale.min !== 'number' || typeof scale.max !== 'number') continue; - const center = (scale.min + scale.max) / 2; - const radius = ((scale.max - scale.min) * factor) / 2 || 1; - chart.setScale(key, { min: center - radius, max: center + radius }); - } - }; - const resetZoom = () => { - const chart = chartRef.current; - if (!chart) return; - zoomRanges.current = {}; - trackZoom.current = false; - chart.setData(chart.data, true); - trackZoom.current = true; - }; - - if (series.length === 0) - return ( -

- 当前 trial 尚无 scalar 数据 -
- ); - return ( -
-
-

- {title} -

-
- - - -
-
-
-

- 图表区域滚轮以指针位置缩放;也可拖拽框选,或使用右上角按钮缩放和复位。 -

-
- ); -} +export { ScalarChart } from '../components/charts/ScalarChart'; diff --git a/web_platform/src/tuning/TuningApp.tsx b/web_platform/src/tuning/TuningApp.tsx index 7379f72f..841d97cc 100644 --- a/web_platform/src/tuning/TuningApp.tsx +++ b/web_platform/src/tuning/TuningApp.tsx @@ -3,9 +3,11 @@ import { Bot, FlaskConical, Play, RefreshCw, Wifi } from 'lucide-react'; import { useShallow } from 'zustand/react/shallow'; import { Badge, Button, Select } from '../components/ui'; import { LocalTrainingClient } from '../training/LocalTrainingClient'; +import { PretrainedSourceSelect } from '../training/PretrainedSourceSelect'; +import { pretrainedSelectionError } from '../training/pretrainedSelection'; import { rememberTrainingConnection } from '../training/storage'; import type { ObjectiveWeights, TuningCreateRequest, TuningMode } from '../training/types'; -import { OBJECTIVE_META, STATE_META } from './domain'; +import { objectiveMeta, OBSTACLE_OBJECTIVES, STATE_META } from './domain'; import { TuningConsole } from './TuningConsole'; import { useTuningStore } from './tuningStore'; @@ -32,13 +34,24 @@ export function TuningApp() { ] as const, ), ); + const [trainingContext, setTrainingContext] = + useState>(); const [testingAgent, setTestingAgent] = useState(false); useEffect(() => { const receive = (event: MessageEvent) => { if (event.origin !== location.origin || typeof event.data !== 'object') return; - const data = event.data as { type?: string; endpoint?: string; token?: string }; + const data = event.data as { + type?: string; + endpoint?: string; + token?: string; + trainingContext?: Pick< + TuningCreateRequest, + 'taskId' | 'taskConfig' | 'seed' | 'pretrainedSourceId' + >; + }; if (data.type === 'mujoco-tuning-credentials' && data.endpoint && data.token) { + setTrainingContext(data.trainingContext); useTuningStore.getState().setConnection(data.endpoint, data.token); rememberTrainingConnection(data.endpoint, data.token); } @@ -129,7 +142,11 @@ export function TuningApp() { )} - {sessionId ? : } + {sessionId ? ( + + ) : ( + + )} {error && (