Cen #8
@@ -1,3 +1,4 @@
|
||||
Please always speak chinese.
|
||||
python 虚拟环境路径在/home/cen/Embodied_Workspace/Mujoco_Projects/mujoco/.venv/bin/activate
|
||||
系统是Ubuntu 24.04 LTS
|
||||
系统是Ubuntu 24.04 LTS
|
||||
每次提交版本前,更新CHANGELOG.md文件
|
||||
+118
@@ -2,6 +2,104 @@
|
||||
|
||||
本项目的重要变更记录在此文件中,版本标签沿用仓库现有的 `V主版本.次版本[.修订版本]` 格式。
|
||||
|
||||
## [未发布]
|
||||
|
||||
## [0.9.4] - 2026-09-09
|
||||
|
||||
- 前端设计优化:升级蓝黑、电光青与少量紫色的 Cyber HUD 视觉,保留亮色主题;统一工作台、欢迎界面、共享控件、浮层和源码/调参窗口材质。
|
||||
- 恢复自调参平台左侧会话与 Trial、中间指标与排行、右侧决策与审批的三栏布局;窄屏纵向排列,保留审批与 Monaco Diff 操作。
|
||||
- 增加运行状态呼吸与交互微光,支持减少动态效果、强制色彩和无模糊能力回退;侧栏采用静态材质,避免大面积实时模糊影响仿真性能。
|
||||
- 优化欢迎页与源代码入口、执行器控件密度和重复参数提示,将训练地图操作文案统一为“同步场景地图”。
|
||||
- 补充双主题材质、状态动效、回退及三栏布局回归测试和设计交付说明;性能数据仅为无头 Chromium 短时特效开关对照,不代表真实 GPU 或重型模型基准。
|
||||
|
||||
## [0.9.3] - 2026-09-08
|
||||
|
||||
- 重构绿色主工作台:紧凑全局入口、可收起资源/上下文侧栏,统一状态、仿真/地图工具、草稿、摄像头与风险槽位;窄屏临时布局不覆盖桌面偏好。
|
||||
- 逐模块整理地图、Body/Joint/Actuator、Python/ONNX、训练与录制的信息层级,保留单位、主动作、兼容限制及危险后果;控制台首次懒挂载后保留输入与任务状态。
|
||||
- 独立调参改为指标/排行主区和会话/决策侧展,保留审批、Monaco 差异、参数护栏及停止入口;主工作台/调参共享深浅主题,换主题不丢图表缩放或编辑草稿。
|
||||
- 统一 Tooltip/Popover/对话框边界、全屏 Portal 与顶层 Escape;修复 Dialog 焦点恢复、Monaco 模型释放顺序、摄像头实际渲染槽位,以及浅色主按钮对比色被共享 CSS 覆盖的问题。
|
||||
- 补深浅主题五档布局、领域/调参/系统截图和交互回归,同步设计文档及旧 E2E 入口。训练/调参布局采用 mock;真实训练、Agent、外部真实策略/上传服务和长期性能不作为本次 UI 验收通过项。
|
||||
|
||||
## [0.9.2] - 2026-09-08
|
||||
|
||||
- 修复custom_boxes因起终点编辑使同步状态持续失效而无法训练:移除训练面板起终点坐标显示与编辑,训练/调参启动时自动重新编译已应用碰撞场景,并自动选择满足净空、最小距离和连通性约束的参考点。
|
||||
|
||||
## [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 +160,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、缺项及兼容性变更的真实组件刷新回归,验证失效时零启动请求及用户明确选择后的请求体。
|
||||
|
||||
Generated
+2
-2
@@ -1,12 +1,12 @@
|
||||
{
|
||||
"name": "mujoco-web-platform",
|
||||
"version": "0.8.3",
|
||||
"version": "0.9.4",
|
||||
"lockfileVersion": 3,
|
||||
"requires": true,
|
||||
"packages": {
|
||||
"": {
|
||||
"name": "mujoco-web-platform",
|
||||
"version": "0.8.3",
|
||||
"version": "0.9.4",
|
||||
"license": "Apache-2.0",
|
||||
"dependencies": {
|
||||
"@monaco-editor/react": "^4.7.0",
|
||||
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "mujoco-web-platform",
|
||||
"version": "0.8.3",
|
||||
"version": "0.9.4",
|
||||
"description": "基于 MuJoCo WebAssembly 的本地机器人仿真与控制平台",
|
||||
"private": true,
|
||||
"type": "module",
|
||||
|
||||
@@ -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 <server-owned-path>` 传入训练器,绝不使用 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必须标准。面板不显示或编辑出生/目标坐标;场景编译器从连通自由栅格自动选取障碍/边界净空>0.55m且间距>=2m的确定性参考点,出生z固定.32,四元数取成功加载时真实初始浮动基座姿态。启动训练和打开调参时均自动重新编译当前已应用碰撞场景。
|
||||
- 出生及目标中心必须距边界>=.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均保持原协议。
|
||||
@@ -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=<URL编码显示名>`
|
||||
|
||||
- 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。
|
||||
|
||||
验证后原始四件产物复制到`<tuning-data-root>/pretrained_sources/<source_id>/`内容寻址只读快照。浏览器提交的`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产物。
|
||||
@@ -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不匹配或快照失效均明确拒绝,绝不静默随机初始化。
|
||||
|
||||
@@ -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)
|
||||
@@ -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"],
|
||||
]
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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))
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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"])
|
||||
@@ -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))
|
||||
|
||||
+150
-12
@@ -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
|
||||
|
||||
@@ -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
|
||||
]
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 = (
|
||||
"<mujoco><worldbody>"
|
||||
+ "".join(
|
||||
'<geom type="box" pos="{}" size="{}"/>'.format(
|
||||
" ".join(map(str, b["pos"])), " ".join(map(str, b["size"]))
|
||||
)
|
||||
for b in layouts[layout]
|
||||
)
|
||||
+ "</worldbody></mujoco>"
|
||||
)
|
||||
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()
|
||||
@@ -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()
|
||||
@@ -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 = '<geom name="terrain_0" type="box" pos="-.5 0 -.1" size="6 6 .1"/>'
|
||||
model = mujoco.MjModel.from_xml_string(
|
||||
'<mujoco><worldbody><body pos=".5 0 0">' + floor + "</body></worldbody></mujoco>"
|
||||
)
|
||||
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(
|
||||
'<mujoco><worldbody><body pos=".5 0 0">' + geoms + "</body></worldbody></mujoco>"
|
||||
)
|
||||
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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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))
|
||||
@@ -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()
|
||||
@@ -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": "<include/>"}},
|
||||
{"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")
|
||||
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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}")
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
# 自定义训练与前视射线避障
|
||||
|
||||
1. 启动仓库本地训练服务,在「控制台 → 强化学习任务」输入令牌并连接。
|
||||
2. 选择「前视射线避障导航」,设置地形、种子、障碍物参数、FOV、探测/安全距离及避障奖励权重,发起训练。界面显示进度、最近价值/策略/熵损失和原始日志。
|
||||
3. 可选的「检查当前场景地图」从全部**已应用**实例的已编译MuJoCo静态碰撞几何导出权威 `custom_boxes/boxes-v1`,支持多实例、平移/旋转、工程地图路径及嵌套body。不重新运行随机预设,不同步未应用草稿。面板不再显示或编辑出生/目标坐标:系统从连通栅格自动选择满足障碍/边界净空0.55m且相距至少2m的参考点;训练每个episode仍独立随机采样。点击训练或打开调参时会从当前已应用碰撞场景自动重新编译并校验,不依赖旧同步快照;绝不删除或移动障碍来腾出安全区。
|
||||
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不生效,重新按下的正常点击仍可设定。
|
||||
@@ -0,0 +1,48 @@
|
||||
# Cyber HUD 界面样式
|
||||
|
||||
本次直接升级现有 React 界面的 JSX/CSS,不增加静态演示入口、外部字体或依赖,保留亮色主题和所有仿真业务行为。
|
||||
|
||||
## 源码
|
||||
|
||||
- `src/styles.css`:蓝黑/电光青语义色、紫色装饰、玻璃材质、交互微光、状态呼吸和可访问性降级。
|
||||
- `src/components/ui/{Button,IconButton,Tabs,CollapsibleSection,Dialog,Popover}.tsx`:共享视觉挂点,保留原有键盘交互与 Portal。
|
||||
- `src/app/App.tsx` 及 `src/app/components/{WorkbenchHeader,SidebarPanel,ViewportHUD,WorkspaceOverlays,SourceEditorDialog}.tsx`:工作台外壳、欢迎卡片、HUD 和编辑窗口。
|
||||
- `src/tuning/{TuningApp,TuningConsole,TuningSessionRail,MetricsComparisonBoard}.tsx`:调参页、导航、指标卡片;Monaco 原生亮/暗代码主题不变。
|
||||
- `e2e/ui-cyber-hud.spec.ts`:新增 5 个材质、主题持久化、动效及回退测试。
|
||||
|
||||
## v0.9.4 调参布局
|
||||
|
||||
自调参会话页恢复三栏:左侧会话与 Trial、中间指标与排行、右侧决策与审批。宽度达到 1280px 时三栏同时显示并分别滚动,更窄窗口按上述顺序纵向排列。原有主题、审批与 Diff 编辑行为保持不变。
|
||||
|
||||
## 材质与性能约束
|
||||
|
||||
- 暗色基底 `#060914` / `#0B1224`,强调 `#36E4F2`,装饰 `#A78BFA`;亮色强调 `#087887`。
|
||||
- 毛玻璃仅用于顶栏、欢迎卡片、小型 HUD 和浮层;侧栏使用静态渐变质感,避免长列表滚动时大面积背景模糊。
|
||||
- 3.6 秒呼吸动效仅绑定真实仿真运行状态;装饰层不拦截指针。
|
||||
- 减少动态效果偏好禁用动画;无背景模糊支持时使用实色表面;强制色彩模式保留清晰边框。
|
||||
|
||||
## 验证记录
|
||||
|
||||
- `npm run typecheck`、`npm run lint`、`npm run build` 通过。
|
||||
- `npm run test`:97 个测试文件、408 项测试通过。
|
||||
- 工作台布局、共享浮层、领域面板、调参系统及新增 HUD 测试:59 项端到端测试通过。
|
||||
- 双主题覆盖 1920×1080、1440×900、1366×768、1024×768、768×800;检查了源码/Diff、弹层、主题同步和视口避让。
|
||||
- 变更源码的 Prettier 检查与 `git diff --check` 通过。
|
||||
- 构建仍提示部分包超过 500kB;本次未改变加载架构。
|
||||
|
||||
性能取证:Headless Chromium、1440×900、24 刚体场景下按“关闭特效→开启→开启→关闭”采样,每段 2.5 秒并滚动侧栏。初版全侧栏模糊约 54 FPS,取消侧栏实时模糊后开/关特效均约 60 FPS,P95 帧间隔约 16.7ms。
|
||||
|
||||
**限制**:这是短时特效开关对照,不是完整重构前后的严格基准,也不代表真实桌面 GPU 或重型模型表现;未执行 Firefox/Safari 实机验证。无模糊回退测试通过修改 CSS 能力分支来模拟,不等同于旧浏览器实测。
|
||||
|
||||
## 本地证据
|
||||
|
||||
- `../test-results/cyber-hud-delivery/workbench-dark.png`
|
||||
- `../test-results/cyber-hud-delivery/workbench-light.png`
|
||||
- `../test-results/cyber-hud-delivery/tuning-dark.png`
|
||||
- `../test-results/cyber-hud-delivery/tuning-light.png`
|
||||
- `../test-results/cyber-hud-delivery/before-dark.png`
|
||||
- `../test-results/cyber-hud-delivery/performance-final.json`
|
||||
|
||||
截图与性能证据为本地生成产物,不纳入源码提交;再次运行 Playwright 可能清理 `test-results/`。详细命令日志保存在 `/tmp/cyber-hud-evidence/`。
|
||||
|
||||
预览:`npm run dev`,主入口 `/`,独立调参入口 `/tuning.html`。
|
||||
+167
-89
@@ -1,4 +1,4 @@
|
||||
import { expect, test } from '@playwright/test';
|
||||
import { expect, test, type Page } from '@playwright/test';
|
||||
import { readFileSync } from 'node:fs';
|
||||
import { fileURLToPath } from 'node:url';
|
||||
import { zipSync } from 'fflate';
|
||||
@@ -6,6 +6,33 @@ import { zipSync } from 'fflate';
|
||||
const fixture = (relative: string) =>
|
||||
fileURLToPath(new URL(`../fixtures/${relative}`, import.meta.url));
|
||||
|
||||
async function openProjectPanel(page: Page) {
|
||||
const show = page.getByRole('button', { name: '显示工程面板', exact: true });
|
||||
if (await show.isVisible()) await show.click();
|
||||
}
|
||||
|
||||
async function expectImportedFiles(page: Page, names: string[]) {
|
||||
await openProjectPanel(page);
|
||||
await page.getByRole('tab', { name: '工程文件', exact: true }).click();
|
||||
const tree = page.getByLabel('工程文件树', { exact: true });
|
||||
for (const name of names) await expect(tree.getByText(name, { exact: true })).toBeVisible();
|
||||
await page.screenshot({ path: test.info().outputPath('imported-project-files.png') });
|
||||
}
|
||||
|
||||
async function expectNotificationDetails(
|
||||
page: Page,
|
||||
messages: RegExp[],
|
||||
absentMessages: RegExp[] = [],
|
||||
) {
|
||||
await page.getByRole('button', { name: '通知中心' }).click();
|
||||
const notifications = page.getByRole('dialog', { name: '通知中心' });
|
||||
for (const summary of await notifications.getByText('事件详情', { exact: true }).all())
|
||||
await summary.click();
|
||||
for (const message of messages) await expect(notifications).toContainText(message);
|
||||
for (const message of absentMessages) await expect(notifications).not.toContainText(message);
|
||||
await page.keyboard.press('Escape');
|
||||
}
|
||||
|
||||
const SIMPLE_MODEL = `
|
||||
<mujoco model="e2e">
|
||||
<compiler angle="radian"/>
|
||||
@@ -79,7 +106,7 @@ const LARGE_MODEL = `
|
||||
|
||||
test('独立自调参工作台不需要加载 MuJoCo 主应用即可打开', async ({ page }) => {
|
||||
await page.goto('/tuning.html');
|
||||
await expect(page.getByRole('heading', { name: 'Go2 奖励函数自调参 Agent' })).toBeVisible();
|
||||
await expect(page.getByRole('heading', { name: 'Go2 自调参' })).toBeVisible();
|
||||
await expect(page.getByText('新建 Unitree-Go2-Flat 调参 Session')).toBeVisible();
|
||||
await expect(page.getByLabel('访问令牌(仅当前标签页)')).toBeVisible();
|
||||
await expect(page.getByRole('button', { name: '启动自调参' })).toBeVisible();
|
||||
@@ -102,20 +129,23 @@ test('显示中文平台骨架并加载单文件模型', async ({ page }) => {
|
||||
const resetCameraBox = await page.getByRole('button', { name: '相机复位' }).boundingBox(),
|
||||
playBox = await page.getByRole('button', { name: '▶ 播放' }).boundingBox();
|
||||
expect(
|
||||
resetCameraBox && playBox && resetCameraBox.x + resetCameraBox.width <= playBox.x,
|
||||
resetCameraBox && playBox && resetCameraBox.y + resetCameraBox.height < playBox.y,
|
||||
).toBeTruthy();
|
||||
await page.getByRole('button', { name: '更多工作台操作' }).click();
|
||||
await expect(page.getByRole('menuitem', { name: '工作台设置' })).toBeVisible();
|
||||
await page.keyboard.press('Escape');
|
||||
await page.setViewportSize({ width: 1440, height: 900 });
|
||||
await expect(page.getByRole('heading', { name: 'MuJoCo Web 仿真平台' })).toBeVisible();
|
||||
await expect(page.getByRole('heading', { name: 'MuJoCo' })).toBeVisible();
|
||||
await expect(page.getByRole('main').getByText('拖放模型工程到此处')).toBeVisible();
|
||||
await expect(page.getByRole('img', { name: 'XYZ 方向指示器' })).toBeVisible();
|
||||
await expect(page.getByRole('button', { name: '切换到白天主题' })).toBeVisible();
|
||||
await page.getByRole('button', { name: '更多工作台操作' }).click();
|
||||
await expect(page.getByRole('menuitem', { name: '切换主题' })).toBeVisible();
|
||||
await page.keyboard.press('Escape');
|
||||
await page.getByRole('button', { name: '布局设置' }).click();
|
||||
await expect(page.getByRole('dialog', { name: '布局设置' })).toBeVisible();
|
||||
await page.keyboard.press('Escape');
|
||||
await page.getByRole('button', { name: '工作台设置' }).click();
|
||||
await page.getByRole('button', { name: '更多工作台操作' }).click();
|
||||
await page.getByRole('menuitem', { name: '工作台设置' }).click();
|
||||
await expect(page.getByRole('dialog', { name: '工作台设置' })).toBeVisible();
|
||||
await page.keyboard.press('Escape');
|
||||
await page.keyboard.press('Control+k');
|
||||
@@ -123,16 +153,20 @@ test('显示中文平台骨架并加载单文件模型', async ({ page }) => {
|
||||
await page.getByLabel('搜索命令').fill('复位相机');
|
||||
await expect(page.getByRole('option', { name: /复位相机/ })).toBeVisible();
|
||||
await page.keyboard.press('Escape');
|
||||
await page.getByRole('button', { name: '进入全屏' }).click();
|
||||
await expect(page.getByRole('button', { name: '退出全屏' })).toBeVisible();
|
||||
await page.getByRole('button', { name: '更多工作台操作' }).click();
|
||||
await page.getByRole('menuitem', { name: '进入全屏' }).click();
|
||||
await expect.poll(() => page.evaluate(() => Boolean(document.fullscreenElement))).toBe(true);
|
||||
await page.keyboard.press('Control+k');
|
||||
await expect(page.getByRole('dialog', { name: '命令面板' })).toBeVisible();
|
||||
await page.keyboard.press('Escape');
|
||||
await page.getByRole('button', { name: '退出全屏' }).click();
|
||||
await page.getByRole('button', { name: '切换到白天主题' }).click();
|
||||
await page.getByRole('button', { name: '更多工作台操作' }).click();
|
||||
await page.getByRole('menuitem', { name: '退出全屏' }).click();
|
||||
await page.getByRole('button', { name: '更多工作台操作' }).click();
|
||||
await page.getByRole('menuitem', { name: '切换主题' }).click();
|
||||
await expect(page.locator('#root > div')).toHaveClass(/theme-light/);
|
||||
await expect(page.getByRole('button', { name: '切换到黑夜主题' })).toBeVisible();
|
||||
await page.getByRole('button', { name: '切换到黑夜主题' }).click();
|
||||
await expect(page.locator('#root > div')).toHaveClass(/theme-light/);
|
||||
await page.getByRole('button', { name: '更多工作台操作' }).click();
|
||||
await page.getByRole('menuitem', { name: '切换主题' }).click();
|
||||
await expect(page.locator('#root > div')).toHaveClass(/theme-dark/);
|
||||
|
||||
await page
|
||||
@@ -144,7 +178,7 @@ test('显示中文平台骨架并加载单文件模型', async ({ page }) => {
|
||||
buffer: Buffer.from(SIMPLE_MODEL),
|
||||
});
|
||||
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await page.getByRole('button', { name: '显示设置' }).click();
|
||||
const displayDialog = page.getByRole('dialog', { name: '视图显示设置' });
|
||||
@@ -160,8 +194,9 @@ test('显示中文平台骨架并加载单文件模型', async ({ page }) => {
|
||||
await page.getByText('事件日志').click();
|
||||
await expect(page.getByRole('dialog', { name: '诊断与事件日志' })).toBeVisible();
|
||||
await page.keyboard.press('Escape');
|
||||
await page.getByRole('button', { name: /FPS .*物理/ }).click();
|
||||
await page.getByRole('button', { name: /FPS \d/ }).click();
|
||||
await expect(page.getByRole('dialog', { name: '性能详情' })).toBeVisible();
|
||||
await expect(page.getByRole('dialog', { name: '性能详情' })).toContainText('WASM已加载');
|
||||
await page.keyboard.press('Escape');
|
||||
await page.getByRole('tab', { name: '控制台' }).click();
|
||||
const tools = page.getByRole('tabpanel', { name: '控制台' });
|
||||
@@ -171,7 +206,7 @@ test('显示中文平台骨架并加载单文件模型', async ({ page }) => {
|
||||
await expect(tools.getByText('motor', { exact: true })).toBeVisible();
|
||||
await expect(tools.getByText('关节:slide', { exact: true })).toBeVisible();
|
||||
await page.getByRole('tab', { name: '数据录制' }).click();
|
||||
await expect(page.getByRole('tabpanel', { name: '数据录制' })).toContainText('仿真遥测记录');
|
||||
await expect(page.getByRole('tabpanel', { name: '数据录制' })).toContainText('仅当前会话保存');
|
||||
await expect(page.locator('main canvas')).toBeVisible();
|
||||
await page.getByRole('tab', { name: '检查器' }).click();
|
||||
const structure = page.getByRole('navigation', { name: '模型结构树' });
|
||||
@@ -192,9 +227,9 @@ test('显示中文平台骨架并加载单文件模型', async ({ page }) => {
|
||||
'aria-pressed',
|
||||
'true',
|
||||
);
|
||||
await expect(page.getByText('已忽略').first()).toBeVisible();
|
||||
await expect(page.getByText('已忽略关节限位').first()).toBeVisible();
|
||||
await page.getByRole('button', { name: '重置关节' }).click();
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeVisible();
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled();
|
||||
});
|
||||
|
||||
test('窄视口默认保留完整视口并可按需打开侧栏', async ({ page }) => {
|
||||
@@ -204,9 +239,7 @@ test('窄视口默认保留完整视口并可按需打开侧栏', async ({ page
|
||||
await expect(page.getByRole('button', { name: '显示工程面板' })).toBeVisible();
|
||||
await expect(page.getByRole('button', { name: '显示右侧面板' })).toBeVisible();
|
||||
await page.getByRole('button', { name: '显示右侧面板' }).click();
|
||||
await expect(
|
||||
page.getByRole('complementary').filter({ hasText: '导入模型后显示检查器' }),
|
||||
).toBeVisible();
|
||||
await expect(page.getByRole('complementary').filter({ hasText: '检查器待命' })).toBeVisible();
|
||||
});
|
||||
|
||||
test('工作区布局与视口显示偏好在刷新后保留', async ({ page }) => {
|
||||
@@ -230,7 +263,7 @@ test('转换后的 MJCF 可编辑并重新载入', async ({ page }) => {
|
||||
.locator('input[type="file"]')
|
||||
.first()
|
||||
.setInputFiles({ name: 'model.xml', mimeType: 'text/xml', buffer: Buffer.from(SIMPLE_MODEL) });
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await page.getByRole('button', { name: '源代码' }).click();
|
||||
const dialog = page.getByRole('dialog', { name: '转换后的 MJCF 编辑器' });
|
||||
await expect(dialog).toBeVisible();
|
||||
@@ -245,7 +278,7 @@ test('转换后的 MJCF 可编辑并重新载入', async ({ page }) => {
|
||||
timeout: 30_000,
|
||||
});
|
||||
await dialog.getByRole('button', { name: '关闭源代码编辑器' }).click();
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await expect(page.getByRole('button', { name: '导出 URDF' })).toHaveCount(0);
|
||||
await expect(page.getByRole('button', { name: '导出 MJCF' })).toHaveCount(0);
|
||||
});
|
||||
@@ -256,7 +289,7 @@ test('关闭已修改的 MJCF 前要求确认', async ({ page }) => {
|
||||
.locator('input[type="file"]')
|
||||
.first()
|
||||
.setInputFiles({ name: 'model.xml', mimeType: 'text/xml', buffer: Buffer.from(SIMPLE_MODEL) });
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await page.getByRole('button', { name: '源代码' }).click();
|
||||
const editorDialog = page.getByRole('dialog', { name: '转换后的 MJCF 编辑器' });
|
||||
await editorDialog.locator('.monaco-editor').click({ position: { x: 240, y: 120 } });
|
||||
@@ -287,8 +320,14 @@ test('加载包含 include、OBJ、STL 与 PNG 的工程', async ({ page }) => {
|
||||
fixture('mjcf_include/triangle.stl'),
|
||||
fixture('mjcf_include/checker.png'),
|
||||
]);
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
|
||||
await expect(page.getByText('5 个文件')).toBeVisible();
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await expectImportedFiles(page, [
|
||||
'model.xml',
|
||||
'world.xml',
|
||||
'triangle.obj',
|
||||
'triangle.stl',
|
||||
'checker.png',
|
||||
]);
|
||||
});
|
||||
|
||||
test('加载引用 OBJ 的 URDF 工程', async ({ page }) => {
|
||||
@@ -301,12 +340,14 @@ test('加载引用 OBJ 的 URDF 工程', async ({ page }) => {
|
||||
await expect(options.getByRole('checkbox', { name: /为关节添加驱动器/ })).toBeChecked();
|
||||
await expect(options.getByRole('checkbox', { name: /添加传感器/ })).toBeChecked();
|
||||
await options.getByRole('button', { name: '转换并加载' }).click();
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
|
||||
await expect(page.getByText('2 个文件')).toBeVisible();
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await expectImportedFiles(page, ['robot.urdf', 'triangle.obj']);
|
||||
await expect(page.getByLabel('URDF 处理方式')).toHaveValue('mjcf');
|
||||
await expect(page.getByLabel('URDF 基座类型')).toHaveValue('floating');
|
||||
await page.getByRole('button', { name: '通知中心' }).click();
|
||||
const notifications = page.getByRole('dialog', { name: '通知中心' });
|
||||
for (const summary of await notifications.getByText('事件详情', { exact: true }).all())
|
||||
await summary.click();
|
||||
await expect(notifications).toContainText(/模型已加载 · \d+ 项兼容调整/);
|
||||
await expect(notifications).toContainText(/URDF 已转换为 MJCF(浮动基座),并整体平移/);
|
||||
await page.keyboard.press('Escape');
|
||||
@@ -321,7 +362,7 @@ test('加载引用 OBJ 的 URDF 工程', async ({ page }) => {
|
||||
await expect(
|
||||
sourceDialog.getByRole('button', { name: '保存并重新载入', exact: true }),
|
||||
).toBeDisabled({ timeout: 30_000 });
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible();
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled();
|
||||
await sourceDialog.getByRole('button', { name: '关闭源代码编辑器' }).click();
|
||||
});
|
||||
|
||||
@@ -340,13 +381,15 @@ test('URDF 自动生成的关节驱动器与摄像头可通过 MuJoCo 编译', a
|
||||
.getByRole('dialog', { name: '配置 URDF 仿真组件' })
|
||||
.getByRole('button', { name: '转换并加载' })
|
||||
.click();
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await expect(page.getByLabel('摄像头画面')).toBeVisible();
|
||||
await page.getByRole('button', { name: '隐藏画面' }).click();
|
||||
await page.getByRole('button', { name: '显示摄像头画面' }).click();
|
||||
await expect(page.getByLabel('摄像头画面')).toBeVisible();
|
||||
await page.getByRole('button', { name: '通知中心' }).click();
|
||||
const notifications = page.getByRole('dialog', { name: '通知中心' });
|
||||
for (const summary of await notifications.getByText('事件详情', { exact: true }).all())
|
||||
await summary.click();
|
||||
await expect(notifications).toContainText('已为 1 个 hinge/slide 关节生成 motor 驱动器');
|
||||
await expect(notifications).toContainText('已将 640×480 摄像头固连到 arm');
|
||||
await page.keyboard.press('Escape');
|
||||
@@ -357,8 +400,8 @@ test('URDF 自动生成的关节驱动器与摄像头可通过 MuJoCo 编译', a
|
||||
.click();
|
||||
await expect(page.getByText('shoulder_motor')).toBeVisible();
|
||||
await expect(page.getByText('关节:shoulder')).toBeVisible();
|
||||
await expect(page.getByText('N·m', { exact: true })).toBeVisible();
|
||||
await page.getByText('常用参数').click();
|
||||
await expect(page.locator('output').filter({ hasText: 'N·m' })).toBeVisible();
|
||||
await page.getByText('增益与输出限幅').click();
|
||||
const kp = page.getByLabel(/kp(MJCF stiffness/),
|
||||
kv = page.getByLabel(/kv(MJCF damping/);
|
||||
await kp.fill('150');
|
||||
@@ -387,7 +430,7 @@ test('转换后的 MJCF 保存时保留 DAE 转换缓存资源', async ({ page }
|
||||
.getByRole('dialog', { name: '配置 URDF 仿真组件' })
|
||||
.getByRole('button', { name: '转换并加载' })
|
||||
.click();
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await page.getByRole('button', { name: '源代码' }).click();
|
||||
const dialog = page.getByRole('dialog', { name: '转换后的 MJCF 编辑器' });
|
||||
await dialog.locator('.monaco-editor').click({ position: { x: 240, y: 120 } });
|
||||
@@ -397,7 +440,7 @@ test('转换后的 MJCF 保存时保留 DAE 转换缓存资源', async ({ page }
|
||||
await expect(dialog.getByRole('button', { name: '保存并重新载入', exact: true })).toBeDisabled({
|
||||
timeout: 30_000,
|
||||
});
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible();
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled();
|
||||
await expect(page.getByText('模型编译失败')).toHaveCount(0);
|
||||
});
|
||||
|
||||
@@ -411,7 +454,7 @@ test('slide 关节向屏幕轴正方向拖动时 qpos 同向增加', async ({ pa
|
||||
mimeType: 'text/xml',
|
||||
buffer: Buffer.from(SLIDE_DIRECTION_MODEL),
|
||||
});
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await page.getByRole('button', { name: '关节拖动' }).click();
|
||||
const canvas = page.locator('main canvas').first(),
|
||||
box = await canvas.boundingBox();
|
||||
@@ -439,7 +482,7 @@ test('可导入并启用 Python 控制器', async ({ page }) => {
|
||||
.locator('input[type="file"]')
|
||||
.first()
|
||||
.setInputFiles({ name: 'model.xml', mimeType: 'text/xml', buffer: Buffer.from(SIMPLE_MODEL) });
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await page.getByRole('tab', { name: '控制台' }).click();
|
||||
await page
|
||||
.getByRole('tabpanel', { name: '控制台' })
|
||||
@@ -450,11 +493,13 @@ test('可导入并启用 Python 控制器', async ({ page }) => {
|
||||
.locator('input[accept=".py,text/x-python"]')
|
||||
.setInputFiles({ name: 'balance.py', mimeType: 'text/x-python', buffer: Buffer.from(python) });
|
||||
await expect(page.getByText('测试 PD 控制器', { exact: true })).toBeVisible({ timeout: 30_000 });
|
||||
await page.getByText('脚本详情', { exact: true }).click();
|
||||
await expect(page.getByText('Python / Pyodide')).toBeVisible();
|
||||
await page.getByRole('button', { name: '启用', exact: true }).click();
|
||||
await expect(
|
||||
page.getByRole('tabpanel', { name: '控制台' }).getByRole('button', { name: /Python 脚本控制/ }),
|
||||
).toContainText('运行');
|
||||
await page.screenshot({ path: test.info().outputPath('python-running.png') });
|
||||
});
|
||||
|
||||
test('中等规模模型持续步进并可重复加载', async ({ page }) => {
|
||||
@@ -462,30 +507,33 @@ test('中等规模模型持续步进并可重复加载', async ({ page }) => {
|
||||
const input = page.locator('input[type="file"]').first();
|
||||
const modelFile = { name: 'large.xml', mimeType: 'text/xml', buffer: Buffer.from(LARGE_MODEL) };
|
||||
await input.setInputFiles(modelFile);
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await page.getByRole('button', { name: '▶ 播放' }).click();
|
||||
await page.waitForTimeout(2_000);
|
||||
await expect(page.locator('footer')).not.toContainText('时间 0.000 s');
|
||||
await expect(page.getByLabel('视口状态')).not.toContainText('时间 0.000 s');
|
||||
|
||||
// 播放过程中重置必须同时暂停底层会话,之后仍可正常播放和暂停。
|
||||
await page.getByRole('button', { name: '重置', exact: true }).click();
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeVisible();
|
||||
await expect(page.locator('footer')).toContainText('时间 0.000 s');
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled();
|
||||
await expect(page.getByLabel('视口状态')).toContainText('时间 0.000 s');
|
||||
await page.getByRole('button', { name: '▶ 播放' }).click();
|
||||
await page.waitForTimeout(500);
|
||||
await page.getByRole('button', { name: '⏸ 暂停' }).click();
|
||||
await page.waitForTimeout(200);
|
||||
const pausedTime = (await page.locator('footer').innerText()).match(/时间 ([\d.]+) s/)?.[1];
|
||||
const pausedTime = (await page.getByLabel('视口状态').innerText()).match(/时间 ([\d.]+) s/)?.[1];
|
||||
expect(Number(pausedTime)).toBeGreaterThan(0);
|
||||
await page.waitForTimeout(500);
|
||||
expect((await page.locator('footer').innerText()).match(/时间 ([\d.]+) s/)?.[1]).toBe(pausedTime);
|
||||
expect((await page.getByLabel('视口状态').innerText()).match(/时间 ([\d.]+) s/)?.[1]).toBe(
|
||||
pausedTime,
|
||||
);
|
||||
|
||||
await input.setInputFiles(modelFile);
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await expect(page.getByRole('alert')).toHaveCount(0);
|
||||
});
|
||||
|
||||
test('认证资产可点击创建场景并拖到画布落位', async ({ page }) => {
|
||||
await page.setViewportSize({ width: 1440, height: 900 });
|
||||
await page.goto('/');
|
||||
await page
|
||||
.locator('input[type="file"]')
|
||||
@@ -521,9 +569,11 @@ test('认证资产可点击创建场景并拖到画布落位', async ({ page })
|
||||
const beforeRotation = await gizmoLine.getAttribute('x2');
|
||||
const canvasBox = await page.locator('main canvas').first().boundingBox();
|
||||
expect(canvasBox).not.toBeNull();
|
||||
await page.mouse.move(canvasBox!.x + 24, canvasBox!.y + 24);
|
||||
// 左上现在是可交互状态槽位;从画布空白中部拖动,继续验证真实 OrbitControls。
|
||||
const dragY = canvasBox!.y + canvasBox!.height * 0.35;
|
||||
await page.mouse.move(canvasBox!.x + 24, dragY);
|
||||
await page.mouse.down({ button: 'left' });
|
||||
await page.mouse.move(canvasBox!.x + 104, canvasBox!.y + 50, { steps: 8 });
|
||||
await page.mouse.move(canvasBox!.x + 104, dragY + 26, { steps: 8 });
|
||||
await page.mouse.up({ button: 'left' });
|
||||
await expect.poll(() => gizmoLine.getAttribute('x2')).not.toBe(beforeRotation);
|
||||
await expect(page.getByLabel('对象名称')).toHaveCount(0);
|
||||
@@ -545,16 +595,16 @@ test('认证资产自动打开地图属性并与参数地形一次编译', async
|
||||
|
||||
const library = page.getByLabel('地图资产库');
|
||||
await library.getByRole('button', { name: '添加基础方盒' }).click();
|
||||
await expect(page.getByText('Map / Object')).toBeVisible();
|
||||
await expect(page.getByLabel('地图物体检查器')).toBeVisible();
|
||||
await expect(page.getByLabel('地图物体检查器')).toBeVisible();
|
||||
await expect(page.getByText('1 项场景更改待应用')).toBeVisible();
|
||||
|
||||
const sceneTree = page.getByLabel('场景资产树');
|
||||
await expect(sceneTree.getByText('基础方盒')).toBeVisible();
|
||||
await sceneTree.getByRole('treeitem', { name: 'box', exact: true }).click();
|
||||
await expect(page.getByText('Robot / Body')).toBeVisible();
|
||||
await expect(page.getByText(/^Body #/)).toBeVisible();
|
||||
await sceneTree.getByRole('treeitem', { name: /基础方盒/ }).click();
|
||||
await expect(page.getByText('Map / Object')).toBeVisible();
|
||||
await expect(page.getByLabel('地图物体检查器')).toBeVisible();
|
||||
|
||||
await library.getByRole('button', { name: '添加随机粗糙地形' }).click();
|
||||
await page.getByLabel('位置 X(m)').fill('4');
|
||||
@@ -585,8 +635,13 @@ test('认证资产自动打开地图属性并与参数地形一次编译', async
|
||||
await page.getByRole('button', { name: '应用场景' }).click();
|
||||
|
||||
await expect(page.getByText('2 项场景更改待应用')).toHaveCount(0, { timeout: 30_000 });
|
||||
await expect(page.getByText(/已加载工程地图“场景 1”(1 个物理几何/)).toBeVisible();
|
||||
await expect(page.getByText(/已加载随机粗糙地形物理地图/)).toBeVisible();
|
||||
await page.getByRole('button', { name: '通知中心' }).click();
|
||||
const notifications = page.getByRole('dialog', { name: '通知中心' });
|
||||
for (const summary of await notifications.getByText('事件详情', { exact: true }).all())
|
||||
await summary.click();
|
||||
await expect(notifications).toContainText(/已加载工程地图“场景 1”(1 个物理几何/);
|
||||
await expect(notifications).toContainText(/已加载随机粗糙地形物理地图/);
|
||||
await page.keyboard.press('Escape');
|
||||
await expect(page.getByLabel('地图来源')).toHaveValue('builtin');
|
||||
|
||||
// 当前仍选中参数地形时,直接点已编译的认证资产也必须反查场景与对象,首次点击即挂载操纵器。
|
||||
@@ -682,8 +737,7 @@ test('工程地图与参数地形共享放置草稿、实例变换和回滚入
|
||||
await expect(page.getByText('2 项场景更改待应用')).toBeVisible();
|
||||
await page.getByRole('button', { name: '应用场景' }).click();
|
||||
await expect(page.getByText('2 项场景更改待应用')).toHaveCount(0, { timeout: 30_000 });
|
||||
await expect(page.getByText(/已加载工程地图“草稿仓库”/)).toBeVisible();
|
||||
await expect(page.getByText(/已加载波浪地形物理地图/)).toBeVisible();
|
||||
await expectNotificationDetails(page, [/已加载工程地图“草稿仓库”/, /已加载波浪地形物理地图/]);
|
||||
|
||||
await page.getByRole('button', { name: '删除地图实例 草稿仓库' }).click();
|
||||
await expect(page.getByText('1 项场景更改待应用')).toBeVisible();
|
||||
@@ -739,23 +793,26 @@ test('参数化地形可连续拖到画布并一次性编译', async ({ page })
|
||||
await expect(page.getByLabel('场景资产树')).toContainText('波浪地形');
|
||||
|
||||
await page.getByLabel('场景资产树').getByRole('treeitem', { name: 'box' }).click();
|
||||
await expect(page.getByText('Robot / Body')).toBeVisible();
|
||||
await expect(page.getByText(/^Body #/)).toBeVisible();
|
||||
const canvasBox = await page.locator('main canvas').first().boundingBox();
|
||||
expect(canvasBox).not.toBeNull();
|
||||
await page.mouse.click(
|
||||
canvasBox!.x + canvasBox!.width * 0.5,
|
||||
canvasBox!.y + canvasBox!.height * 0.78,
|
||||
);
|
||||
await expect(page.getByText('Map / Instance')).toBeVisible();
|
||||
await expect(page.getByText('参数化地形')).toBeVisible();
|
||||
await expect(page.getByLabel('地图视口工具')).toBeVisible();
|
||||
await expect(page.getByText('地图与对象属性')).toBeVisible();
|
||||
await expect(page.getByLabel('物理地图预设')).toBeVisible();
|
||||
|
||||
await page.getByRole('button', { name: '应用场景' }).click();
|
||||
await expect(page.getByText('2 项场景更改待应用')).toHaveCount(0, { timeout: 30_000 });
|
||||
await expect(page.getByText(/已加载随机粗糙地形物理地图/)).toBeVisible({
|
||||
timeout: 30_000,
|
||||
});
|
||||
await expect(page.getByText(/已加载波浪地形物理地图/)).toBeVisible();
|
||||
await page.getByRole('button', { name: '通知中心' }).click();
|
||||
const notifications = page.getByRole('dialog', { name: '通知中心' });
|
||||
for (const summary of await notifications.getByText('事件详情', { exact: true }).all())
|
||||
await summary.click();
|
||||
await expect(notifications).toContainText(/已加载随机粗糙地形物理地图/);
|
||||
await expect(notifications).toContainText(/已加载波浪地形物理地图/);
|
||||
await page.keyboard.press('Escape');
|
||||
});
|
||||
|
||||
test('应用内置 MJCF 楼梯物理地图', async ({ page }) => {
|
||||
@@ -769,13 +826,14 @@ test('应用内置 MJCF 楼梯物理地图', async ({ page }) => {
|
||||
mimeType: 'text/xml',
|
||||
buffer: Buffer.from(SIMPLE_MODEL),
|
||||
});
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
|
||||
await page.getByLabel('地图来源').selectOption('builtin');
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await page.getByLabel('地图资产库').getByRole('button', { name: '添加波浪地形' }).click();
|
||||
await page.getByLabel('物理地图预设').selectOption('stairs');
|
||||
await page.getByLabel('台阶数量').fill('6');
|
||||
await page.getByRole('button', { name: '应用并重新编译' }).click();
|
||||
await expect(page.getByText(/已加载楼梯物理地图/)).toBeVisible({ timeout: 30_000 });
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible();
|
||||
await expect(page.getByLabel('地图草稿状态')).toBeHidden({ timeout: 30_000 });
|
||||
await expectNotificationDetails(page, [/已加载楼梯物理地图/]);
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled();
|
||||
});
|
||||
|
||||
test('依次应用全部系统参数化地形', async ({ page }) => {
|
||||
@@ -800,19 +858,18 @@ test('依次应用全部系统参数化地形', async ({ page }) => {
|
||||
mimeType: 'text/xml',
|
||||
buffer: Buffer.from(SIMPLE_MODEL),
|
||||
});
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
|
||||
await page.getByLabel('地图来源').selectOption('builtin');
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await page.getByLabel('地图资产库').getByRole('button', { name: '添加波浪地形' }).click();
|
||||
for (const [preset, label] of terrains) {
|
||||
await page.getByLabel('物理地图预设').selectOption(preset);
|
||||
await page.getByLabel('地形边长(m)').fill('6');
|
||||
if (preset === 'rough' || preset === 'wave')
|
||||
await page.getByLabel('水平采样间距(m)').fill('0.25');
|
||||
await page.getByRole('button', { name: '应用并重新编译' }).click();
|
||||
await expect(page.getByText(new RegExp(`已加载${label}物理地图`))).toBeVisible({
|
||||
timeout: 30_000,
|
||||
});
|
||||
await expect(page.getByLabel('地图草稿状态')).toBeHidden({ timeout: 30_000 });
|
||||
await expectNotificationDetails(page, [new RegExp(`已加载${label}物理地图`)]);
|
||||
}
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible();
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled();
|
||||
});
|
||||
|
||||
test('导入并应用分层工程地图包', async ({ page }) => {
|
||||
@@ -843,11 +900,18 @@ test('导入并应用分层工程地图包', async ({ page }) => {
|
||||
mimeType: 'application/zip',
|
||||
buffer: Buffer.from(project),
|
||||
});
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
|
||||
await page.getByLabel('地图来源').selectOption({ label: '测试场景' });
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await page
|
||||
.getByLabel('地图资产库')
|
||||
.getByRole('button', { name: '放置工程地图 测试场景' })
|
||||
.click();
|
||||
// 资产库放置默认不移动机器人;显式选择出生点后仍验证原编译路径。
|
||||
await expect(page.getByLabel('地图出生点')).toHaveValue('');
|
||||
await page.getByLabel('地图出生点').selectOption('start');
|
||||
await expect(page.getByLabel('地图出生点')).toHaveValue('start');
|
||||
await page.getByRole('button', { name: '应用并重新编译' }).click();
|
||||
await expect(page.getByText(/已加载工程地图“测试场景”/)).toBeVisible({ timeout: 30_000 });
|
||||
await expect(page.getByLabel('地图草稿状态')).toBeHidden({ timeout: 30_000 });
|
||||
await expectNotificationDetails(page, [/已加载工程地图“测试场景”/], [/视觉地图加载失败/]);
|
||||
await expect(page.getByText(/视觉地图加载失败/)).toHaveCount(0);
|
||||
});
|
||||
|
||||
@@ -877,10 +941,13 @@ test('将受支持的只读物理地图转换为可编辑副本', async ({ page
|
||||
mimeType: 'application/zip',
|
||||
buffer: Buffer.from(project),
|
||||
});
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
|
||||
await page.getByLabel('地图来源').selectOption({ label: '旧版基础场景' });
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await page
|
||||
.getByLabel('地图资产库')
|
||||
.getByRole('button', { name: '放置工程地图 旧版基础场景' })
|
||||
.click();
|
||||
await page.getByRole('button', { name: '应用并重新编译' }).click();
|
||||
const editor = page.getByText('认证资产与场景对象属性').locator('..');
|
||||
const editor = page.getByText('源内容修改影响所有同源实例').locator('../..');
|
||||
await expect(editor.getByText(/保持只读/)).toBeVisible({ timeout: 30_000 });
|
||||
await editor.getByRole('button', { name: '创建可编辑副本' }).click();
|
||||
await expect(page.getByText('已创建可编辑地图副本')).toBeVisible({ timeout: 30_000 });
|
||||
@@ -924,10 +991,15 @@ test('编辑 V3 地图对象并事务式应用', async ({ page }) => {
|
||||
mimeType: 'application/zip',
|
||||
buffer: Buffer.from(project),
|
||||
});
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
|
||||
await page.getByLabel('地图来源').selectOption({ label: '可编辑场景' });
|
||||
await page.getByRole('button', { name: '应用并重新编译' }).click();
|
||||
await expect(page.getByText('认证资产与场景对象属性')).toBeVisible({ timeout: 30_000 });
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await page
|
||||
.getByLabel('地图资产库')
|
||||
.getByRole('button', { name: '放置工程地图 可编辑场景' })
|
||||
.click();
|
||||
await page.getByRole('button', { name: '应用场景', exact: true }).click();
|
||||
await expect(page.getByText('源内容修改影响所有同源实例', { exact: true })).toBeVisible({
|
||||
timeout: 30_000,
|
||||
});
|
||||
const mapTools = page.getByLabel('地图视口工具');
|
||||
await expect(mapTools.getByRole('button', { name: '移动工具 W' })).toHaveAttribute(
|
||||
'aria-pressed',
|
||||
@@ -955,7 +1027,7 @@ test('编辑 V3 地图对象并事务式应用', async ({ page }) => {
|
||||
).toBeVisible({
|
||||
timeout: 30_000,
|
||||
});
|
||||
await expect(draftStatus).toContainText('地图草稿已同步');
|
||||
await expect(draftStatus).toBeHidden();
|
||||
|
||||
await page
|
||||
.getByLabel('地图对象列表')
|
||||
@@ -969,13 +1041,13 @@ test('编辑 V3 地图对象并事务式应用', async ({ page }) => {
|
||||
.getByRole('button', { name: /box · 方盒/ })
|
||||
.click();
|
||||
await expect(page.getByLabel('对象位置X')).toHaveValue('2');
|
||||
await expect(draftStatus).toContainText('地图草稿已同步');
|
||||
await expect(draftStatus).toBeHidden();
|
||||
|
||||
await page.getByRole('button', { name: '新增', exact: true }).click();
|
||||
await expect(page.getByLabel('地图对象列表').getByRole('button')).toHaveCount(2);
|
||||
await page.keyboard.press('Delete');
|
||||
await expect(page.getByLabel('地图对象列表').getByRole('button')).toHaveCount(1);
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible();
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled();
|
||||
});
|
||||
|
||||
test('工程地图编译失败时保留上一仿真会话', async ({ page }) => {
|
||||
@@ -1004,34 +1076,40 @@ test('工程地图编译失败时保留上一仿真会话', async ({ page }) =>
|
||||
mimeType: 'application/zip',
|
||||
buffer: Buffer.from(project),
|
||||
});
|
||||
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
await page.getByRole('button', { name: '▶ 播放' }).click();
|
||||
await page.waitForTimeout(200);
|
||||
const runningTime = Number(
|
||||
(await page.locator('footer').innerText()).match(/时间 ([\d.]+) s/)?.[1] ?? 0,
|
||||
(await page.getByLabel('视口状态').innerText()).match(/时间 ([\d.]+) s/)?.[1] ?? 0,
|
||||
);
|
||||
await page.getByLabel('地图来源').selectOption({ label: '动态错误地图' });
|
||||
await page.getByRole('button', { name: '应用并重新编译' }).click();
|
||||
await page
|
||||
.getByLabel('地图资产库')
|
||||
.getByRole('button', { name: '放置工程地图 动态错误地图' })
|
||||
.click();
|
||||
await page.getByRole('button', { name: '应用场景', exact: true }).click();
|
||||
await expect(page.getByRole('alert')).toContainText('模型编译失败', { timeout: 30_000 });
|
||||
await expect(page.getByLabel('地图来源')).toHaveValue('project:maps/invalid/map.json');
|
||||
await expect(page.getByLabel('地图草稿状态')).toContainText('应用失败');
|
||||
await page.screenshot({ path: test.info().outputPath('map-compile-failed.png') });
|
||||
await expect(page.getByText('1 项场景更改待应用')).toBeVisible();
|
||||
await expect(page.getByLabel('场景资产树')).toContainText('待应用');
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeVisible();
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled();
|
||||
await page.getByRole('button', { name: '关闭错误' }).click();
|
||||
await page.getByRole('button', { name: '▶ 播放' }).click();
|
||||
await expect
|
||||
.poll(async () =>
|
||||
Number((await page.locator('footer').innerText()).match(/时间 ([\d.]+) s/)?.[1] ?? 0),
|
||||
Number((await page.getByLabel('视口状态').innerText()).match(/时间 ([\d.]+) s/)?.[1] ?? 0),
|
||||
)
|
||||
.toBeGreaterThan(runningTime);
|
||||
await page.getByRole('button', { name: '放弃更改' }).click();
|
||||
await expect(page.getByText('1 项场景更改待应用')).toHaveCount(0);
|
||||
await expect(page.getByLabel('地图来源')).toHaveValue('none');
|
||||
await expect(page.getByLabel('场景资产树')).not.toContainText('动态错误地图');
|
||||
await expect(page.getByLabel('地图草稿状态')).toBeHidden();
|
||||
});
|
||||
|
||||
test('无效模型显示中文诊断且保留工程树', async ({ page }) => {
|
||||
await page.goto('/');
|
||||
await page.locator('input[type="file"]').first().setInputFiles(fixture('invalid.xml'));
|
||||
await expect(page.getByRole('alert')).toContainText('模型编译失败', { timeout: 30_000 });
|
||||
await expect(page.getByText('invalid.xml', { exact: false }).first()).toBeVisible();
|
||||
await expectImportedFiles(page, ['invalid.xml']);
|
||||
});
|
||||
|
||||
@@ -0,0 +1,158 @@
|
||||
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(/<mesh\b[^>]*\/>/g, '')
|
||||
.replace(/<geom\b[^>]*\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.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ 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();
|
||||
await tools.getByText('推理详情', { exact: true }).click();
|
||||
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();
|
||||
await section.evaluate((element) => element.scrollIntoView({ block: 'start' }));
|
||||
await page.screenshot({ path: test.info().outputPath('onnx-policy.png') });
|
||||
const targetRow = tools.getByText('当前目标 (X, Y) m', { 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 page.screenshot({ path: test.info().outputPath('onnx-contract-error.png') });
|
||||
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 page.screenshot({ path: test.info().outputPath('training-metrics.png') });
|
||||
await tools.getByRole('button', { name: /训练指标趋势/ }).click();
|
||||
await expect(tools.locator('.uplot')).toHaveCount(0);
|
||||
});
|
||||
@@ -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(/<mesh\b[^>]*\/>/g, '')
|
||||
.replace(/<geom\b[^>]*\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.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ 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) m', { 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.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ 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.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ 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: '▶ 播放' })).toBeEnabled();
|
||||
});
|
||||
}
|
||||
|
||||
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 =
|
||||
'<actuator>' +
|
||||
fixture.jointNames
|
||||
.map((name: string) => `<motor name="${name}_motor" joint="${name}"/>`)
|
||||
.join('') +
|
||||
'</actuator>';
|
||||
await page.goto('/');
|
||||
await page
|
||||
.locator('input[type="file"]')
|
||||
.first()
|
||||
.setInputFiles({
|
||||
name: 'go2.xml',
|
||||
mimeType: 'text/xml',
|
||||
buffer: Buffer.from(go2.replace('</mujoco>', actuators + '</mujoco>')),
|
||||
});
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ 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);
|
||||
});
|
||||
@@ -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(/<mesh\b[^>]*\/>/g, '')
|
||||
.replace(/<geom\b[^>]*\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(/<mesh\b[^>]*\/>/g, '')
|
||||
.replace(/<geom\b[^>]*\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.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ 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);
|
||||
});
|
||||
@@ -0,0 +1,179 @@
|
||||
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<void>((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<void>((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 ui.getByText('上传基础策略', { exact: true }).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 ui.getByText('来源与校验详情', { exact: true }).click();
|
||||
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<string, unknown> | 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<string, unknown>;
|
||||
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<void>((done) => {
|
||||
if (backend.exitCode !== null) done();
|
||||
else backend.once('exit', () => done());
|
||||
});
|
||||
writeFileSync(join(root, 'browser-server.log'), log);
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,99 @@
|
||||
import type { TuningSession } from '../src/training/types';
|
||||
const objectives = {
|
||||
velocity_tracking: 0.35,
|
||||
action_smoothness: 0.2,
|
||||
posture_stability: 0.15,
|
||||
fall_avoidance: 0.15,
|
||||
foot_slip: 0.1,
|
||||
energy: 0.05,
|
||||
};
|
||||
|
||||
export function tuningFixture(): TuningSession {
|
||||
const baseConfig = { weights: { pose: 1 }, params: {} };
|
||||
const resultConfig = { weights: { pose: 1.1 }, params: {} };
|
||||
return {
|
||||
id: 'a'.repeat(32),
|
||||
state: 'awaiting_approval',
|
||||
mode: 'approval',
|
||||
createdAt: '2026-01-01T00:00:00Z',
|
||||
updatedAt: '2026-01-01T00:02:00Z',
|
||||
config: {
|
||||
taskId: 'Unitree-Go2-Flat',
|
||||
mode: 'approval',
|
||||
runName: '界面验收会话',
|
||||
numEnvs: 16,
|
||||
seed: 42,
|
||||
gpuIds: [0],
|
||||
trialCount: 2,
|
||||
initialIterations: 300,
|
||||
middleIterations: 900,
|
||||
finalIterations: 2000,
|
||||
evalNumEnvs: 8,
|
||||
evalSteps: 10,
|
||||
objectiveWeights: objectives,
|
||||
fallbackEnabled: false,
|
||||
rungs: [300, 900, 2000],
|
||||
promote: [2, 2, 2],
|
||||
},
|
||||
objectiveWeights: objectives,
|
||||
message: 'paused',
|
||||
bestTrialId: '2'.repeat(32),
|
||||
consecutiveNoImprove: 0,
|
||||
fallbackEnabled: false,
|
||||
trials: [
|
||||
{
|
||||
id: '1'.repeat(32),
|
||||
sessionId: 'a'.repeat(32),
|
||||
number: 0,
|
||||
state: 'completed',
|
||||
rung: 0,
|
||||
targetIterations: 300,
|
||||
rewardConfig: baseConfig,
|
||||
score: 0,
|
||||
eligible: true,
|
||||
createdAt: '2026-01-01T00:00:00Z',
|
||||
message: 'done',
|
||||
},
|
||||
{
|
||||
id: '2'.repeat(32),
|
||||
sessionId: 'a'.repeat(32),
|
||||
number: 1,
|
||||
state: 'completed',
|
||||
rung: 0,
|
||||
targetIterations: 300,
|
||||
rewardConfig: resultConfig,
|
||||
proposalId: '3'.repeat(32),
|
||||
score: 0.12,
|
||||
eligible: true,
|
||||
evaluation: {
|
||||
metrics: { linear_velocity_rmse: 0.2 },
|
||||
score: { score: 0.12, eligible: true, components: { velocity_tracking: 0.2 } },
|
||||
},
|
||||
createdAt: '2026-01-01T00:01:00Z',
|
||||
message: 'done',
|
||||
},
|
||||
],
|
||||
proposals: [
|
||||
{
|
||||
id: '3'.repeat(32),
|
||||
sessionId: 'a'.repeat(32),
|
||||
baseTrialId: '1'.repeat(32),
|
||||
state: 'pending',
|
||||
source: 'agent',
|
||||
patch: { weights: { pose: 1.1 }, params: {} },
|
||||
rationale: '提高姿态奖励以降低躯干倾角。',
|
||||
expectedImpact: { posture_stability: '姿态误差预计下降 8%' },
|
||||
confidence: 0.82,
|
||||
createdAt: '2026-01-01T00:00:30Z',
|
||||
},
|
||||
],
|
||||
audit: [],
|
||||
control: {
|
||||
runPolicy: 'step',
|
||||
dispatchTokens: 0,
|
||||
constraintsRevision: 0,
|
||||
constraints: {},
|
||||
effectiveAfterCurrent: false,
|
||||
},
|
||||
};
|
||||
}
|
||||
@@ -0,0 +1,114 @@
|
||||
import { expect, test } from '@playwright/test';
|
||||
|
||||
const MODEL = `<mujoco model="hud"><worldbody><light pos="0 0 3"/>
|
||||
<geom type="plane" size="3 3 .1"/><body pos="0 0 1"><freejoint/>
|
||||
<geom type="box" size=".2 .2 .2" mass="1"/></body></worldbody></mujoco>`;
|
||||
|
||||
for (const theme of ['dark', 'light']) {
|
||||
test(`Cyber HUD 材质与可读性 ${theme}`, async ({ page }, info) => {
|
||||
await page.setViewportSize({ width: 1440, height: 900 });
|
||||
await page.addInitScript((value) => {
|
||||
if (!localStorage.getItem('mujoco-platform-theme')) {
|
||||
localStorage.setItem('mujoco-platform-theme', value);
|
||||
}
|
||||
}, theme);
|
||||
await page.goto('/');
|
||||
const welcome = page.getByRole('region', { name: '导入模型工程' });
|
||||
await expect(welcome).toBeVisible();
|
||||
await expect(page.locator('.cyber-workspace')).toHaveClass(new RegExp(`theme-${theme}`));
|
||||
await expect(welcome).toHaveCSS('backdrop-filter', 'blur(12px) saturate(1.2)');
|
||||
const primary = page.getByRole('button', { name: '选择文件', exact: true });
|
||||
await primary.focus();
|
||||
await expect(primary).toBeFocused();
|
||||
await expect(primary).toHaveCSS('outline-style', 'solid');
|
||||
// 在真实浏览器中验证主按钮与正文的计算色对比度,而非仅比较 token。
|
||||
const contrast = await primary.evaluate((button) => {
|
||||
function luminance(color: string) {
|
||||
const rgb = color
|
||||
.match(/[\d.]+/g)!
|
||||
.slice(0, 3)
|
||||
.map(Number);
|
||||
const linear = rgb.map((value) => {
|
||||
const v = value / 255;
|
||||
return v <= 0.04045 ? v / 12.92 : ((v + 0.055) / 1.055) ** 2.4;
|
||||
});
|
||||
return linear[0] * 0.2126 + linear[1] * 0.7152 + linear[2] * 0.0722;
|
||||
}
|
||||
const style = getComputedStyle(button);
|
||||
const a = luminance(style.color),
|
||||
b = luminance(style.backgroundColor);
|
||||
return (Math.max(a, b) + 0.05) / (Math.min(a, b) + 0.05);
|
||||
});
|
||||
expect(contrast).toBeGreaterThanOrEqual(4.5);
|
||||
const disabled = page.getByRole('button', { name: '源代码', exact: true });
|
||||
await expect(disabled).toBeDisabled();
|
||||
await expect(disabled).toHaveCSS('box-shadow', 'none');
|
||||
await page.screenshot({
|
||||
path: info.outputPath(`workbench-${theme}.png`),
|
||||
animations: 'disabled',
|
||||
});
|
||||
await page.goto('/tuning.html');
|
||||
await expect(page.locator('.cyber-workspace')).toHaveClass(new RegExp(`theme-${theme}`));
|
||||
await page
|
||||
.getByRole('button', { name: `切换到${theme === 'dark' ? '白天' : '黑夜'}主题` })
|
||||
.click();
|
||||
await page.reload();
|
||||
await expect(page.locator('.cyber-workspace')).toHaveClass(
|
||||
new RegExp(`theme-${theme === 'dark' ? 'light' : 'dark'}`),
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
test('运行状态呼吸遵循减少动态效果,装饰不拦截交互', async ({ page }) => {
|
||||
await page.goto('/');
|
||||
await page.locator('#mujoco-project-files').setInputFiles({
|
||||
name: 'hud.xml',
|
||||
mimeType: 'text/xml',
|
||||
buffer: Buffer.from(MODEL),
|
||||
});
|
||||
const play = page.getByRole('button', { name: '▶ 播放' });
|
||||
await expect(play).toBeEnabled({ timeout: 30_000 });
|
||||
await play.click();
|
||||
const signal = page.locator('.hud-signal');
|
||||
await expect(signal).toHaveAttribute('data-running', 'true');
|
||||
const animation = () => signal.evaluate((el) => getComputedStyle(el, '::after').animationName);
|
||||
expect(await animation()).toBe('hud-breathe');
|
||||
expect(await signal.evaluate((el) => getComputedStyle(el, '::after').pointerEvents)).toBe('none');
|
||||
await page.emulateMedia({ reducedMotion: 'reduce' });
|
||||
await expect.poll(animation).toBe('none');
|
||||
await page.getByRole('button', { name: '⏸ 暂停' }).click();
|
||||
await expect(signal).toHaveAttribute('data-running', 'false');
|
||||
await page.emulateMedia({ forcedColors: 'active' });
|
||||
await expect(page.locator('.engineering-glass').first()).toHaveCSS('backdrop-filter', 'none');
|
||||
});
|
||||
|
||||
for (const theme of ['dark', 'light']) {
|
||||
test(`无模糊能力时使用实色回退 ${theme}`, async ({ page }) => {
|
||||
// Chromium 支持模糊:使构建 CSS 的能力分支为 false,验证回退规则本身。
|
||||
let replaced = false;
|
||||
await page.route('**/*.css', async (route) => {
|
||||
const response = await route.fetch();
|
||||
const css = await response.text();
|
||||
const fallback = css.replace(/backdrop-filter:\s*blur\(1px\)/g, 'cyber-unsupported: none');
|
||||
replaced ||= css !== fallback;
|
||||
await route.fulfill({ response, body: fallback });
|
||||
});
|
||||
await page.addInitScript(
|
||||
(value) => localStorage.setItem('mujoco-platform-theme', value),
|
||||
theme,
|
||||
);
|
||||
await page.goto('/');
|
||||
const welcome = page.getByRole('region', { name: '导入模型工程' });
|
||||
await expect(welcome).toBeVisible();
|
||||
expect(replaced).toBe(true);
|
||||
await expect(welcome).toHaveCSS('backdrop-filter', 'none');
|
||||
await expect(welcome).toHaveCSS(
|
||||
'background-color',
|
||||
theme === 'dark' ? 'rgb(11, 18, 36)' : 'rgb(251, 252, 254)',
|
||||
);
|
||||
await page.getByRole('button', { name: '打开命令面板' }).click();
|
||||
await expect(page.getByRole('dialog', { name: '命令面板' })).toBeVisible();
|
||||
await page.keyboard.press('Escape');
|
||||
await expect(page.getByRole('dialog', { name: '命令面板' })).toBeHidden();
|
||||
});
|
||||
}
|
||||
@@ -0,0 +1,191 @@
|
||||
import { expect, test, type Page } from '@playwright/test';
|
||||
|
||||
const MODEL = `<mujoco model="domain-panels"><worldbody><geom type="plane" size="3 3 .1"/><body name="box" pos="0 0 1"><joint name="slide" type="slide" axis="1 0 0" range="-1 1"/><geom type="box" size=".2 .2 .2" mass="1"/></body></worldbody><actuator><motor name="motor" joint="slide" ctrlrange="-2 2"/></actuator></mujoco>`;
|
||||
async function load(page: Page) {
|
||||
await page.goto('/');
|
||||
await page
|
||||
.locator('#mujoco-project-files')
|
||||
.setInputFiles({ name: 'domain.xml', mimeType: 'text/xml', buffer: Buffer.from(MODEL) });
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled();
|
||||
}
|
||||
async function left(page: Page) {
|
||||
const button = page.getByRole('button', { name: '显示工程面板', exact: true });
|
||||
if (await button.isVisible()) await button.click();
|
||||
}
|
||||
async function right(page: Page) {
|
||||
const button = page.getByRole('button', { name: '显示右侧面板', exact: true });
|
||||
if (await button.isVisible()) await button.click();
|
||||
}
|
||||
async function screenshot(page: Page, name: string) {
|
||||
expect(await page.evaluate(() => document.documentElement.scrollWidth <= innerWidth)).toBe(true);
|
||||
await page.screenshot({
|
||||
path: test.info().outputPath(name + '.png'),
|
||||
mask: [page.getByLabel('视口状态')],
|
||||
});
|
||||
}
|
||||
|
||||
for (const theme of ['dark', 'light']) {
|
||||
for (const [width, height] of [
|
||||
[1920, 1080],
|
||||
[1440, 900],
|
||||
[1366, 768],
|
||||
[1024, 768],
|
||||
[768, 800],
|
||||
]) {
|
||||
test(`领域检查器与录制 ${theme} ${width}`, async ({ page }) => {
|
||||
const errors: string[] = [];
|
||||
page.on('pageerror', (error) => errors.push(error.message));
|
||||
await page.setViewportSize({ width, height });
|
||||
await page.addInitScript(
|
||||
(value) => localStorage.setItem('mujoco-platform-theme', value),
|
||||
theme,
|
||||
);
|
||||
await load(page);
|
||||
await left(page);
|
||||
await page
|
||||
.getByRole('navigation', { name: '模型结构树' })
|
||||
.getByRole('treeitem', { name: /box/ })
|
||||
.first()
|
||||
.click();
|
||||
await expect(page.getByText(/^Body #/)).toBeVisible();
|
||||
await expect(page.getByRole('button', { name: '标识与层级' })).toHaveAttribute(
|
||||
'aria-expanded',
|
||||
'false',
|
||||
);
|
||||
await screenshot(page, 'body');
|
||||
await page
|
||||
.getByRole('tabpanel', { name: '检查器' })
|
||||
.getByRole('button', { name: /slide/ })
|
||||
.click();
|
||||
await expect(page.getByRole('slider', { name: 'slide' })).toBeVisible();
|
||||
await page.getByRole('slider', { name: 'slide' }).focus();
|
||||
await page.keyboard.press('ArrowRight');
|
||||
await expect(page.getByRole('slider', { name: 'slide' })).not.toHaveValue('0');
|
||||
await page.getByText('增益与输出限幅', { exact: true }).click();
|
||||
await expect(page.getByText(/参数立即作用于模型/)).toBeVisible();
|
||||
await screenshot(page, 'joint-actuator');
|
||||
await page.getByRole('tab', { name: '数据录制' }).click();
|
||||
await page.getByRole('button', { name: '开始记录' }).click();
|
||||
await expect(page.getByText('记录中', { exact: true })).toBeVisible();
|
||||
await expect(page.getByLabel('采样频率', { exact: true })).toBeDisabled();
|
||||
await page.getByRole('button', { name: '▶ 播放' }).click();
|
||||
await page.waitForTimeout(150);
|
||||
await page.getByRole('button', { name: '⏸ 暂停' }).click();
|
||||
await page.getByRole('button', { name: '停止记录' }).click();
|
||||
await page.getByText('实时运动状态', { exact: true }).click();
|
||||
await expect(page.getByRole('button', { name: '导出 CSV' })).toBeEnabled();
|
||||
await screenshot(page, 'recording');
|
||||
await left(page);
|
||||
await page.getByLabel('地图资产库').getByRole('button', { name: '添加基础方盒' }).click();
|
||||
await page.getByRole('tab', { name: '检查器' }).click();
|
||||
await expect(page.getByLabel('对象参数sizeX', { exact: true })).toBeVisible();
|
||||
await expect(page.getByText('表面材质', { exact: true }).locator('..')).not.toHaveAttribute(
|
||||
'open',
|
||||
);
|
||||
await page.getByLabel('对象参数sizeX', { exact: true }).fill('1.5');
|
||||
await page.getByLabel('对象参数sizeX', { exact: true }).press('Enter');
|
||||
await expect(page.getByLabel('地图草稿状态')).toContainText('未保存');
|
||||
await screenshot(page, 'map-draft');
|
||||
if (width === 1440) {
|
||||
await page
|
||||
.getByLabel('对象参数sizeX', { exact: true })
|
||||
.locator('..')
|
||||
.getByRole('button')
|
||||
.focus();
|
||||
await expect(page.getByRole('tooltip')).toContainText('Shift 精调');
|
||||
const tip = await page.getByRole('tooltip').boundingBox();
|
||||
expect(tip!.x + tip!.width).toBeLessThanOrEqual(width);
|
||||
await screenshot(page, 'map-keyboard-tooltip');
|
||||
await page.keyboard.press('Escape');
|
||||
await expect(page.getByRole('tooltip')).toHaveCount(0);
|
||||
}
|
||||
expect(errors).toEqual([]);
|
||||
});
|
||||
}
|
||||
|
||||
test(`训练折叠详情与跨视图状态 ${theme}`, async ({ page }) => {
|
||||
await page.setViewportSize({ width: 1440, height: 900 });
|
||||
await page.addInitScript(
|
||||
(value) => localStorage.setItem('mujoco-platform-theme', value),
|
||||
theme,
|
||||
);
|
||||
let starts = 0;
|
||||
let offline = true;
|
||||
const job = {
|
||||
id: 'a'.repeat(32),
|
||||
taskId: 'Unitree-Go2-Flat',
|
||||
state: 'running',
|
||||
progress: 0.4,
|
||||
iteration: 4,
|
||||
maxIterations: 10,
|
||||
message: '模拟训练中,未启动真实作业',
|
||||
artifactReady: false,
|
||||
logs: ['Learning iteration 4 / 10', 'Mean value loss: 0.9'],
|
||||
};
|
||||
await page.route('http://127.0.0.1:8765/**', async (route) => {
|
||||
const request = route.request();
|
||||
if (offline) return route.fulfill({ status: 503, json: { error: '模拟服务断连' } });
|
||||
if (request.url().endsWith('/health'))
|
||||
return route.fulfill({
|
||||
json: {
|
||||
ready: true,
|
||||
trainerRoot: '/mock/trainer',
|
||||
tasks: ['Unitree-Go2-Flat'],
|
||||
taskMetadata: [
|
||||
{
|
||||
id: 'Unitree-Go2-Flat',
|
||||
name: 'Go2 平地',
|
||||
terrainPresets: ['plane'],
|
||||
terrainParameters: { size: { min: 4, max: 30, default: 8 } },
|
||||
sensorParameters: {},
|
||||
browserCompatible: true,
|
||||
},
|
||||
],
|
||||
},
|
||||
});
|
||||
if (request.url().endsWith('/presets')) return route.fulfill({ json: { presets: [] } });
|
||||
if (request.method() === 'POST') starts++;
|
||||
return route.fulfill({ json: job });
|
||||
});
|
||||
await load(page);
|
||||
await right(page);
|
||||
await page.getByRole('tab', { name: '控制台' }).click();
|
||||
const tools = page.getByRole('tabpanel', { name: '控制台' });
|
||||
const group = tools.getByRole('button', { name: /强化学习任务/ });
|
||||
await group.click();
|
||||
await page.getByLabel('训练服务访问令牌').fill('mock-token');
|
||||
await tools.getByRole('button', { name: '连接', exact: true }).click();
|
||||
await expect(tools.getByRole('alert')).toContainText('模拟服务断连');
|
||||
await screenshot(page, 'training-disconnected');
|
||||
offline = false;
|
||||
await tools.getByRole('button', { name: '连接', exact: true }).click();
|
||||
await page.getByLabel('训练地形', { exact: true }).selectOption('plane');
|
||||
await tools.getByRole('button', { name: '地形详细参数' }).click();
|
||||
await page.getByLabel('地图尺寸 m', { exact: true }).fill('100');
|
||||
await tools.getByRole('button', { name: '地形详细参数' }).click();
|
||||
await tools.getByRole('button', { name: '发起本地训练' }).click();
|
||||
await expect(tools.getByRole('alert')).toContainText('超出允许范围');
|
||||
await expect(tools.getByRole('alert')).toBeInViewport();
|
||||
await expect(page.getByLabel('地图尺寸 m', { exact: true })).toBeVisible();
|
||||
expect(starts).toBe(0);
|
||||
await screenshot(page, 'training-validation');
|
||||
await page.getByLabel('地图尺寸 m', { exact: true }).fill('12');
|
||||
await page.getByLabel('运行名称', { exact: true }).fill('preserved');
|
||||
await group.click();
|
||||
await page.getByRole('tab', { name: '检查器' }).click();
|
||||
await page.getByRole('tab', { name: '数据录制' }).click();
|
||||
await page.getByRole('tab', { name: '控制台' }).click();
|
||||
await group.click();
|
||||
await expect(page.getByLabel('运行名称', { exact: true })).toHaveValue('preserved');
|
||||
await expect(page.getByLabel('训练服务访问令牌')).toHaveValue('mock-token');
|
||||
await tools.getByRole('button', { name: '发起本地训练' }).click();
|
||||
await expect(tools.getByRole('button', { name: '停止训练' })).toBeVisible();
|
||||
await screenshot(page, 'training-running');
|
||||
await group.click();
|
||||
await expect(group).toContainText('训练中');
|
||||
await page.getByRole('tab', { name: '检查器' }).click();
|
||||
await page.getByRole('tab', { name: '控制台' }).click();
|
||||
await expect(group).toContainText('训练中');
|
||||
expect(starts).toBe(1);
|
||||
});
|
||||
}
|
||||
@@ -0,0 +1,106 @@
|
||||
import { expect, test, type Locator } from '@playwright/test';
|
||||
|
||||
async function expectWithinViewport(locator: Locator, width: number, height: number) {
|
||||
await expect(locator).toBeVisible();
|
||||
await expect
|
||||
.poll(async () => {
|
||||
const box = await locator.boundingBox();
|
||||
return Boolean(
|
||||
box &&
|
||||
box.x >= 7 &&
|
||||
box.y >= 7 &&
|
||||
box.x + box.width <= width - 7 &&
|
||||
box.y + box.height <= height - 7,
|
||||
);
|
||||
})
|
||||
.toBe(true);
|
||||
}
|
||||
|
||||
for (const theme of ['dark', 'light']) {
|
||||
for (const [width, height] of [
|
||||
[1920, 1080],
|
||||
[1440, 900],
|
||||
[1366, 768],
|
||||
[1024, 768],
|
||||
[768, 800],
|
||||
]) {
|
||||
test(`共享浮层边界与主题 ${theme} ${width}×${height}`, async ({ page }) => {
|
||||
await page.setViewportSize({ width, height });
|
||||
await page.addInitScript(
|
||||
(value) => localStorage.setItem('mujoco-platform-theme', value),
|
||||
theme,
|
||||
);
|
||||
await page.goto('/');
|
||||
const trigger = page.getByRole('button', { name: /隐藏工程面板|显示工程面板/ });
|
||||
await trigger.focus();
|
||||
const tooltip = page.getByRole('tooltip');
|
||||
await expectWithinViewport(tooltip, width, height);
|
||||
await expect(trigger).toHaveAttribute(
|
||||
'aria-describedby',
|
||||
(await tooltip.getAttribute('id')) as string,
|
||||
);
|
||||
expect(
|
||||
await tooltip.evaluate((node) => node.closest('.theme-dark, .theme-light')?.className),
|
||||
).toContain(`theme-${theme}`);
|
||||
await page.keyboard.press('Escape');
|
||||
await expect(tooltip).toHaveCount(0);
|
||||
await page.getByRole('button', { name: /FPS \d/ }).click();
|
||||
const popover = page.getByRole('dialog', { name: '性能详情' });
|
||||
await expectWithinViewport(popover, width, height);
|
||||
await page.screenshot({
|
||||
path: test.info().outputPath(`performance-${theme}-${width}.png`),
|
||||
fullPage: true,
|
||||
});
|
||||
await page.setViewportSize({ width: width - 40, height: height - 40 });
|
||||
await expectWithinViewport(popover, width - 40, height - 40);
|
||||
await page.keyboard.press('Escape');
|
||||
await expect(popover).toHaveCount(0);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
test('边缘 hover 提示可移入,disabled 控件说明可经键盘访问', async ({ page }) => {
|
||||
await page.setViewportSize({ width: 1440, height: 900 });
|
||||
await page.goto('/');
|
||||
const theme = page.getByRole('button', { name: '打开命令面板' });
|
||||
await theme.hover();
|
||||
const tooltip = page.getByRole('tooltip');
|
||||
await expectWithinViewport(tooltip, 1440, 900);
|
||||
await tooltip.hover();
|
||||
await expect(tooltip).toBeVisible();
|
||||
await page.keyboard.press('Escape');
|
||||
await expect(tooltip).toHaveCount(0);
|
||||
const step = page.getByRole('button', { name: '单步', exact: true });
|
||||
await expect(step).toBeDisabled();
|
||||
await step.locator('..').focus();
|
||||
await expect(tooltip).toHaveText('单步(暂停时可用)');
|
||||
expect(await theme.evaluate((node) => node.parentElement?.hasAttribute('tabindex'))).toBe(false);
|
||||
});
|
||||
|
||||
test('全屏中的对话框 Tooltip 不裁剪且 Escape 仅关闭顶层', async ({ page }) => {
|
||||
await page.setViewportSize({ width: 1440, height: 900 });
|
||||
await page.goto('/');
|
||||
await page.getByRole('button', { name: '更多工作台操作' }).click();
|
||||
await page.getByRole('menuitem', { name: '进入全屏' }).click();
|
||||
await page.getByRole('button', { name: '更多工作台操作' }).click();
|
||||
await page.getByRole('menuitem', { name: '工作台设置' }).click();
|
||||
const dialog = page.getByRole('dialog', { name: '工作台设置' });
|
||||
await expect(dialog).toBeVisible();
|
||||
const close = dialog.getByRole('button', { name: '关闭', exact: true });
|
||||
await close.focus();
|
||||
const tooltip = page.getByRole('tooltip');
|
||||
await expectWithinViewport(tooltip, 1440, 900);
|
||||
expect(await tooltip.evaluate((node) => document.fullscreenElement?.contains(node))).toBe(true);
|
||||
expect(await tooltip.evaluate((node) => Boolean(node.closest('[role="dialog"]')))).toBe(false);
|
||||
await page.screenshot({
|
||||
path: test.info().outputPath('fullscreen-dialog-tooltip.png'),
|
||||
fullPage: true,
|
||||
});
|
||||
await page.keyboard.press('Escape');
|
||||
await expect(tooltip).toHaveCount(0);
|
||||
await expect(dialog).toBeVisible();
|
||||
await page.keyboard.press('Escape');
|
||||
await expect(dialog).toHaveCount(0);
|
||||
await page.getByRole('button', { name: '更多工作台操作' }).click();
|
||||
await page.getByRole('menuitem', { name: '退出全屏' }).click();
|
||||
});
|
||||
@@ -0,0 +1,354 @@
|
||||
import { expect, test, type Page } from '@playwright/test';
|
||||
import { zipSync } from 'fflate';
|
||||
|
||||
const MODEL = `<mujoco model="layout-camera"><worldbody>
|
||||
<light pos="0 0 3"/><geom type="plane" size="3 3 .1"/>
|
||||
<camera pos="3 -3 2" xyaxes="1 1 0 -1 1 3"/>
|
||||
<body name="box" pos="0 0 1"><joint name="slide" type="slide" axis="1 0 0" range="-1 1"/>
|
||||
<geom type="box" size=".2 .2 .2" mass="1"/></body>
|
||||
</worldbody></mujoco>`;
|
||||
async function importModel(page: Page) {
|
||||
await page
|
||||
.locator('#mujoco-project-files')
|
||||
.setInputFiles({ name: 'layout.xml', mimeType: 'text/xml', buffer: Buffer.from(MODEL) });
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled({ timeout: 30_000 });
|
||||
}
|
||||
async function openLeft(page: Page) {
|
||||
const button = page.getByRole('button', { name: '显示工程面板', exact: true });
|
||||
if (await button.isVisible()) await button.click();
|
||||
}
|
||||
async function expectSlots(page: Page) {
|
||||
await expect
|
||||
.poll(() =>
|
||||
page.locator('main').evaluate((main) => {
|
||||
const bounds = main.getBoundingClientRect();
|
||||
const slots = [...main.querySelectorAll<HTMLElement>('[data-overlay-slot]')]
|
||||
.map((node) => ({ name: node.dataset.overlaySlot, rect: node.getBoundingClientRect() }))
|
||||
.filter(({ rect }) => rect.width > 0 && rect.height > 0);
|
||||
const problems: string[] = [];
|
||||
for (const { name, rect } of slots) {
|
||||
if (
|
||||
rect.left < bounds.left - 1 ||
|
||||
rect.right > bounds.right + 1 ||
|
||||
rect.top < bounds.top - 1 ||
|
||||
rect.bottom > bounds.bottom + 1
|
||||
)
|
||||
problems.push(`越界 ${name}`);
|
||||
}
|
||||
for (let a = 0; a < slots.length; a++)
|
||||
for (let b = a + 1; b < slots.length; b++) {
|
||||
const x = slots[a],
|
||||
y = slots[b];
|
||||
if (
|
||||
Math.min(x.rect.right, y.rect.right) - Math.max(x.rect.left, y.rect.left) > 1 &&
|
||||
Math.min(x.rect.bottom, y.rect.bottom) - Math.max(x.rect.top, y.rect.top) > 1
|
||||
)
|
||||
problems.push(`重叠 ${x.name}/${y.name}`);
|
||||
}
|
||||
if (document.documentElement.scrollWidth > window.innerWidth) problems.push('页面横向溢出');
|
||||
return problems;
|
||||
}),
|
||||
)
|
||||
.toEqual([]);
|
||||
}
|
||||
|
||||
for (const theme of ['dark', 'light'])
|
||||
for (const [width, height] of [
|
||||
[1920, 1080],
|
||||
[1440, 900],
|
||||
[1366, 768],
|
||||
[1024, 768],
|
||||
[768, 800],
|
||||
]) {
|
||||
test(`工作台槽位 ${theme} ${width}×${height} 摄像头/草稿/通知`, async ({ page }) => {
|
||||
const errors: string[] = [];
|
||||
page.on('pageerror', (error) => errors.push(error.message));
|
||||
await page.setViewportSize({ width, height });
|
||||
await page.addInitScript(
|
||||
(value) => localStorage.setItem('mujoco-platform-theme', value),
|
||||
theme,
|
||||
);
|
||||
await page.goto('/');
|
||||
await expect(
|
||||
page.getByRole('button', {
|
||||
name: width >= 1440 ? '隐藏工程面板' : '显示工程面板',
|
||||
exact: true,
|
||||
}),
|
||||
).toBeVisible();
|
||||
await expect(
|
||||
page.getByRole('button', {
|
||||
name: width >= 1024 ? '隐藏右侧面板' : '显示右侧面板',
|
||||
exact: true,
|
||||
}),
|
||||
).toBeVisible();
|
||||
await expect(page.locator('header')).not.toContainText('播放');
|
||||
await expect(page.getByLabel('地图草稿状态')).toHaveCount(0);
|
||||
await expectSlots(page);
|
||||
await page.screenshot({ path: test.info().outputPath(`empty-${theme}-${width}.png`) });
|
||||
await importModel(page);
|
||||
await openLeft(page);
|
||||
await page.getByLabel('地图资产库').getByRole('button', { name: '添加基础方盒' }).click();
|
||||
await expect(page.getByLabel('地图草稿状态')).toContainText('1 项未保存改动');
|
||||
await expect(page.getByLabel('摄像头画面')).toBeVisible();
|
||||
if (width < 1024)
|
||||
await page.getByRole('button', { name: '隐藏右侧面板', exact: true }).click();
|
||||
await expect(page.locator('[data-overlay-slot="notices"]')).toContainText('模型加载完成');
|
||||
await expectSlots(page);
|
||||
if (width >= 1024)
|
||||
expect((await page.getByRole('main').boundingBox())!.width).toBeGreaterThanOrEqual(480);
|
||||
await page.screenshot({
|
||||
path: test.info().outputPath(`draft-camera-${theme}-${width}.png`),
|
||||
mask: [page.getByLabel('视口状态')],
|
||||
});
|
||||
await page.getByRole('button', { name: '隐藏画面' }).click();
|
||||
await expect(page.getByLabel('摄像头画面')).toHaveCount(0);
|
||||
await page.getByRole('button', { name: '显示摄像头画面' }).click();
|
||||
await expectSlots(page);
|
||||
// 摄像头与草稿占位时,浮层仍须脱离侧栏裁剪并保持主题和键盘关闭。
|
||||
const performance = page.getByRole('button', { name: /FPS \d/ });
|
||||
await performance.click();
|
||||
const popover = page.getByRole('dialog', { name: '性能详情' });
|
||||
await expect(popover).toBeVisible();
|
||||
expect(
|
||||
await popover.evaluate((node) => {
|
||||
const rect = node.getBoundingClientRect();
|
||||
return (
|
||||
rect.left >= 7 &&
|
||||
rect.top >= 7 &&
|
||||
rect.right <= innerWidth - 7 &&
|
||||
rect.bottom <= innerHeight - 7
|
||||
);
|
||||
}),
|
||||
).toBe(true);
|
||||
expect(
|
||||
await popover.evaluate((node) => node.closest('.theme-dark, .theme-light')?.className),
|
||||
).toContain(`theme-${theme}`);
|
||||
await page.screenshot({
|
||||
path: test.info().outputPath(`draft-camera-popover-${theme}-${width}.png`),
|
||||
mask: [page.getByLabel('视口状态')],
|
||||
});
|
||||
await page.keyboard.press('Escape');
|
||||
await expect(popover).toHaveCount(0);
|
||||
await expect(performance).toBeFocused();
|
||||
if (width === 1024 || width === 768) {
|
||||
await page.getByRole('button', { name: '更多工作台操作' }).click();
|
||||
await page.getByRole('menuitem', { name: '进入全屏' }).click();
|
||||
await expect
|
||||
.poll(() => page.evaluate(() => Boolean(document.fullscreenElement)))
|
||||
.toBe(true);
|
||||
await expectSlots(page);
|
||||
await page.screenshot({
|
||||
path: test.info().outputPath('fullscreen-draft-camera.png'),
|
||||
mask: [page.getByLabel('视口状态')],
|
||||
});
|
||||
await page.getByRole('button', { name: '更多工作台操作' }).click();
|
||||
await page.getByRole('menuitem', { name: '工作台设置' }).click();
|
||||
const settings = page.getByRole('dialog', { name: '工作台设置' });
|
||||
await settings.getByRole('button', { name: '关闭', exact: true }).focus();
|
||||
const tooltip = page.getByRole('tooltip');
|
||||
await expect(tooltip).toBeVisible();
|
||||
expect(
|
||||
await tooltip.evaluate((node) => {
|
||||
const rect = node.getBoundingClientRect();
|
||||
return (
|
||||
Boolean(document.fullscreenElement?.contains(node)) &&
|
||||
rect.left >= 7 &&
|
||||
rect.top >= 7 &&
|
||||
rect.right <= innerWidth - 7 &&
|
||||
rect.bottom <= innerHeight - 7
|
||||
);
|
||||
}),
|
||||
).toBe(true);
|
||||
await page.screenshot({
|
||||
path: test.info().outputPath(`fullscreen-dialog-camera-${theme}-${width}.png`),
|
||||
mask: [page.getByLabel('视口状态')],
|
||||
});
|
||||
await page.keyboard.press('Escape');
|
||||
await expect(tooltip).toHaveCount(0);
|
||||
await expect(settings).toBeVisible();
|
||||
await page.keyboard.press('Escape');
|
||||
await expect(settings).toHaveCount(0);
|
||||
await expect(page.getByLabel('地图草稿状态')).toContainText('未保存改动');
|
||||
await expect(page.getByLabel('摄像头画面')).toBeVisible();
|
||||
await expectSlots(page);
|
||||
await page.getByRole('button', { name: '更多工作台操作' }).click();
|
||||
await page.getByRole('menuitem', { name: '退出全屏' }).click();
|
||||
}
|
||||
expect(errors).toEqual([]);
|
||||
});
|
||||
}
|
||||
|
||||
test('导入→大纲选择→键盘编辑关节→运行,侧栏切换保留输入与草稿', async ({ page }) => {
|
||||
await page.setViewportSize({ width: 1440, height: 900 });
|
||||
await page.goto('/');
|
||||
await importModel(page);
|
||||
await page.getByRole('treeitem', { name: /slide/ }).click();
|
||||
const joint = page.getByRole('slider', { name: /slide/ });
|
||||
await joint.focus();
|
||||
await page.keyboard.press('ArrowRight');
|
||||
await expect(joint).not.toHaveValue('0');
|
||||
const value = await joint.inputValue();
|
||||
await page.getByRole('button', { name: '隐藏右侧面板', exact: true }).click();
|
||||
await page.getByRole('button', { name: '显示右侧面板', exact: true }).click();
|
||||
await expect(joint).toHaveValue(value);
|
||||
await page.getByRole('button', { name: '隐藏右侧面板', exact: true }).blur();
|
||||
await page.keyboard.press('Space');
|
||||
await expect(page.getByRole('button', { name: '⏸ 暂停' })).toBeEnabled();
|
||||
await expect(page.getByLabel('视口状态')).not.toContainText('时间 0.000 s');
|
||||
await page.keyboard.press('Space');
|
||||
await expect(page.getByLabel('视口状态')).toContainText('已暂停');
|
||||
await page.getByLabel('地图资产库').getByRole('button', { name: '添加基础方盒' }).click();
|
||||
await page.getByLabel('对象位置X').fill('2');
|
||||
await page.getByLabel('对象位置X').blur();
|
||||
await page.getByRole('button', { name: '隐藏右侧面板', exact: true }).click();
|
||||
await page.getByRole('button', { name: '显示右侧面板', exact: true }).click();
|
||||
await expect(page.getByLabel('对象位置X')).toHaveValue('2');
|
||||
await expect(page.getByLabel('地图草稿状态')).toContainText('未保存改动');
|
||||
await page.keyboard.press('Control+s');
|
||||
await expect(page.getByLabel('地图草稿状态')).toBeHidden({ timeout: 30_000 });
|
||||
await page
|
||||
.getByLabel('地图对象列表')
|
||||
.getByRole('button', { name: /基础方盒/ })
|
||||
.click();
|
||||
await page.getByLabel('对象位置X').fill('3');
|
||||
await page.getByLabel('地图草稿状态').getByRole('button', { name: '丢弃地图草稿' }).click();
|
||||
await page
|
||||
.getByLabel('地图对象列表')
|
||||
.getByRole('button', { name: /基础方盒/ })
|
||||
.click();
|
||||
await expect(page.getByLabel('对象位置X')).toHaveValue('2');
|
||||
await expect(page.getByLabel('地图草稿状态')).toBeHidden();
|
||||
await page.screenshot({
|
||||
path: test.info().outputPath('closed-loop-applied.png'),
|
||||
mask: [page.getByLabel('视口状态')],
|
||||
});
|
||||
});
|
||||
|
||||
test('宽度预算、键盘调整与窄屏临时折叠不覆盖桌面偏好', async ({ page }) => {
|
||||
await page.setViewportSize({ width: 1440, height: 900 });
|
||||
await page.addInitScript(() => {
|
||||
localStorage.setItem('mujoco-platform-layout', '{"left":true,"right":true}');
|
||||
localStorage.setItem('mujoco-left-sidebar-width', '500');
|
||||
localStorage.setItem('mujoco-right-sidebar-width', '500');
|
||||
});
|
||||
await page.goto('/');
|
||||
expect((await page.getByRole('main').boundingBox())!.width).toBeGreaterThanOrEqual(480);
|
||||
await page.setViewportSize({ width: 1024, height: 768 });
|
||||
await expect
|
||||
.poll(async () => (await page.getByRole('main').boundingBox())!.width)
|
||||
.toBeGreaterThanOrEqual(480);
|
||||
await page.setViewportSize({ width: 768, height: 800 });
|
||||
await expect(page.getByRole('button', { name: '显示工程面板', exact: true })).toBeVisible();
|
||||
await page.getByRole('button', { name: '显示工程面板', exact: true }).click();
|
||||
await page.getByRole('button', { name: '显示右侧面板', exact: true }).click();
|
||||
await expect(page.getByRole('button', { name: '显示工程面板', exact: true })).toBeVisible();
|
||||
expect(await page.evaluate(() => localStorage.getItem('mujoco-left-sidebar-width'))).toBe('500');
|
||||
expect(await page.evaluate(() => localStorage.getItem('mujoco-platform-layout'))).toBe(
|
||||
'{"left":true,"right":true}',
|
||||
);
|
||||
await page.setViewportSize({ width: 1920, height: 1080 });
|
||||
const left = page.getByRole('separator', { name: '调整工程面板宽度' });
|
||||
await expect.poll(async () => (await left.locator('..').boundingBox())!.width).toBe(500);
|
||||
await left.focus();
|
||||
await page.keyboard.press('ArrowLeft');
|
||||
expect(await page.evaluate(() => localStorage.getItem('mujoco-left-sidebar-width'))).toBe('484');
|
||||
});
|
||||
|
||||
test('编译失败时摄像头、风险摘要与重试/放弃入口不重叠', async ({ page }) => {
|
||||
await page.setViewportSize({ width: 1024, height: 768 });
|
||||
await page.goto('/');
|
||||
const project = zipSync({
|
||||
'model.xml': Buffer.from(MODEL),
|
||||
'maps/invalid/map.json': Buffer.from(
|
||||
JSON.stringify({
|
||||
schemaVersion: 1,
|
||||
id: 'invalid',
|
||||
name: '错误地图',
|
||||
coordinateSystem: { units: 'm', up: 'Z', forward: '+X' },
|
||||
physics: { source: 'world.xml' },
|
||||
spawnPoints: [],
|
||||
}),
|
||||
),
|
||||
'maps/invalid/world.xml': Buffer.from(
|
||||
'<mujoco><worldbody><body><joint/><geom type="box" size="1 1 1"/></body></worldbody></mujoco>',
|
||||
),
|
||||
});
|
||||
await page.locator('#mujoco-project-files').setInputFiles({
|
||||
name: 'invalid.zip',
|
||||
mimeType: 'application/zip',
|
||||
buffer: Buffer.from(project),
|
||||
});
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled();
|
||||
await openLeft(page);
|
||||
await page
|
||||
.getByLabel('地图资产库')
|
||||
.getByRole('button', { name: '放置工程地图 错误地图' })
|
||||
.click();
|
||||
await page.getByLabel('地图草稿状态').getByRole('button', { name: '提交地图草稿' }).click();
|
||||
await expect(page.getByRole('alert')).toContainText('模型编译失败');
|
||||
await expect(page.getByLabel('地图草稿状态')).toContainText('应用失败');
|
||||
await expect(page.getByLabel('摄像头画面')).toBeVisible();
|
||||
await expectSlots(page);
|
||||
await page.getByRole('button', { name: '技术详情', exact: true }).click();
|
||||
await expectSlots(page);
|
||||
await page.screenshot({ path: test.info().outputPath('camera-draft-error-1024.png') });
|
||||
await page.getByRole('button', { name: '关闭错误' }).click();
|
||||
await expect(page.getByLabel('地图草稿状态')).toContainText('应用失败');
|
||||
await page.getByLabel('地图草稿状态').getByRole('button', { name: '丢弃地图草稿' }).click();
|
||||
await expect(page.getByLabel('地图草稿状态')).toBeHidden();
|
||||
});
|
||||
|
||||
for (const theme of ['dark', 'light']) {
|
||||
test(`主动作实际对比度与减少动效 ${theme}`, async ({ page }) => {
|
||||
await page.setViewportSize({ width: 768, height: 800 });
|
||||
await page.emulateMedia({ reducedMotion: 'reduce' });
|
||||
await page.addInitScript(
|
||||
(value) => localStorage.setItem('mujoco-platform-theme', value),
|
||||
theme,
|
||||
);
|
||||
await page.goto('/');
|
||||
const primary = page.getByRole('button', { name: '选择文件', exact: true });
|
||||
await expect(primary).toBeVisible();
|
||||
const contrast = () =>
|
||||
primary.evaluate((node) => {
|
||||
const style = getComputedStyle(node);
|
||||
const luminance = (color: string) => {
|
||||
const values = color
|
||||
.match(/[\d.]+/g)!
|
||||
.slice(0, 3)
|
||||
.map(Number)
|
||||
.map((value) => {
|
||||
const channel = value / 255;
|
||||
return channel <= 0.04045 ? channel / 12.92 : ((channel + 0.055) / 1.055) ** 2.4;
|
||||
});
|
||||
return values[0] * 0.2126 + values[1] * 0.7152 + values[2] * 0.0722;
|
||||
};
|
||||
const fg = luminance(style.color),
|
||||
bg = luminance(style.backgroundColor);
|
||||
return (Math.max(fg, bg) + 0.05) / (Math.min(fg, bg) + 0.05);
|
||||
});
|
||||
await expect.poll(contrast).toBeGreaterThanOrEqual(4.5);
|
||||
await primary.hover();
|
||||
await expect.poll(contrast).toBeGreaterThanOrEqual(4.5);
|
||||
await primary.focus();
|
||||
expect(
|
||||
await primary.evaluate((node) =>
|
||||
Number.parseFloat(getComputedStyle(node).transitionDuration),
|
||||
),
|
||||
).toBeLessThanOrEqual(0.00001);
|
||||
await importModel(page);
|
||||
await openLeft(page);
|
||||
await page.getByLabel('地图资产库').getByRole('button', { name: '添加基础方盒' }).click();
|
||||
await page.getByRole('button', { name: '隐藏右侧面板', exact: true }).click();
|
||||
await expectSlots(page);
|
||||
expect(
|
||||
await page
|
||||
.locator('.draft-dirty-dot')
|
||||
.evaluate((node) => Number.parseFloat(getComputedStyle(node).animationDuration)),
|
||||
).toBeLessThanOrEqual(0.00001);
|
||||
await page.screenshot({
|
||||
path: test.info().outputPath(`reduced-motion-${theme}.png`),
|
||||
mask: [page.getByLabel('视口状态')],
|
||||
});
|
||||
});
|
||||
}
|
||||
@@ -0,0 +1,299 @@
|
||||
import { expect, test, type Page } from '@playwright/test';
|
||||
import { tuningFixture } from './tuningFixture';
|
||||
|
||||
const sizes = [
|
||||
[1920, 1080],
|
||||
[1440, 900],
|
||||
[1366, 768],
|
||||
[1024, 768],
|
||||
[768, 800],
|
||||
] as const;
|
||||
async function noOverflow(page: Page) {
|
||||
expect(await page.evaluate(() => document.documentElement.scrollWidth <= innerWidth)).toBe(true);
|
||||
const overflowing = await page
|
||||
.locator('#root')
|
||||
.evaluate((root) =>
|
||||
[...root.querySelectorAll<HTMLElement>('main, nav, header, [role="dialog"]')]
|
||||
.filter(
|
||||
(el) => el.getClientRects().length && el.getBoundingClientRect().right > innerWidth + 1,
|
||||
)
|
||||
.map((el) => el.tagName),
|
||||
);
|
||||
expect(overflowing).toEqual([]);
|
||||
}
|
||||
for (const theme of ['dark', 'light'] as const)
|
||||
for (const [width, height] of sizes) {
|
||||
test(`调参与系统 ${theme} ${width}`, async ({ page }, info) => {
|
||||
await page.setViewportSize({ width, height });
|
||||
await page.addInitScript(
|
||||
(value) => localStorage.setItem('mujoco-platform-theme', value),
|
||||
theme,
|
||||
);
|
||||
const errors: string[] = [];
|
||||
page.on('pageerror', (error) => errors.push(error.message));
|
||||
const session = tuningFixture();
|
||||
const actions: { method: string; path: string; body: unknown }[] = [];
|
||||
await page.route('**/api/tuning/**', async (route) => {
|
||||
const req = route.request(),
|
||||
path = new URL(req.url()).pathname;
|
||||
expect(req.headers().authorization).toBe('Bearer ui-mock-token');
|
||||
if (req.method() !== 'GET')
|
||||
actions.push({ method: req.method(), path, body: req.postDataJSON() });
|
||||
let body: unknown = session;
|
||||
if (path.endsWith('/capabilities'))
|
||||
body = { configured: true, ready: true, model: 'UI Mock Agent', pretrainedSources: [] };
|
||||
else if (path === '/api/tuning/sessions') body = { sessions: [session] };
|
||||
else if (path.endsWith('/metrics'))
|
||||
body = {
|
||||
series: [
|
||||
{
|
||||
tag: 'Train/mean_reward',
|
||||
points: Array.from({ length: 30 }, (_, step) => ({
|
||||
step,
|
||||
wallTime: step,
|
||||
value: step / 10 + Math.sin(step) / 4,
|
||||
})),
|
||||
},
|
||||
],
|
||||
};
|
||||
else if (path.endsWith('/approve')) {
|
||||
session.proposals[0].state = 'approved';
|
||||
session.state = 'paused';
|
||||
} else if (req.method() === 'DELETE') session.state = 'cancelled';
|
||||
await route.fulfill({ json: body });
|
||||
});
|
||||
await page.goto('/tuning.html');
|
||||
await expect(page.getByRole('heading', { name: 'Go2 自调参' })).toBeVisible();
|
||||
await expect(page.locator('#root > div')).toHaveClass(new RegExp(`theme-${theme}`));
|
||||
await noOverflow(page);
|
||||
await page.screenshot({ path: info.outputPath('tuning-empty.png') });
|
||||
await page.getByLabel('训练服务地址').fill('http://127.0.0.1:8765');
|
||||
await page.getByLabel('访问令牌(仅当前标签页)').fill('ui-mock-token');
|
||||
await page.getByRole('button', { name: '连接/刷新' }).click();
|
||||
await expect(page.getByRole('img', { name: /收敛曲线/ })).toBeVisible();
|
||||
await expect(page.getByRole('button', { name: '停止 Session' })).toBeVisible();
|
||||
const sessions = page.getByRole('region', { name: '会话与 Trial' });
|
||||
const metrics = page.getByRole('main', { name: '指标与排行' });
|
||||
const decisions = page.getByRole('complementary', { name: '决策与审批' });
|
||||
await expect(sessions).toBeVisible();
|
||||
await expect(metrics).toBeVisible();
|
||||
await expect(decisions).toBeVisible();
|
||||
const left = (await sessions.boundingBox())!;
|
||||
const center = (await metrics.boundingBox())!;
|
||||
const right = (await decisions.boundingBox())!;
|
||||
if (width >= 1280) {
|
||||
expect(Math.abs(left.y - center.y)).toBeLessThan(2);
|
||||
expect(Math.abs(center.y - right.y)).toBeLessThan(2);
|
||||
expect(left.x + left.width).toBeLessThanOrEqual(center.x + 1);
|
||||
expect(center.x + center.width).toBeLessThanOrEqual(right.x + 1);
|
||||
} else {
|
||||
expect(left.y + left.height).toBeLessThanOrEqual(center.y + 1);
|
||||
expect(center.y + center.height).toBeLessThanOrEqual(right.y + 1);
|
||||
}
|
||||
await page.getByRole('button', { name: '收敛曲线放大', exact: true }).click();
|
||||
await page
|
||||
.getByRole('button', { name: `切换到${theme === 'dark' ? '白天' : '黑夜'}主题` })
|
||||
.click();
|
||||
await expect(page.locator('#root > div')).toHaveClass(
|
||||
new RegExp(`theme-${theme === 'dark' ? 'light' : 'dark'}`),
|
||||
);
|
||||
await page
|
||||
.getByRole('button', { name: `切换到${theme === 'dark' ? '黑夜' : '白天'}主题` })
|
||||
.click();
|
||||
await noOverflow(page);
|
||||
await page.screenshot({ path: info.outputPath('tuning-metrics.png') });
|
||||
await page.getByRole('heading', { name: 'Trial Leaderboard' }).scrollIntoViewIfNeeded();
|
||||
await noOverflow(page);
|
||||
await page.screenshot({ path: info.outputPath('tuning-leaderboard.png') });
|
||||
await page.getByRole('button', { name: '导入', exact: true }).scrollIntoViewIfNeeded();
|
||||
await expect(page.getByRole('button', { name: '导入', exact: true })).toBeVisible();
|
||||
await expect(page.getByText('导入会替换主工作台策略;无主窗口时下载文件。')).toBeVisible();
|
||||
await noOverflow(page);
|
||||
await page.screenshot({ path: info.outputPath('tuning-sessions.png') });
|
||||
await page.getByRole('button', { name: 'Monaco Diff 审查' }).scrollIntoViewIfNeeded();
|
||||
await noOverflow(page);
|
||||
await page.screenshot({ path: info.outputPath('tuning-decisions.png') });
|
||||
await page.getByRole('button', { name: 'Monaco Diff 审查' }).click();
|
||||
const review = page.getByRole('dialog', { name: /Reward Merge Patch 审查/ });
|
||||
await expect(review.locator('.monaco-diff-editor')).toBeVisible();
|
||||
await expect(
|
||||
review.locator(theme === 'light' ? '.monaco-editor.vs' : '.monaco-editor.vs-dark').first(),
|
||||
).toBeVisible();
|
||||
await review
|
||||
.locator('.monaco-editor')
|
||||
.last()
|
||||
.click({ position: { x: 180, y: 80 } });
|
||||
await page.keyboard.press('Control+Home');
|
||||
await page.keyboard.press('ArrowDown');
|
||||
await page.keyboard.press('ArrowDown');
|
||||
await page.keyboard.press('Home');
|
||||
await page.keyboard.press('Shift+End');
|
||||
await page.keyboard.insertText(' "pose": 1.2');
|
||||
await page.evaluate(
|
||||
(value) => {
|
||||
localStorage.setItem('mujoco-platform-theme', value);
|
||||
window.dispatchEvent(new StorageEvent('storage', { key: 'mujoco-platform-theme' }));
|
||||
},
|
||||
theme === 'light' ? 'dark' : 'light',
|
||||
);
|
||||
await expect(
|
||||
review.locator(theme === 'light' ? '.monaco-editor.vs-dark' : '.monaco-editor.vs').last(),
|
||||
).toBeVisible();
|
||||
await page.evaluate((value) => {
|
||||
localStorage.setItem('mujoco-platform-theme', value);
|
||||
window.dispatchEvent(new StorageEvent('storage', { key: 'mujoco-platform-theme' }));
|
||||
}, theme);
|
||||
await expect(review.locator('.monaco-editor').last()).toContainText('1.2');
|
||||
await page.getByLabel('Proposal 审批反馈').fill('界面验收,不启动真实训练');
|
||||
await noOverflow(page);
|
||||
await page.screenshot({ path: info.outputPath('tuning-review.png') });
|
||||
await page.getByRole('button', { name: '批准编辑后的 Patch' }).click();
|
||||
await expect(review).toBeHidden();
|
||||
expect(actions.find((item) => item.path.endsWith('/approve'))?.body).toEqual({
|
||||
feedback: '界面验收,不启动真实训练',
|
||||
patch: { weights: { pose: 1.2 }, params: {} },
|
||||
});
|
||||
await page.getByRole('button', { name: '参数护栏' }).click();
|
||||
await expect(page.getByRole('dialog', { name: /参数安全护栏/ })).toBeVisible();
|
||||
await noOverflow(page);
|
||||
await page.screenshot({ path: info.outputPath('tuning-constraints.png') });
|
||||
await page.keyboard.press('Escape');
|
||||
await page.getByRole('button', { name: '停止 Session' }).click();
|
||||
expect(actions.filter((item) => item.method === 'DELETE')).toHaveLength(1);
|
||||
expect(
|
||||
actions.filter((item) => item.path === '/api/tuning/sessions' && item.method === 'POST'),
|
||||
).toHaveLength(0);
|
||||
await page.goto('/');
|
||||
await expect(page.getByRole('region', { name: '导入模型工程' })).toBeVisible();
|
||||
await page.getByRole('button', { name: '更多工作台操作' }).click();
|
||||
await page.getByRole('menuitem', { name: '工作台设置' }).click();
|
||||
await expect(page.getByRole('dialog', { name: '工作台设置' })).toBeVisible();
|
||||
await noOverflow(page);
|
||||
await page.screenshot({ path: info.outputPath('settings.png') });
|
||||
await page.keyboard.press('Escape');
|
||||
await page.getByRole('button', { name: '布局设置' }).click();
|
||||
await page.screenshot({ path: info.outputPath('layout.png') });
|
||||
await page.keyboard.press('Escape');
|
||||
await page.keyboard.press('Control+k');
|
||||
await page.getByLabel('搜索命令').fill('快捷');
|
||||
await expect(page.getByLabel('搜索命令')).toBeFocused();
|
||||
await page.keyboard.press('Enter');
|
||||
await expect(page.getByRole('dialog', { name: '快捷键与视口操作' })).toBeVisible();
|
||||
await page.screenshot({ path: info.outputPath('help.png') });
|
||||
await page.keyboard.press('Escape');
|
||||
await page.keyboard.press('Control+k');
|
||||
await page.getByLabel('搜索命令').fill('无匹配命令');
|
||||
await expect(page.getByText('没有匹配的命令')).toBeVisible();
|
||||
await page.screenshot({ path: info.outputPath('command-empty.png') });
|
||||
expect(errors).toEqual([]);
|
||||
});
|
||||
}
|
||||
|
||||
for (const theme of ['dark', 'light'] as const)
|
||||
test(`源码与导入风险 ${theme}`, async ({ page }, info) => {
|
||||
await page.setViewportSize({ width: 768, height: 800 });
|
||||
await page.addInitScript(
|
||||
(value) => localStorage.setItem('mujoco-platform-theme', value),
|
||||
theme,
|
||||
);
|
||||
const errors: string[] = [];
|
||||
page.on('pageerror', (error) => errors.push(error.message));
|
||||
await page.goto('/');
|
||||
await page.locator('#mujoco-project-files').setInputFiles({
|
||||
name: 'model.urdf',
|
||||
mimeType: 'text/xml',
|
||||
buffer: Buffer.from(
|
||||
'<robot name="jointed"><link name="base"><inertial><mass value="1"/><origin xyz="0 0 0"/><inertia ixx=".1" iyy=".1" izz=".1" ixy="0" ixz="0" iyz="0"/></inertial><visual><geometry><box size=".4 .4 .2"/></geometry></visual></link><link name="arm"><inertial><mass value=".2"/><origin xyz="0 0 .25"/><inertia ixx=".01" iyy=".01" izz=".01" ixy="0" ixz="0" iyz="0"/></inertial><visual><origin xyz="0 0 .25"/><geometry><box size=".1 .1 .5"/></geometry></visual></link><joint name="shoulder" type="revolute"><parent link="base"/><child link="arm"/><origin xyz="0 0 .1"/><axis xyz="0 1 0"/><limit lower="-1" upper="1" effort="10" velocity="2"/></joint></robot>',
|
||||
),
|
||||
});
|
||||
const urdf = page.getByRole('dialog', { name: '配置 URDF 仿真组件' });
|
||||
await expect(urdf.getByText(/生成不限幅 motor 控制输入/)).toBeVisible();
|
||||
await page.screenshot({ path: info.outputPath('urdf-options.png') });
|
||||
await urdf.getByRole('button', { name: '转换并加载' }).click();
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled();
|
||||
await page.getByRole('button', { name: '通知中心' }).click();
|
||||
const notifications = page.getByRole('dialog', { name: '通知中心' });
|
||||
for (const summary of await notifications.getByText('事件详情', { exact: true }).all())
|
||||
await summary.click();
|
||||
await page.screenshot({ path: info.outputPath('notifications.png') });
|
||||
await page.getByRole('button', { name: '事件日志', exact: true }).click();
|
||||
await expect(page.getByRole('dialog', { name: '诊断与事件日志' })).toBeVisible();
|
||||
await page.screenshot({ path: info.outputPath('diagnostics.png') });
|
||||
await page.keyboard.press('Escape');
|
||||
await page.getByRole('button', { name: '源代码', exact: true }).click();
|
||||
const source = page.getByRole('dialog', { name: '转换后的 MJCF 编辑器' });
|
||||
await expect(source.locator('.monaco-editor')).toBeVisible();
|
||||
await noOverflow(page);
|
||||
await page.screenshot({ path: info.outputPath('source.png') });
|
||||
await source.locator('.monaco-editor').click({ position: { x: 200, y: 100 } });
|
||||
await page.keyboard.press('Control+End');
|
||||
await page.keyboard.insertText('\n');
|
||||
await expect(source.getByText('已修改', { exact: true })).toBeVisible();
|
||||
await page.keyboard.press('Escape');
|
||||
const confirm = page.getByRole('dialog', { name: '放弃未保存的修改?' });
|
||||
await expect(confirm.getByText(/无法恢复/)).toBeVisible();
|
||||
await page.screenshot({ path: info.outputPath('confirm.png') });
|
||||
await page.keyboard.press('Escape');
|
||||
await expect(confirm).toBeHidden();
|
||||
await expect(source).toBeVisible();
|
||||
await source.getByRole('button', { name: '关闭源代码编辑器' }).click();
|
||||
await confirm.getByRole('button', { name: '放弃修改' }).click();
|
||||
await expect(source).toBeHidden();
|
||||
expect(errors).toEqual([]);
|
||||
});
|
||||
|
||||
test('主工作台与独立调参标签页双向同步主题', async ({ page, context }) => {
|
||||
await page.goto('/');
|
||||
const tuning = await context.newPage();
|
||||
await tuning.goto('/tuning.html');
|
||||
await page.getByRole('button', { name: '更多工作台操作' }).click();
|
||||
await page.getByRole('menuitem', { name: '切换主题' }).click();
|
||||
await expect(tuning.locator('#root > div')).toHaveClass(/theme-light/);
|
||||
await tuning.getByRole('button', { name: '切换到黑夜主题' }).click();
|
||||
await expect(page.locator('#root > div')).toHaveClass(/theme-dark/);
|
||||
await tuning.close();
|
||||
});
|
||||
|
||||
for (const theme of ['dark', 'light'] as const)
|
||||
test(`连接错误与入口选择 ${theme}`, async ({ page }, info) => {
|
||||
await page.setViewportSize({ width: 768, height: 800 });
|
||||
await page.addInitScript(
|
||||
(value) => localStorage.setItem('mujoco-platform-theme', value),
|
||||
theme,
|
||||
);
|
||||
await page.route('**/api/tuning/**', (route) =>
|
||||
route.fulfill({ status: 503, json: { error: '训练服务未就绪,请检查本地服务' } }),
|
||||
);
|
||||
await page.goto('/tuning.html');
|
||||
await page.getByLabel('访问令牌(仅当前标签页)').fill('ui-mock-token');
|
||||
await page.getByRole('button', { name: '连接/刷新' }).click();
|
||||
await expect(page.getByRole('alert')).toBeInViewport();
|
||||
await expect(page.getByRole('button', { name: '关闭错误' })).toBeVisible();
|
||||
await page.screenshot({ path: info.outputPath('tuning-error.png') });
|
||||
await page.goto('/');
|
||||
await page.locator('#mujoco-project-files').setInputFiles(
|
||||
['first.xml', 'second.xml'].map((name) => ({
|
||||
name,
|
||||
mimeType: 'text/xml',
|
||||
buffer: Buffer.from(
|
||||
'<mujoco><worldbody><geom type="sphere" size=".1"/></worldbody></mujoco>',
|
||||
),
|
||||
})),
|
||||
);
|
||||
const entries = page.getByRole('dialog', { name: '选择模型入口' });
|
||||
await expect(entries).toBeVisible();
|
||||
await page.screenshot({ path: info.outputPath('entry-selection.png') });
|
||||
await entries.getByRole('button', { name: /first.xml/ }).click();
|
||||
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeEnabled();
|
||||
await page.locator('#mujoco-project-files').setInputFiles({
|
||||
name: 'bad.xml',
|
||||
mimeType: 'text/xml',
|
||||
buffer: Buffer.from(
|
||||
'<mujoco><worldbody><geom type="sphere" size="-1"/></worldbody></mujoco>',
|
||||
),
|
||||
});
|
||||
await expect(page.getByRole('alert')).toContainText('模型编译失败');
|
||||
await page.getByRole('button', { name: '技术详情', exact: true }).click();
|
||||
await noOverflow(page);
|
||||
await page.screenshot({ path: info.outputPath('compile-error.png') });
|
||||
});
|
||||
@@ -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拒绝及旧会话保留回归。
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+15
-33
@@ -10,6 +10,14 @@
|
||||
/>
|
||||
<link rel="icon" href="data:," />
|
||||
<title>MuJoCo Web 仿真平台</title>
|
||||
<script>
|
||||
try {
|
||||
document.documentElement.dataset.bootTheme =
|
||||
localStorage.getItem('mujoco-platform-theme') === 'light' ? 'light' : 'dark';
|
||||
} catch {
|
||||
/* 默认深色,启动不依赖存储权限。 */
|
||||
}
|
||||
</script>
|
||||
<style>
|
||||
html,
|
||||
body,
|
||||
@@ -19,52 +27,26 @@
|
||||
}
|
||||
body {
|
||||
background: #09111e;
|
||||
color: #f1f5f9;
|
||||
}
|
||||
html[data-boot-theme='light'] body {
|
||||
background: #f7f9fc;
|
||||
color: #172b41;
|
||||
}
|
||||
.boot-screen {
|
||||
display: grid;
|
||||
height: 100%;
|
||||
place-items: center;
|
||||
color: #f1f5f9;
|
||||
font:
|
||||
13px Inter,
|
||||
system-ui,
|
||||
13px system-ui,
|
||||
sans-serif;
|
||||
}
|
||||
.boot-mark {
|
||||
display: grid;
|
||||
width: 42px;
|
||||
height: 42px;
|
||||
margin: 0 auto 14px;
|
||||
place-items: center;
|
||||
border: 1px solid #2b604f;
|
||||
border-radius: 13px;
|
||||
background: #123b31;
|
||||
color: #38d39f;
|
||||
font-weight: 800;
|
||||
box-shadow: 0 16px 48px rgb(0 0 0 / 35%);
|
||||
animation: boot-pulse 1.4s ease-in-out infinite;
|
||||
}
|
||||
.boot-caption {
|
||||
color: #8fa0b5;
|
||||
font-size: 11px;
|
||||
letter-spacing: 0.04em;
|
||||
text-align: center;
|
||||
}
|
||||
@keyframes boot-pulse {
|
||||
50% {
|
||||
transform: translateY(-2px);
|
||||
box-shadow: 0 18px 54px rgb(56 211 159 / 15%);
|
||||
}
|
||||
}
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<div id="root">
|
||||
<div class="boot-screen" role="status" aria-label="正在启动仿真工作台">
|
||||
<div>
|
||||
<div class="boot-mark">M</div>
|
||||
<div class="boot-caption">正在启动本地仿真工作台…</div>
|
||||
</div>
|
||||
正在启动本地仿真工作台…
|
||||
</div>
|
||||
</div>
|
||||
<script type="module" src="/src/main.tsx"></script>
|
||||
|
||||
+550
-408
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,18 @@
|
||||
import { render, screen } from '@testing-library/react';
|
||||
import { ErrorBoundary } from './ErrorBoundary';
|
||||
it('致命错误保留主题、完整错误与重新加载风险', () => {
|
||||
vi.spyOn(console, 'error').mockImplementation(() => {});
|
||||
localStorage.setItem('mujoco-platform-theme', 'light');
|
||||
function Broken(): never {
|
||||
throw new Error('渲染失败');
|
||||
}
|
||||
render(
|
||||
<ErrorBoundary>
|
||||
<Broken />
|
||||
</ErrorBoundary>,
|
||||
);
|
||||
expect(screen.getByRole('main')).toHaveClass('theme-light');
|
||||
expect(screen.getByText('渲染失败')).toBeVisible();
|
||||
expect(screen.getByText(/重新加载会丢失/)).toBeVisible();
|
||||
expect(screen.getByRole('button', { name: '重新加载' })).toBeEnabled();
|
||||
});
|
||||
@@ -1,3 +1,4 @@
|
||||
import { useThemePreference } from './hooks/useThemePreference';
|
||||
import { Component, type ErrorInfo, type ReactNode } from 'react';
|
||||
import { Button } from '../components/ui';
|
||||
export class ErrorBoundary extends Component<{ children: ReactNode }, { error?: Error }> {
|
||||
@@ -9,20 +10,26 @@ export class ErrorBoundary extends Component<{ children: ReactNode }, { error?:
|
||||
console.error('React fatal error', error, info);
|
||||
}
|
||||
render() {
|
||||
return this.state.error ? (
|
||||
<main className="grid h-screen place-items-center bg-app text-text-primary">
|
||||
<section className="max-w-xl rounded-xl border border-danger-border bg-panel p-6 shadow-xl">
|
||||
<h1 className="text-xl font-semibold">界面发生致命错误</h1>
|
||||
<pre className="mt-3 whitespace-pre-wrap text-sm text-danger">
|
||||
{this.state.error.message}
|
||||
</pre>
|
||||
<Button variant="danger" className="mt-4" onClick={() => location.reload()}>
|
||||
重新加载
|
||||
</Button>
|
||||
</section>
|
||||
</main>
|
||||
) : (
|
||||
this.props.children
|
||||
);
|
||||
return this.state.error ? <FatalError error={this.state.error} /> : this.props.children;
|
||||
}
|
||||
}
|
||||
|
||||
function FatalError({ error }: { error: Error }) {
|
||||
const [theme] = useThemePreference();
|
||||
return (
|
||||
<main
|
||||
className={`theme-${theme} grid min-h-screen place-items-center bg-app p-4 text-text-primary`}
|
||||
>
|
||||
<section className="min-w-0 max-w-xl rounded-lg border border-danger-border bg-panel p-6 shadow-xl">
|
||||
<h1 className="text-xl font-semibold">界面发生致命错误</h1>
|
||||
<pre className="mt-3 max-h-[50vh] overflow-auto whitespace-pre-wrap break-words text-sm text-danger">
|
||||
{error.message}
|
||||
</pre>
|
||||
<p className="mt-3 text-xs text-warning">重新加载会丢失当前浏览器会话中未保存的修改。</p>
|
||||
<Button variant="danger" className="mt-4" onClick={() => location.reload()}>
|
||||
重新加载
|
||||
</Button>
|
||||
</section>
|
||||
</main>
|
||||
);
|
||||
}
|
||||
|
||||
@@ -35,6 +35,12 @@ export function CommandPalette({
|
||||
useEffect(() => {
|
||||
if (open) requestAnimationFrame(() => input.current?.focus());
|
||||
}, [open]);
|
||||
useEffect(() => {
|
||||
if (open && highlighted >= 0)
|
||||
document
|
||||
.getElementById(`${listId}-${filtered[highlighted].id}`)
|
||||
?.scrollIntoView?.({ block: 'nearest' });
|
||||
}, [open, highlighted, filtered, listId]);
|
||||
const close = () => {
|
||||
setQuery('');
|
||||
setActive(0);
|
||||
@@ -105,8 +111,8 @@ export function CommandPalette({
|
||||
>
|
||||
<span className="flex h-5 w-5 items-center justify-center">{command.icon}</span>
|
||||
<span className="min-w-0 flex-1">
|
||||
<span className="block truncate font-medium">{command.label}</span>
|
||||
<span className="block text-[10px] text-text-tertiary">{command.group}</span>
|
||||
<span className="block break-words font-medium">{command.label}</span>
|
||||
<span className="block text-xs text-text-tertiary">{command.group}</span>
|
||||
</span>
|
||||
{command.shortcut && <Kbd>{command.shortcut}</Kbd>}
|
||||
</button>
|
||||
|
||||
@@ -22,7 +22,7 @@ export function DiagnosticNotice({
|
||||
<div className="min-w-0 flex-1">
|
||||
<h2 className="text-sm font-semibold text-text-primary">{value.summary}</h2>
|
||||
{value.path && (
|
||||
<p className="mt-0.5 truncate text-xs text-text-tertiary" title={value.path}>
|
||||
<p className="mt-0.5 break-all text-xs text-text-tertiary" title={value.path}>
|
||||
路径:{value.path}
|
||||
</p>
|
||||
)}
|
||||
@@ -41,7 +41,7 @@ export function DiagnosticNotice({
|
||||
</IconButton>
|
||||
</div>
|
||||
{expanded && (
|
||||
<pre className="max-h-36 overflow-auto border-t border-danger-border bg-danger-soft p-3 text-xs text-danger">
|
||||
<pre className="max-h-36 overflow-auto whitespace-pre-wrap break-words border-t border-danger-border bg-danger-soft p-3 text-xs text-danger">
|
||||
{value.detail}
|
||||
</pre>
|
||||
)}
|
||||
|
||||
@@ -1,8 +1,24 @@
|
||||
import { useState } from 'react';
|
||||
import { CheckCircle2, Info, TriangleAlert, XCircle } from 'lucide-react';
|
||||
import { useMemo, useState } from 'react';
|
||||
import {
|
||||
CheckCircle2,
|
||||
FileInput,
|
||||
Info,
|
||||
TerminalSquare,
|
||||
TriangleAlert,
|
||||
XCircle,
|
||||
} from 'lucide-react';
|
||||
import { Button, CopyButton, Dialog, Tabs } from '../../components/ui';
|
||||
import type { WorkbenchNotification } from './NotificationCenter';
|
||||
type Filter = 'all' | 'warning' | 'danger';
|
||||
|
||||
type Filter = 'all' | 'import' | 'warning' | 'danger';
|
||||
|
||||
const toneMeta = {
|
||||
danger: { icon: XCircle, color: 'text-danger', label: 'ERROR' },
|
||||
warning: { icon: TriangleAlert, color: 'text-warning', label: 'WARN' },
|
||||
success: { icon: CheckCircle2, color: 'text-success', label: 'OK' },
|
||||
info: { icon: Info, color: 'text-accent', label: 'INFO' },
|
||||
} as const;
|
||||
|
||||
export function DiagnosticsDrawer({
|
||||
open,
|
||||
items,
|
||||
@@ -15,37 +31,59 @@ export function DiagnosticsDrawer({
|
||||
onClear: () => void;
|
||||
}) {
|
||||
const [filter, setFilter] = useState<Filter>('all');
|
||||
const groups = useMemo(
|
||||
() => ({
|
||||
all: items,
|
||||
import: items.filter((item) => item.category === 'import'),
|
||||
warning: items.filter((item) => item.tone === 'warning'),
|
||||
danger: items.filter((item) => item.tone === 'danger'),
|
||||
}),
|
||||
[items],
|
||||
);
|
||||
|
||||
const content = (value: Filter) => {
|
||||
const filtered = items.filter((item) => value === 'all' || item.tone === value);
|
||||
const filtered = groups[value];
|
||||
return (
|
||||
<div className="space-y-2">
|
||||
<div className="diagnostics-terminal panel-scroll min-h-32 max-h-[45vh] overflow-auto rounded-lg border border-border-subtle bg-input/75">
|
||||
{filtered.length ? (
|
||||
filtered.map((item) => {
|
||||
const Icon =
|
||||
item.tone === 'danger'
|
||||
? XCircle
|
||||
: item.tone === 'warning'
|
||||
? TriangleAlert
|
||||
: item.tone === 'success'
|
||||
? CheckCircle2
|
||||
: Info;
|
||||
const meta = toneMeta[item.tone];
|
||||
const Icon = meta.icon;
|
||||
const date = new Date(item.at);
|
||||
const validDate = Number.isFinite(date.getTime());
|
||||
return (
|
||||
<article key={item.id} className="rounded-lg border border-border bg-surface p-3">
|
||||
<div className="flex items-start gap-2">
|
||||
<Icon
|
||||
className={`mt-0.5 h-4 w-4 ${item.tone === 'danger' ? 'text-danger' : item.tone === 'warning' ? 'text-warning' : 'text-success'}`}
|
||||
/>
|
||||
<div className="min-w-0 flex-1">
|
||||
<h3 className="text-xs font-semibold">{item.title}</h3>
|
||||
<time className="text-[10px] text-text-tertiary">
|
||||
{new Date(item.at).toLocaleString('zh-CN')}
|
||||
</time>
|
||||
{item.detail && (
|
||||
<pre className="mt-2 whitespace-pre-wrap text-[10px] leading-4 text-text-secondary">
|
||||
{item.detail}
|
||||
</pre>
|
||||
<article
|
||||
key={item.id}
|
||||
className="grid grid-cols-[auto_minmax(0,1fr)] sm:grid-cols-[auto_minmax(0,1fr)_auto] gap-2 border-b border-border-subtle px-3 py-2.5 last:border-0 hover:bg-element-hover/45"
|
||||
>
|
||||
<Icon className={`mt-0.5 h-3.5 w-3.5 ${meta.color}`} aria-hidden="true" />
|
||||
<div className="min-w-0">
|
||||
<div className="flex min-w-0 items-baseline gap-2">
|
||||
<span className={`technical-value text-xs font-bold ${meta.color}`}>
|
||||
{meta.label}
|
||||
</span>
|
||||
<h3 className="break-words text-xs font-semibold text-text-primary">
|
||||
{item.title}
|
||||
</h3>
|
||||
{item.category && (
|
||||
<span className="technical-value rounded border border-border-subtle px-1 text-xs uppercase text-text-tertiary">
|
||||
{item.category}
|
||||
</span>
|
||||
)}
|
||||
</div>
|
||||
{item.detail && (
|
||||
<pre className="mt-1 min-w-0 whitespace-pre-wrap break-words font-mono text-xs leading-4 text-text-tertiary">
|
||||
{item.detail}
|
||||
</pre>
|
||||
)}
|
||||
</div>
|
||||
<div className="flex items-start gap-1.5">
|
||||
<time
|
||||
dateTime={validDate ? date.toISOString() : undefined}
|
||||
className="technical-value whitespace-nowrap text-xs text-text-tertiary"
|
||||
>
|
||||
{validDate ? date.toLocaleTimeString('zh-CN', { hour12: false }) : '—'}
|
||||
</time>
|
||||
{item.detail && (
|
||||
<CopyButton value={`${item.title}\n${item.detail}`} label="复制事件详情" />
|
||||
)}
|
||||
@@ -54,19 +92,28 @@ export function DiagnosticsDrawer({
|
||||
);
|
||||
})
|
||||
) : (
|
||||
<p className="p-8 text-center text-xs text-text-tertiary">没有符合条件的事件</p>
|
||||
<div className="grid min-h-32 place-items-center text-xs text-text-tertiary">
|
||||
没有符合条件的事件
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
return (
|
||||
<Dialog
|
||||
open={open}
|
||||
onClose={onClose}
|
||||
title="诊断与事件日志"
|
||||
className="max-w-2xl"
|
||||
placement="bottom"
|
||||
className="bg-panel"
|
||||
contentClassName="pt-3"
|
||||
footer={
|
||||
<div className="flex justify-end">
|
||||
<div className="flex items-center justify-between">
|
||||
<span className="flex items-center gap-1.5 text-xs text-text-tertiary">
|
||||
<TerminalSquare className="h-3.5 w-3.5 text-accent" />
|
||||
日志仅保存在当前浏览器会话;清空后不可恢复
|
||||
</span>
|
||||
<Button variant="danger" disabled={!items.length} onClick={onClear}>
|
||||
清空事件
|
||||
</Button>
|
||||
@@ -79,15 +126,21 @@ export function DiagnosticsDrawer({
|
||||
onValueChange={setFilter}
|
||||
keepMounted={false}
|
||||
items={[
|
||||
{ value: 'all', label: `全部 ${items.length}`, content: content('all') },
|
||||
{ value: 'all', label: `全部 ${groups.all.length}`, content: content('all') },
|
||||
{
|
||||
value: 'import',
|
||||
label: `导入 ${groups.import.length}`,
|
||||
icon: <FileInput className="h-3 w-3" />,
|
||||
content: content('import'),
|
||||
},
|
||||
{
|
||||
value: 'warning',
|
||||
label: `警告 ${items.filter((item) => item.tone === 'warning').length}`,
|
||||
label: `警告 ${groups.warning.length}`,
|
||||
content: content('warning'),
|
||||
},
|
||||
{
|
||||
value: 'danger',
|
||||
label: `错误 ${items.filter((item) => item.tone === 'danger').length}`,
|
||||
label: `错误 ${groups.danger.length}`,
|
||||
content: content('danger'),
|
||||
},
|
||||
]}
|
||||
|
||||
@@ -9,7 +9,7 @@ describe('EntrySelectionDialog', () => {
|
||||
],
|
||||
select = vi.fn();
|
||||
const { rerender } = render(<EntrySelectionDialog entries={entries} onSelect={select} />);
|
||||
const entry = screen.getByRole('button', { name: '模型 A' });
|
||||
const entry = screen.getByRole('button', { name: '模型 A a.xml' });
|
||||
entry.focus();
|
||||
rerender(<EntrySelectionDialog entries={[...entries]} onSelect={select} />);
|
||||
expect(entry).toHaveFocus();
|
||||
|
||||
@@ -15,11 +15,17 @@ export function EntrySelectionDialog({
|
||||
{entries.map((entry) => (
|
||||
<Button
|
||||
key={entry.path}
|
||||
className="w-full justify-start overflow-hidden"
|
||||
aria-label={`${entry.label} ${entry.path}`}
|
||||
className="!h-auto min-h-8 w-full justify-start py-2 text-left"
|
||||
onClick={() => onSelect(entry.path)}
|
||||
icon={<FileCode2 className="h-4 w-4" />}
|
||||
>
|
||||
<span className="truncate">{entry.label}</span>
|
||||
<span className="min-w-0">
|
||||
<span className="block break-words">{entry.label}</span>
|
||||
<span className="block break-all font-mono text-xs text-text-tertiary">
|
||||
{entry.path}
|
||||
</span>
|
||||
</span>
|
||||
</Button>
|
||||
))}
|
||||
</div>
|
||||
|
||||
@@ -17,7 +17,7 @@ export function ErrorRecoveryPanel({
|
||||
return (
|
||||
<section
|
||||
role="alert"
|
||||
className="absolute bottom-4 left-1/2 z-30 w-[min(42rem,calc(100%-2rem))] -translate-x-1/2 overflow-hidden rounded-xl border border-danger-border bg-panel shadow-2xl"
|
||||
className="w-full shrink-0 overflow-hidden rounded-xl border border-danger-border bg-panel shadow-2xl"
|
||||
>
|
||||
<div className="flex items-start gap-3 p-3">
|
||||
<span className="mt-0.5 grid h-7 w-7 shrink-0 place-items-center rounded-full bg-danger-soft text-danger">
|
||||
|
||||
@@ -52,6 +52,7 @@ describe('工作台反馈组件', () => {
|
||||
it('展示格式化状态数据', () => {
|
||||
render(<StatusBar time={1.25} fps={60} stepMs={0.5} memoryMb={10} loaded overBudget={false} />);
|
||||
expect(screen.getByText(/时间 1.250 s/)).toBeVisible();
|
||||
expect(screen.getByText(/WASM 已加载/)).toBeVisible();
|
||||
fireEvent.click(screen.getByRole('button', { name: /FPS 60/ }));
|
||||
expect(screen.getByRole('dialog', { name: '性能详情' })).toHaveTextContent('WASM已加载');
|
||||
});
|
||||
});
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import { Columns3, Focus, PanelLeft, PanelRight, RotateCcw } from 'lucide-react';
|
||||
import { Button, Dialog } from '../../components/ui';
|
||||
import { Button, Dialog, Tooltip } from '../../components/ui';
|
||||
export type LayoutPreset = 'default' | 'viewport' | 'project' | 'control';
|
||||
const presets = [
|
||||
{ value: 'default' as const, label: '默认布局', detail: '左右面板均衡显示', icon: Columns3 },
|
||||
@@ -49,19 +49,20 @@ export function LayoutSettingsDialog({
|
||||
<h3 className="mb-2 mt-4 text-xs font-semibold">布局预设</h3>
|
||||
<div className="grid grid-cols-2 gap-2">
|
||||
{presets.map((item) => (
|
||||
<button
|
||||
key={item.value}
|
||||
onClick={() => onPreset(item.value)}
|
||||
className="flex gap-2 rounded-lg border border-border bg-surface p-3 text-left hover:border-accent hover:bg-accent-soft focus-visible:ring-2 focus-visible:ring-accent/30"
|
||||
>
|
||||
<item.icon className="h-4 w-4 shrink-0 text-accent" />
|
||||
<span>
|
||||
<span className="block text-xs font-medium">{item.label}</span>
|
||||
<span className="mt-0.5 block text-[10px] text-text-tertiary">{item.detail}</span>
|
||||
</span>
|
||||
</button>
|
||||
<Tooltip key={item.value} content={item.detail}>
|
||||
<button
|
||||
onClick={() => onPreset(item.value)}
|
||||
className="flex gap-2 rounded-lg border border-border bg-surface p-3 text-left hover:border-accent hover:bg-accent-soft focus-visible:ring-2 focus-visible:ring-accent/30"
|
||||
>
|
||||
<item.icon className="h-4 w-4 shrink-0 text-accent" />
|
||||
<span>
|
||||
<span className="block text-xs font-medium">{item.label}</span>
|
||||
</span>
|
||||
</button>
|
||||
</Tooltip>
|
||||
))}
|
||||
</div>
|
||||
<p className="mt-3 text-xs text-text-tertiary">窄屏一次显示一侧面板,不覆盖桌面宽度偏好。</p>
|
||||
<Button
|
||||
className="mt-4 w-full"
|
||||
onClick={onReset}
|
||||
|
||||
@@ -64,7 +64,18 @@ describe('MapViewportToolbar', () => {
|
||||
});
|
||||
|
||||
describe('MapDraftStatusOverlay', () => {
|
||||
it('常驻显示未保存数量并提供提交与丢弃入口', () => {
|
||||
it('提交中与失败摘要不能隐藏,失败保留重试和放弃入口', () => {
|
||||
const props = { visible: true, changeCount: 2, onCommit: noop, onDiscard: noop };
|
||||
const { rerender } = render(<MapDraftStatusOverlay {...props} loading />);
|
||||
expect(screen.getByRole('status')).toHaveTextContent('正在应用草稿');
|
||||
expect(screen.getByRole('button', { name: '提交地图草稿' })).toBeDisabled();
|
||||
rerender(<MapDraftStatusOverlay {...props} loading={false} error="编译失败,草稿已保留" />);
|
||||
expect(screen.getByRole('status')).toHaveTextContent('应用失败');
|
||||
expect(screen.getByRole('status')).toHaveTextContent('编译失败,草稿已保留');
|
||||
expect(screen.getByRole('button', { name: '提交地图草稿' })).toBeEnabled();
|
||||
expect(screen.getByRole('button', { name: '丢弃地图草稿' })).toBeEnabled();
|
||||
});
|
||||
it('显示未保存数量并提供提交与丢弃入口', () => {
|
||||
const onCommit = vi.fn();
|
||||
const onDiscard = vi.fn();
|
||||
render(
|
||||
@@ -83,7 +94,7 @@ describe('MapDraftStatusOverlay', () => {
|
||||
expect(onDiscard).toHaveBeenCalledOnce();
|
||||
});
|
||||
|
||||
it('无草稿时显示已同步并禁用动作', () => {
|
||||
it('无草稿时不占位', () => {
|
||||
render(
|
||||
<MapDraftStatusOverlay
|
||||
visible
|
||||
@@ -93,7 +104,6 @@ describe('MapDraftStatusOverlay', () => {
|
||||
onDiscard={noop}
|
||||
/>,
|
||||
);
|
||||
expect(screen.getByText('地图草稿已同步')).toBeVisible();
|
||||
expect(screen.getByRole('button', { name: '提交地图草稿' })).toBeDisabled();
|
||||
expect(screen.queryByLabelText('地图草稿状态')).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
@@ -61,7 +61,7 @@ export function MapViewportToolbar({
|
||||
return (
|
||||
<div
|
||||
aria-label="地图视口工具"
|
||||
className="engineering-glass map-tool-enter absolute left-1/2 top-3 z-20 flex max-w-[calc(100%-24px)] -translate-x-1/2 items-center gap-1.5 rounded-xl border p-1.5"
|
||||
className="flex max-w-full flex-wrap items-center justify-center gap-1 p-1"
|
||||
>
|
||||
<button
|
||||
type="button"
|
||||
@@ -69,7 +69,7 @@ export function MapViewportToolbar({
|
||||
aria-pressed={interactionActive}
|
||||
disabled={loading}
|
||||
onClick={onActivate}
|
||||
className={`hidden h-8 items-center gap-1.5 rounded-lg px-2 text-[10px] font-semibold transition-colors sm:flex ${interactionActive ? 'bg-accent-soft text-accent' : 'bg-surface/80 text-text-secondary hover:bg-element-hover'}`}
|
||||
className={`flex h-8 items-center gap-1.5 rounded-lg px-2 text-xs font-semibold transition-colors ${interactionActive ? 'bg-accent-soft text-accent' : 'bg-surface/80 text-text-secondary hover:bg-element-hover'}`}
|
||||
>
|
||||
<MapPinned className="h-3.5 w-3.5" aria-hidden="true" />
|
||||
地图编辑
|
||||
@@ -91,8 +91,8 @@ export function MapViewportToolbar({
|
||||
<Grid3X3 className="h-3.5 w-3.5" />
|
||||
</IconButton>
|
||||
<span aria-hidden="true" className="mx-0.5 h-5 w-px bg-border" />
|
||||
<label className="flex items-center gap-1 text-[10px] text-text-tertiary">
|
||||
<span className="hidden lg:inline">贴地</span>
|
||||
<label className="flex items-center gap-1 text-xs text-text-tertiary">
|
||||
<span>贴地</span>
|
||||
<Select
|
||||
aria-label="贴地检测模式"
|
||||
className="w-[112px]"
|
||||
@@ -129,27 +129,34 @@ export function MapDraftStatusOverlay({
|
||||
loading,
|
||||
onCommit,
|
||||
onDiscard,
|
||||
error,
|
||||
}: {
|
||||
visible: boolean;
|
||||
changeCount: number;
|
||||
loading: boolean;
|
||||
error?: string;
|
||||
onCommit: () => void;
|
||||
onDiscard: () => void;
|
||||
}) {
|
||||
if (!visible) return null;
|
||||
const dirty = changeCount > 0;
|
||||
if (!visible || (!dirty && !loading && !error)) return null;
|
||||
return (
|
||||
<div
|
||||
aria-label="地图草稿状态"
|
||||
role="status"
|
||||
className="engineering-glass absolute bottom-3 left-1/2 z-20 flex max-w-[calc(100%-24px)] -translate-x-1/2 items-center gap-2 rounded-xl border px-2 py-1.5"
|
||||
className="engineering-glass flex max-w-full flex-wrap items-center justify-center gap-2 rounded-xl border px-2 py-1.5"
|
||||
>
|
||||
<span
|
||||
aria-hidden="true"
|
||||
className={`h-2 w-2 shrink-0 rounded-full ${dirty ? 'draft-dirty-dot bg-warning' : 'bg-success'}`}
|
||||
/>
|
||||
<span className="min-w-0 whitespace-nowrap text-[10px] font-medium text-text-secondary sm:text-[11px]">
|
||||
{dirty ? `${changeCount} 项未保存改动` : '地图草稿已同步'}
|
||||
<span className="min-w-0 text-xs font-medium text-text-secondary">
|
||||
{loading
|
||||
? '正在应用草稿…'
|
||||
: error
|
||||
? `应用失败 · ${changeCount} 项未保存改动`
|
||||
: `${changeCount} 项未保存改动`}
|
||||
{error && <span className="block text-danger">{error}</span>}
|
||||
</span>
|
||||
<Button
|
||||
variant="ghost"
|
||||
@@ -158,7 +165,7 @@ export function MapDraftStatusOverlay({
|
||||
icon={<Trash2 className="h-3 w-3" />}
|
||||
onClick={onDiscard}
|
||||
>
|
||||
<span className="hidden sm:inline">丢弃</span>
|
||||
放弃
|
||||
</Button>
|
||||
<Button
|
||||
variant="primary"
|
||||
@@ -167,7 +174,7 @@ export function MapDraftStatusOverlay({
|
||||
icon={<Check className="h-3 w-3" />}
|
||||
onClick={onCommit}
|
||||
>
|
||||
<span className="hidden sm:inline">提交</span>
|
||||
应用
|
||||
</Button>
|
||||
</div>
|
||||
);
|
||||
@@ -191,7 +198,7 @@ export function MapAssetDropIndicator({ target }: { target?: MapAssetDropTarget
|
||||
<span className="grid h-8 w-8 place-items-center rounded-full border border-current bg-panel/85 backdrop-blur">
|
||||
<MapPinned className="h-4 w-4" />
|
||||
</span>
|
||||
<span className="absolute left-1/2 top-full mt-2 -translate-x-1/2 whitespace-nowrap rounded-lg border border-border-strong bg-panel/90 px-2 py-1 text-[10px] font-medium shadow-xl backdrop-blur">
|
||||
<span className="absolute left-1/2 top-full mt-2 -translate-x-1/2 whitespace-nowrap rounded-lg border border-border-strong bg-panel/90 px-2 py-1 text-xs font-medium shadow-xl backdrop-blur">
|
||||
{valid
|
||||
? `释放放置 · ${target.position![0].toFixed(1)}, ${target.position![1].toFixed(1)}`
|
||||
: '请拖到 3D 地面'}
|
||||
|
||||
@@ -83,24 +83,30 @@ function props(
|
||||
}
|
||||
|
||||
describe('ModelControlsSidebar', () => {
|
||||
it('没有选择时显示场景摘要,并保留基础地图快速入口', () => {
|
||||
it('没有选择时仅显示收敛的场景摘要', () => {
|
||||
render(<ModelControlsSidebar {...props()} />);
|
||||
expect(screen.getByText('未选择对象')).toBeVisible();
|
||||
expect(screen.getByText('模型摘要')).toBeVisible();
|
||||
expect(screen.getByText('未选择地图实例')).toBeVisible();
|
||||
expect(screen.getByRole('button', { name: '模型摘要' })).toHaveAttribute(
|
||||
'aria-expanded',
|
||||
'false',
|
||||
);
|
||||
fireEvent.click(screen.getByRole('button', { name: '模型摘要' }));
|
||||
expect(screen.getByText('qpos / qvel')).toBeVisible();
|
||||
expect(screen.queryByText('未选择地图实例')).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it('按照统一选择自动路由 Body 与 Joint 检查器', () => {
|
||||
const view = render(
|
||||
<ModelControlsSidebar {...props({ selection: { kind: 'body', bodyId: 2 } })} />,
|
||||
);
|
||||
expect(screen.getByText('Robot / Body')).toBeVisible();
|
||||
expect(screen.getByText('arm', { selector: '[title="arm"]' })).toBeVisible();
|
||||
expect(screen.getByText('Body #2')).toBeVisible();
|
||||
expect(screen.getByText('arm')).toBeVisible();
|
||||
|
||||
view.rerender(
|
||||
<ModelControlsSidebar {...props({ selection: { kind: 'joint', jointId: 7, bodyId: 2 } })} />,
|
||||
);
|
||||
expect(screen.getByText('Robot / Joint')).toBeVisible();
|
||||
expect(screen.getByRole('slider', { name: 'arm_joint' })).toBeVisible();
|
||||
expect(screen.getByText('Hinge · arm')).toBeVisible();
|
||||
expect(screen.getByText('关联 Actuator')).toBeVisible();
|
||||
});
|
||||
@@ -153,8 +159,8 @@ describe('ModelControlsSidebar', () => {
|
||||
})}
|
||||
/>,
|
||||
);
|
||||
expect(screen.getByText('Map / Instance')).toBeVisible();
|
||||
expect(screen.getByText('随机粗糙地形', { selector: '[title="随机粗糙地形"]' })).toBeVisible();
|
||||
expect(screen.getByText('参数化地形')).toBeVisible();
|
||||
expect(screen.getByText('随机粗糙地形', { selector: 'header span' })).toBeVisible();
|
||||
fireEvent.click(screen.getByRole('tab', { name: '控制台' }));
|
||||
fireEvent.click(screen.getByRole('tab', { name: '数据录制' }));
|
||||
expect(onWorkspaceToolChange.mock.calls).toEqual([['controls'], ['data']]);
|
||||
|
||||
@@ -105,7 +105,10 @@ export function ModelControlsSidebar(props: ModelControlsProps) {
|
||||
{props.workspaceTool ? (
|
||||
props.workspaceTools
|
||||
) : (
|
||||
<div className="p-4 text-sm text-text-tertiary">导入模型后显示检查器</div>
|
||||
<div className="m-2.5 flex h-14 items-center justify-center gap-2 border border-dashed border-border-subtle text-xs text-text-tertiary">
|
||||
<PanelRight className="h-4 w-4" aria-hidden="true" />
|
||||
检查器待命
|
||||
</div>
|
||||
)}
|
||||
</SidebarPanel>
|
||||
);
|
||||
@@ -231,7 +234,6 @@ export function ModelControlsSidebar(props: ModelControlsProps) {
|
||||
onBaseMode={props.onBaseMode}
|
||||
onShowCollision={props.onShowCollision}
|
||||
/>
|
||||
{props.mapSelection.kind === 'none' && mapPanel}
|
||||
</>
|
||||
);
|
||||
}
|
||||
@@ -247,7 +249,9 @@ export function ModelControlsSidebar(props: ModelControlsProps) {
|
||||
>
|
||||
{inspector}
|
||||
</div>
|
||||
{props.workspaceTool && props.workspaceTools}
|
||||
<div hidden={!props.workspaceTool} className="min-h-0 flex-1 overflow-auto panel-scroll">
|
||||
{props.workspaceTools}
|
||||
</div>
|
||||
</SidebarPanel>
|
||||
);
|
||||
}
|
||||
|
||||
@@ -4,7 +4,11 @@ import { Badge, IconButton, Popover } from '../../components/ui';
|
||||
export interface WorkbenchNotification {
|
||||
id: number;
|
||||
title: string;
|
||||
/** 通知中心/诊断抽屉中的完整技术详情。 */
|
||||
detail?: string;
|
||||
/** Toast 使用的一行摘要,避免把日志常驻在视口上。 */
|
||||
message?: string;
|
||||
category?: 'import' | 'compile' | 'runtime' | 'system';
|
||||
tone: 'success' | 'warning' | 'danger' | 'info';
|
||||
at: number;
|
||||
}
|
||||
@@ -33,13 +37,13 @@ export function NotificationCenter({
|
||||
)}
|
||||
>
|
||||
{({ close }) => (
|
||||
<div className="w-80 overflow-hidden rounded-lg border border-border bg-surface-elevated shadow-xl">
|
||||
<div className="w-80 max-w-full overflow-hidden rounded-lg border border-border bg-surface-elevated shadow-xl">
|
||||
<header className="flex h-9 items-center justify-between border-b border-border px-3">
|
||||
<h2 className="text-xs font-semibold">通知</h2>
|
||||
<div className="flex gap-2">
|
||||
{onOpenLog && (
|
||||
<button
|
||||
className="text-[10px] text-accent"
|
||||
className="min-h-7 text-xs text-accent"
|
||||
onClick={() => {
|
||||
close();
|
||||
onOpenLog();
|
||||
@@ -50,7 +54,7 @@ export function NotificationCenter({
|
||||
)}
|
||||
{items.length > 0 && (
|
||||
<button
|
||||
className="flex items-center gap-1 text-[10px] text-text-tertiary hover:text-danger"
|
||||
className="flex min-h-7 items-center gap-1 text-xs text-text-tertiary hover:text-danger"
|
||||
onClick={onClear}
|
||||
>
|
||||
<Trash2 className="h-3 w-3" />
|
||||
@@ -73,7 +77,7 @@ export function NotificationCenter({
|
||||
/>
|
||||
<div className="min-w-0 flex-1">
|
||||
<div className="flex items-center gap-2">
|
||||
<h3 className="truncate text-xs font-medium">{item.title}</h3>
|
||||
<h3 className="break-words text-xs font-medium">{item.title}</h3>
|
||||
<Badge>
|
||||
{new Date(item.at).toLocaleTimeString('zh-CN', {
|
||||
hour: '2-digit',
|
||||
@@ -82,9 +86,12 @@ export function NotificationCenter({
|
||||
</Badge>
|
||||
</div>
|
||||
{item.detail && (
|
||||
<p className="mt-1 line-clamp-3 text-[10px] leading-4 text-text-tertiary">
|
||||
{item.detail}
|
||||
</p>
|
||||
<details className="domain-details mt-1">
|
||||
<summary>事件详情</summary>
|
||||
<p className="whitespace-pre-wrap break-words text-xs text-text-secondary">
|
||||
{item.detail}
|
||||
</p>
|
||||
</details>
|
||||
)}
|
||||
</div>
|
||||
<IconButton
|
||||
@@ -127,13 +134,15 @@ export function ToastViewport({
|
||||
return (
|
||||
<div
|
||||
role="status"
|
||||
className="pointer-events-auto absolute right-4 top-4 z-30 flex w-80 gap-2 rounded-lg border border-border bg-surface-elevated p-3 shadow-xl"
|
||||
className="pointer-events-auto flex w-full gap-2 rounded-lg border border-border bg-surface-elevated p-3 shadow-xl"
|
||||
>
|
||||
<Icon className="h-4 w-4 shrink-0 text-accent" />
|
||||
<div className="min-w-0 flex-1">
|
||||
<p className="text-xs font-medium">{item.title}</p>
|
||||
{item.detail && (
|
||||
<p className="mt-1 line-clamp-2 text-[10px] text-text-tertiary">{item.detail}</p>
|
||||
{(item.message ?? item.detail) && (
|
||||
<p className="mt-1 line-clamp-3 text-xs text-text-tertiary">
|
||||
{item.message ?? item.detail}
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
@@ -1,36 +1,34 @@
|
||||
import { Activity, ChevronUp, Cpu, MemoryStick, TriangleAlert } from 'lucide-react';
|
||||
import { Activity, ChevronUp, MemoryStick, TriangleAlert } from 'lucide-react';
|
||||
import { Badge, Popover, PropertyRow, Separator } from '../../components/ui';
|
||||
export function PerformancePopover({
|
||||
fps,
|
||||
stepMs,
|
||||
memoryMb,
|
||||
overBudget,
|
||||
loaded,
|
||||
}: {
|
||||
fps: number;
|
||||
stepMs: number;
|
||||
memoryMb?: number;
|
||||
overBudget: boolean;
|
||||
loaded?: boolean;
|
||||
}) {
|
||||
return (
|
||||
<Popover
|
||||
label="性能详情"
|
||||
placement="top-left"
|
||||
placement="bottom-left"
|
||||
trigger={({ open, toggle }) => (
|
||||
<button
|
||||
type="button"
|
||||
aria-haspopup="dialog"
|
||||
aria-expanded={open}
|
||||
onClick={toggle}
|
||||
className="flex h-6 items-center gap-3 rounded px-1.5 hover:bg-element-hover focus-visible:ring-2 focus-visible:ring-accent/30"
|
||||
className="flex h-7 items-center gap-3 rounded px-1.5 hover:bg-element-hover focus-visible:ring-2 focus-visible:ring-accent/30"
|
||||
>
|
||||
<span className="flex items-center gap-1.5">
|
||||
<Activity className="h-3 w-3" />
|
||||
FPS {fps.toFixed(0)}
|
||||
</span>
|
||||
<span className="flex items-center gap-1.5">
|
||||
<Cpu className="h-3 w-3" />
|
||||
物理 {stepMs.toFixed(2)} ms
|
||||
</span>
|
||||
<ChevronUp className={`h-3 w-3 transition-transform ${open ? 'rotate-180' : ''}`} />
|
||||
</button>
|
||||
)}
|
||||
@@ -43,6 +41,9 @@ export function PerformancePopover({
|
||||
{overBudget ? '预算超限' : '运行正常'}
|
||||
</Badge>
|
||||
</div>
|
||||
{loaded !== undefined && (
|
||||
<PropertyRow label="WASM" value={loaded ? '已加载' : '未加载'} />
|
||||
)}
|
||||
<PropertyRow label="渲染帧率" value={`${fps.toFixed(0)} FPS`} />
|
||||
<PropertyRow label="物理步进" value={`${stepMs.toFixed(2)} ms`} />
|
||||
<PropertyRow
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import { ChevronRight, FolderRoot } from 'lucide-react';
|
||||
import type { ModelEntry } from '../../project/types';
|
||||
import { SearchableCombobox } from '../../components/ui';
|
||||
import { SearchableCombobox, Tooltip } from '../../components/ui';
|
||||
export function ProjectBreadcrumb({
|
||||
projectName,
|
||||
entries,
|
||||
@@ -17,21 +17,24 @@ export function ProjectBreadcrumb({
|
||||
const parts = selectedEntry?.split('/').filter(Boolean) ?? [];
|
||||
return (
|
||||
<div className="border-b border-border bg-surface px-3 py-2">
|
||||
<div
|
||||
aria-label="当前工程路径"
|
||||
className="flex min-w-0 items-center gap-1 text-[10px] text-text-tertiary"
|
||||
>
|
||||
<FolderRoot className="h-3 w-3 shrink-0 text-accent" />
|
||||
<span className="truncate">{projectName}</span>
|
||||
{parts.map((part, index) => (
|
||||
<span key={`${part}-${index}`} className="contents">
|
||||
<ChevronRight className="h-3 w-3 shrink-0" />
|
||||
<span className={`truncate ${index === parts.length - 1 ? 'text-text-primary' : ''}`}>
|
||||
{part}
|
||||
<Tooltip content={[projectName, selectedEntry].filter(Boolean).join('/')}>
|
||||
<div
|
||||
tabIndex={0}
|
||||
aria-label="当前工程路径"
|
||||
className="flex min-w-0 items-center gap-1 text-xs text-text-tertiary"
|
||||
>
|
||||
<FolderRoot className="h-3 w-3 shrink-0 text-accent" />
|
||||
<span className="truncate">{projectName}</span>
|
||||
{parts.map((part, index) => (
|
||||
<span key={`${part}-${index}`} className="contents">
|
||||
<ChevronRight className="h-3 w-3 shrink-0" />
|
||||
<span className={`truncate ${index === parts.length - 1 ? 'text-text-primary' : ''}`}>
|
||||
{part}
|
||||
</span>
|
||||
</span>
|
||||
</span>
|
||||
))}
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
</Tooltip>
|
||||
{entries.length > 1 && (
|
||||
<div className="mt-2">
|
||||
<SearchableCombobox
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import { useState } from 'react';
|
||||
import { FolderTree, Library, Search } from 'lucide-react';
|
||||
import { FolderTree, Library } from 'lucide-react';
|
||||
import type { MapEntry, ModelEntry } from '../../project/types';
|
||||
import {
|
||||
countProjectSearchResults,
|
||||
@@ -106,7 +106,7 @@ export function ProjectSidebar({
|
||||
const sceneMatches = countSceneSearchResults(snapshot, placedMaps, editorDocuments, sceneQuery);
|
||||
|
||||
return (
|
||||
<SidebarPanel title="场景大纲与资产中心" side="left" visible={visible} icon={<Library />}>
|
||||
<SidebarPanel title="场景与资源" side="left" visible={visible} icon={<Library />}>
|
||||
{projectName ? (
|
||||
<>
|
||||
<div className="flex shrink-0 items-center gap-2 border-b border-border px-3 py-2.5">
|
||||
@@ -114,12 +114,6 @@ export function ProjectSidebar({
|
||||
<div className="truncate text-sm font-medium text-accent" title={projectName}>
|
||||
{projectName}
|
||||
</div>
|
||||
<div className="mt-0.5 truncate text-[10px] text-text-tertiary">
|
||||
{snapshot
|
||||
? `${snapshot.bodies.filter((body) => body.id > 0 && !body.name.startsWith('__platform_map_')).length} Body`
|
||||
: '模型未加载'}{' '}
|
||||
· {placedMaps.length} 个地图实例 · {files.length} 个文件
|
||||
</div>
|
||||
</div>
|
||||
<Button variant="danger" onClick={onRemove} disabled={loading}>
|
||||
移除
|
||||
@@ -136,10 +130,6 @@ export function ProjectSidebar({
|
||||
secondClassName="flex flex-col"
|
||||
first={
|
||||
<>
|
||||
<div className="flex h-8 shrink-0 items-center gap-2 px-3 text-[11px] font-semibold uppercase tracking-wide text-text-tertiary">
|
||||
<Search className="h-3.5 w-3.5" aria-hidden="true" />
|
||||
场景大纲
|
||||
</div>
|
||||
<TreeSearchField
|
||||
value={sceneQuery}
|
||||
onChange={setSceneQuery}
|
||||
|
||||
@@ -23,7 +23,7 @@ export function RightSidebarTabs({
|
||||
<div
|
||||
role="tablist"
|
||||
aria-label="右侧工作区视图"
|
||||
className="grid shrink-0 grid-cols-3 gap-1 border-b border-border bg-panel-muted/40 p-2"
|
||||
className="grid shrink-0 grid-cols-3 gap-1 border-b border-border-subtle bg-panel-muted/35 p-1"
|
||||
>
|
||||
{VIEWS.map((view, index) => {
|
||||
const Icon = view.icon;
|
||||
@@ -35,9 +35,9 @@ export function RightSidebarTabs({
|
||||
role="tab"
|
||||
aria-selected={selected}
|
||||
tabIndex={selected ? 0 : -1}
|
||||
className={`flex min-w-0 items-center justify-center gap-1.5 rounded-md px-1 py-2 text-[11px] font-medium transition-colors ${
|
||||
className={`flex h-8 min-w-0 items-center justify-center gap-1 rounded px-1 text-xs font-medium transition-colors ${
|
||||
selected
|
||||
? 'bg-panel text-accent shadow-sm'
|
||||
? 'border border-border-subtle bg-panel text-accent shadow-sm tool-active-glow'
|
||||
: 'text-text-tertiary hover:bg-element-hover hover:text-text-primary'
|
||||
}`}
|
||||
onClick={() => onChange(view.value)}
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import { useMemo } from 'react';
|
||||
import { Box, Bot, Layers3, MapPinned, Trash2, Zap } from 'lucide-react';
|
||||
import { Button, EmptySearchState, SearchHighlight } from '../../components/ui';
|
||||
import { Button, EmptySearchState, SearchHighlight, Tooltip } from '../../components/ui';
|
||||
import type { EditableMapDocument } from '../../map/editor/types';
|
||||
import type { PlacedMapAsset } from '../../map/types';
|
||||
import {
|
||||
@@ -136,7 +136,7 @@ export function SceneOutliner({
|
||||
{pendingSceneChangeCount > 0 && (
|
||||
<div
|
||||
role="status"
|
||||
className="mb-2 rounded-lg border border-accent/30 bg-accent-soft p-2.5 text-[10px] text-text-secondary"
|
||||
className="mb-2 rounded-lg border border-accent/30 bg-accent-soft p-2.5 text-xs text-text-secondary"
|
||||
>
|
||||
<div className="flex items-center gap-1.5 font-medium text-accent">
|
||||
<Zap className="h-3.5 w-3.5" aria-hidden="true" />
|
||||
@@ -158,7 +158,7 @@ export function SceneOutliner({
|
||||
<summary className="flex cursor-pointer select-none items-center gap-2 rounded px-1.5 py-1.5 text-xs font-semibold text-text-primary hover:bg-element-hover">
|
||||
<Bot className="h-3.5 w-3.5 text-accent" aria-hidden="true" />
|
||||
<span className="min-w-0 flex-1 truncate">机器人</span>
|
||||
<span className="technical-value text-[9px] font-normal text-text-tertiary">
|
||||
<span className="technical-value text-xs font-normal text-text-tertiary">
|
||||
{robotBodies.length} Body
|
||||
</span>
|
||||
</summary>
|
||||
@@ -182,7 +182,7 @@ export function SceneOutliner({
|
||||
<summary className="flex cursor-pointer select-none items-center gap-2 rounded px-1.5 py-1.5 text-xs font-semibold text-text-primary hover:bg-element-hover">
|
||||
<Layers3 className="h-3.5 w-3.5 text-accent" aria-hidden="true" />
|
||||
<span className="min-w-0 flex-1 truncate">地图与环境</span>
|
||||
<span className="technical-value text-[9px] font-normal text-text-tertiary">
|
||||
<span className="technical-value text-xs font-normal text-text-tertiary">
|
||||
{maps.length} 实例
|
||||
</span>
|
||||
</summary>
|
||||
@@ -213,20 +213,22 @@ export function SceneOutliner({
|
||||
<SearchHighlight text={asset.name} query={query} />
|
||||
</span>
|
||||
{pending.has(asset.id) && (
|
||||
<span className="rounded bg-warning/10 px-1 text-[9px] text-warning">
|
||||
<span className="rounded bg-warning/10 px-1 text-xs text-warning">
|
||||
待应用
|
||||
</span>
|
||||
)}
|
||||
</button>
|
||||
<button
|
||||
type="button"
|
||||
aria-label={`删除地图实例 ${asset.name}`}
|
||||
className="mr-1 rounded p-1 text-text-tertiary hover:bg-danger/10 hover:text-danger"
|
||||
disabled={loading}
|
||||
onClick={() => onRemoveMap(asset.id)}
|
||||
>
|
||||
<Trash2 className="h-3 w-3" aria-hidden="true" />
|
||||
</button>
|
||||
<Tooltip content={`从场景移除 ${asset.name};应用前可放弃更改。`}>
|
||||
<button
|
||||
type="button"
|
||||
aria-label={`删除地图实例 ${asset.name}`}
|
||||
className="mr-1 grid h-7 w-7 shrink-0 place-items-center rounded p-1 text-text-tertiary hover:bg-danger/10 hover:text-danger"
|
||||
disabled={loading}
|
||||
onClick={() => onRemoveMap(asset.id)}
|
||||
>
|
||||
<Trash2 className="h-3 w-3" aria-hidden="true" />
|
||||
</button>
|
||||
</Tooltip>
|
||||
</div>
|
||||
{objects.length > 0 && (
|
||||
<ul role="group" className="ml-3 border-l border-border pl-1">
|
||||
@@ -242,7 +244,7 @@ export function SceneOutliner({
|
||||
role="treeitem"
|
||||
aria-selected={selected}
|
||||
disabled={loading}
|
||||
className={`flex w-full items-center gap-1.5 rounded px-1.5 py-1 text-left text-[11px] ${
|
||||
className={`flex w-full items-center gap-1.5 rounded px-1.5 py-1 text-left text-xs ${
|
||||
selected
|
||||
? 'bg-accent-soft text-accent'
|
||||
: 'text-text-secondary hover:bg-element-hover'
|
||||
@@ -253,7 +255,7 @@ export function SceneOutliner({
|
||||
<span className="min-w-0 flex-1 truncate">
|
||||
<SearchHighlight text={object.name} query={query} />
|
||||
</span>
|
||||
<span className="text-[9px] text-text-tertiary">{object.type}</span>
|
||||
<span className="text-xs text-text-tertiary">{object.type}</span>
|
||||
</button>
|
||||
</li>
|
||||
);
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
import { fireEvent, render, screen } from '@testing-library/react';
|
||||
import { ShortcutHelpDialog } from './ShortcutHelpDialog';
|
||||
import { TreeSearchField } from './TreeSearchField';
|
||||
import { ViewportHUD } from './ViewportHUD';
|
||||
import { EmptyWorkspace } from './WorkspaceOverlays';
|
||||
import { ViewerDisplayPopover } from './ViewerDisplayPopover';
|
||||
import { DEFAULT_VIEWER_DISPLAY_OPTIONS } from '../../viewer/displayOptions';
|
||||
@@ -25,24 +24,15 @@ describe('第二批工作台组件', () => {
|
||||
fireEvent.keyDown(document, { key: 'Escape' });
|
||||
expect(close).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
it('视口 HUD 复用状态并给出当前模式的鼠标提示', () => {
|
||||
render(
|
||||
<ViewportHUD
|
||||
ready
|
||||
paused={false}
|
||||
mode="joint"
|
||||
selection={{ bodyId: 2, bodyName: 'arm', geomId: 3, geomType: 1, position: [0, 0, 0] }}
|
||||
/>,
|
||||
);
|
||||
expect(screen.getByLabelText('视口状态')).toHaveTextContent('仿真中');
|
||||
expect(screen.getByLabelText('视口状态')).toHaveTextContent('关节拖动');
|
||||
expect(screen.getByLabelText('视口状态')).toHaveTextContent('arm');
|
||||
expect(screen.getByLabelText('视口操作提示')).toHaveTextContent('左键拖动关节');
|
||||
expect(screen.getByLabelText('视口操作提示')).toHaveTextContent('右键平移');
|
||||
});
|
||||
it('空工作区解释导入到仿真的三步流程', () => {
|
||||
render(<EmptyWorkspace />);
|
||||
expect(screen.getByRole('region', { name: '导入模型工程' })).toBeVisible();
|
||||
expect(screen.getByText('Local Simulation Workspace')).toBeVisible();
|
||||
expect(screen.getByText('本地解析、编译并运行机器人模型,无需上传资源')).toBeVisible();
|
||||
expect(screen.getByLabelText('支持格式')).toHaveTextContent('OBJ / STL / DAE');
|
||||
expect(screen.getByRole('button', { name: '选择文件' })).toBeVisible();
|
||||
expect(screen.getByRole('button', { name: '选择文件夹' })).toBeVisible();
|
||||
expect(screen.queryByText('导入与仿真流程')).not.toBeInTheDocument();
|
||||
expect(screen.getByRole('list', { name: '仿真工作流程' })).toHaveTextContent('导入');
|
||||
expect(screen.getByRole('list', { name: '仿真工作流程' })).toHaveTextContent('检查与配置');
|
||||
expect(screen.getByRole('list', { name: '仿真工作流程' })).toHaveTextContent('运行与调试');
|
||||
|
||||
@@ -47,8 +47,12 @@ export function SettingsDialog({
|
||||
</section>
|
||||
<section>
|
||||
<h3 className="mb-2 text-xs font-semibold">模型与控制</h3>
|
||||
<p className="mb-2 text-xs text-warning">
|
||||
外力强度立即作用于施力工具,可能改变模型运动。
|
||||
</p>
|
||||
<PropertyRow
|
||||
label="角度单位"
|
||||
description="只切换显示单位,不改变模型关节限位"
|
||||
value={
|
||||
<Select
|
||||
aria-label="设置角度单位"
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import { Dialog, Kbd, Separator } from '../../components/ui';
|
||||
const shortcuts = [
|
||||
['Ctrl / Cmd + K', '打开命令面板'],
|
||||
['Space', '播放 / 暂停'],
|
||||
['R', '重置仿真(非地图编辑)'],
|
||||
['1', '选择模式'],
|
||||
@@ -17,11 +18,14 @@ const shortcuts = [
|
||||
export function ShortcutHelpDialog({ open, onClose }: { open: boolean; onClose: () => void }) {
|
||||
return (
|
||||
<Dialog open={open} onClose={onClose} title="快捷键与视口操作">
|
||||
<p className="mb-4 text-xs text-text-tertiary">
|
||||
输入框与 Monaco 编辑器聚焦时,地图和仿真快捷键不抢占输入;源码 Ctrl/Cmd+S 保存并重新载入。
|
||||
</p>
|
||||
<section>
|
||||
<h3 className="mb-2 text-xs font-semibold text-text-primary">键盘快捷键</h3>
|
||||
<dl className="space-y-2">
|
||||
{shortcuts.map(([key, label]) => (
|
||||
<div key={key} className="flex items-center justify-between text-xs">
|
||||
<div key={key} className="flex flex-wrap items-center justify-between gap-2 text-xs">
|
||||
<dt className="text-text-secondary">{label}</dt>
|
||||
<dd>
|
||||
<Kbd>{key}</Kbd>
|
||||
|
||||
@@ -18,10 +18,10 @@ export function SidebarPanel({
|
||||
return (
|
||||
<ResizablePanel side={side} storageKey={`mujoco-${side}-sidebar-width`} visible={visible}>
|
||||
<aside
|
||||
className={`flex h-full w-full min-w-0 flex-col overflow-hidden bg-panel ${side === 'left' ? 'border-r' : 'border-l'} border-border`}
|
||||
className={`cyber-sidebar flex h-full w-full min-w-0 flex-col overflow-hidden ${side === 'left' ? 'border-r' : 'border-l'} border-border-subtle `}
|
||||
>
|
||||
<h2 className="flex h-10 shrink-0 items-center gap-2 border-b border-border bg-gradient-to-r from-panel via-surface/70 to-panel px-3 text-sm font-semibold tracking-tight text-text-primary shadow-[inset_0_-1px_0_rgb(255_255_255/0.025)]">
|
||||
<span className="grid h-6 w-6 place-items-center rounded-md bg-accent-soft text-accent [&>svg]:h-3.5 [&>svg]:w-3.5">
|
||||
<h2 className="flex h-9 shrink-0 items-center gap-2 border-b border-border-subtle px-3 text-xs font-semibold text-text-primary">
|
||||
<span className="text-accent [&>svg]:h-4 [&>svg]:w-4">
|
||||
{icon ?? <Settings2 aria-hidden="true" className="h-3.5 w-3.5" />}
|
||||
</span>
|
||||
{title}
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
import { fireEvent, render, screen } from '@testing-library/react';
|
||||
import { SimulationControls } from './SimulationControls';
|
||||
describe('SimulationControls', () => {
|
||||
it('迁移后保留播放、单步、重置、速度及就绪限制', () => {
|
||||
const pause = vi.fn(),
|
||||
step = vi.fn(),
|
||||
reset = vi.fn(),
|
||||
speed = vi.fn();
|
||||
const props = {
|
||||
ready: true,
|
||||
paused: true,
|
||||
loading: false,
|
||||
speed: 1,
|
||||
onTogglePause: pause,
|
||||
onStep: step,
|
||||
onReset: reset,
|
||||
onSpeed: speed,
|
||||
};
|
||||
const { rerender } = render(<SimulationControls {...props} />);
|
||||
fireEvent.click(screen.getByRole('button', { name: '▶ 播放' }));
|
||||
fireEvent.click(screen.getByRole('button', { name: '单步' }));
|
||||
fireEvent.click(screen.getByRole('button', { name: '重置' }));
|
||||
fireEvent.change(screen.getByLabelText('仿真速度'), { target: { value: '2' } });
|
||||
expect(pause).toHaveBeenCalledOnce();
|
||||
expect(step).toHaveBeenCalledOnce();
|
||||
expect(reset).toHaveBeenCalledOnce();
|
||||
expect(speed).toHaveBeenCalledWith(2);
|
||||
rerender(<SimulationControls {...props} paused={false} />);
|
||||
expect(screen.getByRole('button', { name: '⏸ 暂停' })).toBeEnabled();
|
||||
expect(screen.getByRole('button', { name: '单步' })).toBeDisabled();
|
||||
rerender(<SimulationControls {...props} ready={false} />);
|
||||
expect(screen.getByRole('button', { name: '▶ 播放' })).toBeDisabled();
|
||||
expect(screen.getByRole('button', { name: '重置' })).toBeDisabled();
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,64 @@
|
||||
import { Pause, Play, RotateCcw, StepForward } from 'lucide-react';
|
||||
import { Button, IconButton, Select } from '../../components/ui';
|
||||
export function SimulationControls({
|
||||
paused,
|
||||
ready,
|
||||
speed,
|
||||
loading,
|
||||
onTogglePause,
|
||||
onStep,
|
||||
onReset,
|
||||
onSpeed,
|
||||
}: {
|
||||
paused: boolean;
|
||||
ready: boolean;
|
||||
speed: number;
|
||||
loading: boolean;
|
||||
onTogglePause: () => void;
|
||||
onStep: () => void;
|
||||
onReset: () => void;
|
||||
onSpeed: (speed: number) => void;
|
||||
}) {
|
||||
return (
|
||||
<div aria-label="仿真控制" className="flex flex-wrap items-center justify-center gap-1">
|
||||
<Button
|
||||
variant="ghost"
|
||||
onClick={onTogglePause}
|
||||
disabled={!ready}
|
||||
aria-label={paused ? '▶ 播放' : '⏸ 暂停'}
|
||||
icon={paused ? <Play className="h-3.5 w-3.5" /> : <Pause className="h-3.5 w-3.5" />}
|
||||
>
|
||||
{paused ? '播放' : '暂停'}
|
||||
</Button>
|
||||
<IconButton
|
||||
tooltip="单步(暂停时可用)"
|
||||
aria-label="单步"
|
||||
onClick={onStep}
|
||||
disabled={!ready || !paused}
|
||||
>
|
||||
<StepForward className="h-3.5 w-3.5" />
|
||||
</IconButton>
|
||||
<IconButton
|
||||
tooltip="重置仿真(R;地图编辑时 R 为缩放)"
|
||||
aria-label="重置"
|
||||
onClick={onReset}
|
||||
disabled={!ready}
|
||||
>
|
||||
<RotateCcw className="h-3.5 w-3.5" />
|
||||
</IconButton>
|
||||
<Select
|
||||
aria-label="仿真速度"
|
||||
value={speed}
|
||||
disabled={loading}
|
||||
onChange={(event) => onSpeed(Number(event.target.value))}
|
||||
className="w-[70px]"
|
||||
>
|
||||
<option value={0.25}>0.25×</option>
|
||||
<option value={0.5}>0.5×</option>
|
||||
<option value={1}>1×</option>
|
||||
<option value={2}>2×</option>
|
||||
<option value={4}>4×</option>
|
||||
</Select>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
import { fireEvent, render, screen } from '@testing-library/react';
|
||||
import { SourceEditorDialog } from './SourceEditorDialog';
|
||||
vi.mock('./monacoSetup', () => ({}));
|
||||
vi.mock('@monaco-editor/react', () => ({
|
||||
default: ({
|
||||
value,
|
||||
onChange,
|
||||
theme,
|
||||
}: {
|
||||
value: string;
|
||||
onChange(value: string): void;
|
||||
theme: string;
|
||||
}) => (
|
||||
<textarea
|
||||
aria-label="源码"
|
||||
data-theme={theme}
|
||||
value={value}
|
||||
onChange={(event) => onChange(event.target.value)}
|
||||
/>
|
||||
),
|
||||
}));
|
||||
it('保存失败保留草稿与风险,嵌套确认 Escape 只关闭确认', async () => {
|
||||
const close = vi.fn(),
|
||||
save = vi.fn().mockRejectedValue(new Error('编译失败,上一模型保留'));
|
||||
render(
|
||||
<SourceEditorDialog
|
||||
open
|
||||
code="<mujoco/>"
|
||||
filePath="model.xml"
|
||||
theme="light"
|
||||
onClose={close}
|
||||
onSave={save}
|
||||
/>,
|
||||
);
|
||||
const editor = screen.getByLabelText('源码');
|
||||
expect(editor).toHaveAttribute('data-theme', 'light');
|
||||
expect(screen.getByText(/保存将重新编译/)).toBeVisible();
|
||||
fireEvent.change(editor, { target: { value: '<mujoco model="new"/>' } });
|
||||
editor.focus();
|
||||
fireEvent.keyDown(editor, { key: 's', ctrlKey: true });
|
||||
expect(await screen.findByRole('alert')).toHaveTextContent('编译失败,上一模型保留');
|
||||
expect(editor).toHaveValue('<mujoco model="new"/>');
|
||||
expect(save).toHaveBeenCalledWith('model.xml', '<mujoco model="new"/>');
|
||||
fireEvent.click(screen.getByRole('button', { name: '关闭源代码编辑器' }));
|
||||
expect(screen.getByRole('dialog', { name: '放弃未保存的修改?' })).toBeVisible();
|
||||
fireEvent.keyDown(document, { key: 'Escape' });
|
||||
expect(screen.queryByRole('dialog', { name: '放弃未保存的修改?' })).not.toBeInTheDocument();
|
||||
expect(close).not.toHaveBeenCalled();
|
||||
});
|
||||
@@ -1,3 +1,4 @@
|
||||
import { FloatingLayerContext, useFloatingLayer } from '../../components/ui/floating';
|
||||
import './monacoSetup';
|
||||
import Editor from '@monaco-editor/react';
|
||||
import {
|
||||
@@ -42,6 +43,7 @@ export function SourceEditorDialog({
|
||||
const [code, setCode] = useState(sourceCode),
|
||||
[savedCode, setSavedCode] = useState(sourceCode),
|
||||
[saving, setSaving] = useState(false),
|
||||
[saveError, setSaveError] = useState<string>(),
|
||||
[copied, setCopied] = useState(false),
|
||||
[maximized, setMaximized] = useState(false),
|
||||
[discardOpen, setDiscardOpen] = useState(false),
|
||||
@@ -61,13 +63,17 @@ export function SourceEditorDialog({
|
||||
const save = useCallback(async () => {
|
||||
if (!dirty || problem) return;
|
||||
setSaving(true);
|
||||
setSaveError(undefined);
|
||||
try {
|
||||
await onSave(filePath, code);
|
||||
setSavedCode(code);
|
||||
} catch (value) {
|
||||
setSaveError(value instanceof Error ? value.message : String(value));
|
||||
} finally {
|
||||
setSaving(false);
|
||||
}
|
||||
}, [code, dirty, filePath, onSave, problem]);
|
||||
const layer = useFloatingLayer({ open, roots: () => [dialog.current], dismiss: requestClose });
|
||||
useEffect(() => {
|
||||
if (!open) return;
|
||||
previousFocus.current =
|
||||
@@ -80,7 +86,7 @@ export function SourceEditorDialog({
|
||||
}, [open]);
|
||||
useEffect(() => {
|
||||
const key = (event: KeyboardEvent) => {
|
||||
if (discardOpen) return;
|
||||
if (!open || discardOpen || !dialog.current?.contains(document.activeElement)) return;
|
||||
if (
|
||||
(event.ctrlKey || event.metaKey) &&
|
||||
event.key.toLowerCase() === 's' &&
|
||||
@@ -89,14 +95,11 @@ export function SourceEditorDialog({
|
||||
) {
|
||||
event.preventDefault();
|
||||
void save();
|
||||
} else if (event.key === 'Escape') {
|
||||
event.preventDefault();
|
||||
requestClose();
|
||||
}
|
||||
};
|
||||
window.addEventListener('keydown', key);
|
||||
return () => window.removeEventListener('keydown', key);
|
||||
}, [dirty, discardOpen, problem, requestClose, save]);
|
||||
}, [open, dirty, discardOpen, problem, requestClose, save]);
|
||||
const copy = async () => {
|
||||
await navigator.clipboard.writeText(code);
|
||||
setCopied(true);
|
||||
@@ -124,7 +127,7 @@ export function SourceEditorDialog({
|
||||
};
|
||||
if (!open) return null;
|
||||
return (
|
||||
<>
|
||||
<FloatingLayerContext.Provider value={layer.context}>
|
||||
<div className="fixed inset-0 z-[390] pointer-events-none" role="presentation">
|
||||
<section
|
||||
ref={dialog}
|
||||
@@ -133,12 +136,20 @@ export function SourceEditorDialog({
|
||||
aria-modal="false"
|
||||
aria-label="转换后的 MJCF 编辑器"
|
||||
style={
|
||||
maximized ? undefined : { left: position.x, top: position.y, width: 900, height: 650 }
|
||||
maximized
|
||||
? { zIndex: layer.zIndex }
|
||||
: {
|
||||
zIndex: layer.zIndex,
|
||||
left: `clamp(16px, ${position.x}px, max(16px, calc(100vw - 916px)))`,
|
||||
top: `clamp(16px, ${position.y}px, max(16px, calc(100vh - 666px)))`,
|
||||
width: 'min(900px, calc(100vw - 32px))',
|
||||
height: 'min(650px, calc(100vh - 32px))',
|
||||
}
|
||||
}
|
||||
className={`source-editor-window pointer-events-auto fixed flex min-h-[360px] min-w-[520px] flex-col overflow-hidden border border-border-strong bg-panel shadow-2xl ${maximized ? 'inset-0 h-full w-full' : 'resize'}`}
|
||||
className={`source-editor-window pointer-events-auto fixed flex min-h-64 min-w-0 max-w-[calc(100vw-32px)] max-h-[calc(100vh-32px)] flex-col overflow-hidden border border-border-strong ${maximized ? 'inset-0 h-full w-full' : 'resize'}`}
|
||||
>
|
||||
<header
|
||||
className="flex h-11 shrink-0 cursor-move select-none items-center gap-3 border-b border-border bg-surface px-3"
|
||||
className="flex min-h-11 flex-wrap shrink-0 py-2 cursor-move select-none items-center gap-3 border-b border-border bg-surface px-3"
|
||||
onPointerDown={pointerDown}
|
||||
onPointerMove={pointerMove}
|
||||
onPointerUp={() => {
|
||||
@@ -151,16 +162,16 @@ export function SourceEditorDialog({
|
||||
<div className="truncate font-mono text-xs font-semibold text-text-primary">
|
||||
转换后的 MJCF
|
||||
</div>
|
||||
<div className="truncate font-mono text-[9px] text-text-tertiary" title={filePath}>
|
||||
<div className="break-all font-mono text-xs text-text-tertiary" title={filePath}>
|
||||
{filePath}
|
||||
</div>
|
||||
</div>
|
||||
<span className="text-[10px] text-text-tertiary">{contentSize(code)}</span>
|
||||
<span className="rounded bg-accent-soft px-1.5 py-0.5 text-[9px] font-semibold text-accent">
|
||||
<span className="text-xs text-text-tertiary">{contentSize(code)}</span>
|
||||
<span className="rounded bg-accent-soft px-1.5 py-0.5 text-xs font-semibold text-accent">
|
||||
缓存文件 · 可编辑
|
||||
</span>
|
||||
{dirty && (
|
||||
<span className="rounded bg-warning-soft px-1.5 py-0.5 text-[9px] font-semibold text-warning">
|
||||
<span className="rounded bg-warning-soft px-1.5 py-0.5 text-xs font-semibold text-warning">
|
||||
已修改
|
||||
</span>
|
||||
)}
|
||||
@@ -193,6 +204,14 @@ export function SourceEditorDialog({
|
||||
<X className="h-4 w-4" />
|
||||
</IconButton>
|
||||
</header>
|
||||
<p className="border-b border-border px-3 py-2 text-xs text-warning">
|
||||
保存将重新编译并载入模型;未保存的源码关闭后无法恢复。
|
||||
</p>
|
||||
{saveError && (
|
||||
<p role="alert" className="px-3 py-2 text-xs text-danger">
|
||||
{saveError}
|
||||
</p>
|
||||
)}
|
||||
<div className="min-h-0 flex-1 bg-input">
|
||||
<Editor
|
||||
height="100%"
|
||||
@@ -218,8 +237,8 @@ export function SourceEditorDialog({
|
||||
}}
|
||||
/>
|
||||
</div>
|
||||
<footer className="flex h-7 shrink-0 items-center justify-between gap-3 border-t border-border bg-surface px-3 text-[10px]">
|
||||
<div className={problem ? 'truncate text-warning' : 'text-success'}>
|
||||
<footer className="flex min-h-7 flex-wrap shrink-0 items-center justify-between gap-3 border-t border-border bg-surface px-3 text-xs">
|
||||
<div className={problem ? 'break-words text-warning' : 'text-success'}>
|
||||
{problem ? `XML 错误:${problem}` : '✓ XML 结构正常'}
|
||||
</div>
|
||||
<div className="flex items-center gap-2 font-mono text-text-tertiary">
|
||||
@@ -243,6 +262,6 @@ export function SourceEditorDialog({
|
||||
当前 MJCF 源码包含未保存的修改。关闭后,这些修改将无法恢复。
|
||||
</p>
|
||||
</ConfirmDialog>
|
||||
</>
|
||||
</FloatingLayerContext.Provider>
|
||||
);
|
||||
}
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
import type { ReactNode } from 'react';
|
||||
import { Box, Clock3, MemoryStick, TriangleAlert } from 'lucide-react';
|
||||
import { Clock3, TriangleAlert } from 'lucide-react';
|
||||
import { useShallow } from 'zustand/react/shallow';
|
||||
import { Kbd } from '../../components/ui';
|
||||
import { useAppStore } from '../../stores/useAppStore';
|
||||
import { PerformancePopover } from './PerformancePopover';
|
||||
export interface StatusBarProps {
|
||||
@@ -12,47 +10,30 @@ export interface StatusBarProps {
|
||||
loaded: boolean;
|
||||
overBudget: boolean;
|
||||
}
|
||||
function Item({
|
||||
icon: Icon,
|
||||
children,
|
||||
className = '',
|
||||
}: {
|
||||
icon: typeof Clock3;
|
||||
children: ReactNode;
|
||||
className?: string;
|
||||
}) {
|
||||
return (
|
||||
<span className={`items-center gap-1.5 ${className || 'flex'}`}>
|
||||
<Icon aria-hidden="true" className="h-3 w-3 text-text-tertiary" />
|
||||
{children}
|
||||
</span>
|
||||
);
|
||||
}
|
||||
export function StatusBar({ time, fps, stepMs, memoryMb, loaded, overBudget }: StatusBarProps) {
|
||||
return (
|
||||
<footer className="technical-value relative z-30 flex h-7 shrink-0 items-center gap-3 overflow-hidden border-t border-border bg-panel px-3 text-[11px] text-text-tertiary lg:gap-5">
|
||||
<Item icon={Clock3}>时间 {time?.toFixed(3) ?? '—'} s</Item>
|
||||
<PerformancePopover fps={fps} stepMs={stepMs} memoryMb={memoryMb} overBudget={overBudget} />
|
||||
<Item icon={MemoryStick} className="hidden items-center gap-1.5 md:flex">
|
||||
内存 {memoryMb === undefined ? '—' : `${memoryMb.toFixed(1)} MiB`}
|
||||
</Item>
|
||||
<Item icon={Box} className="hidden items-center gap-1.5 sm:flex">
|
||||
WASM {loaded ? '已加载' : '未加载'}
|
||||
</Item>
|
||||
<div className="technical-value flex flex-wrap items-center gap-2 text-xs text-text-secondary">
|
||||
<span className="flex items-center gap-1 whitespace-nowrap">
|
||||
<Clock3 aria-hidden="true" className="h-3.5 w-3.5" />
|
||||
时间 {time?.toFixed(3) ?? '—'} s
|
||||
</span>
|
||||
<PerformancePopover
|
||||
fps={fps}
|
||||
stepMs={stepMs}
|
||||
memoryMb={memoryMb}
|
||||
loaded={loaded}
|
||||
overBudget={overBudget}
|
||||
/>
|
||||
{overBudget && (
|
||||
<span className="hidden min-w-0 items-center gap-1 truncate text-warning lg:flex">
|
||||
<TriangleAlert className="h-3 w-3 shrink-0" />
|
||||
主线程超出步进预算,已限制追帧
|
||||
<span className="flex items-center gap-1 text-warning">
|
||||
<TriangleAlert aria-hidden="true" className="h-3.5 w-3.5" />
|
||||
步进预算超限
|
||||
</span>
|
||||
)}
|
||||
<span className="ml-auto hidden items-center gap-1.5 xl:flex">
|
||||
<Kbd>Space</Kbd> 播放/暂停 · <Kbd>R</Kbd> 重置 · <Kbd>1/2/3</Kbd> 模式
|
||||
</span>
|
||||
</footer>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
/** 仅让状态栏订阅高频性能数据,避免带动整个工作台重渲染。 */
|
||||
/** 仅浮层订阅高频性能数据,不把 FPS/耗时提升到 App。 */
|
||||
export function StoreStatusBar() {
|
||||
const metrics = useAppStore(
|
||||
useShallow((state) => ({
|
||||
|
||||
@@ -28,7 +28,6 @@ export function ToolbarOverflowMenu({
|
||||
return (
|
||||
<DropdownMenu
|
||||
label="更多工作台操作"
|
||||
className="xl:hidden"
|
||||
items={[
|
||||
{
|
||||
id: 'commands',
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import { useState, type ReactNode } from 'react';
|
||||
import { Camera, Settings2 } from 'lucide-react';
|
||||
import type { CameraDirection, UrdfEnhancementOptions } from '../../project/urdfToMjcf';
|
||||
import { Button, Dialog, Select } from '../../components/ui';
|
||||
import { Button, Dialog, Select, Tooltip } from '../../components/ui';
|
||||
|
||||
function OptionCard({
|
||||
checked,
|
||||
@@ -84,9 +84,8 @@ export function UrdfImportOptionsDialog({
|
||||
</div>
|
||||
}
|
||||
>
|
||||
<p className="text-sm text-text-secondary">
|
||||
导入 <strong className="text-text-primary">{path}</strong>{' '}
|
||||
后,是否自动补充以下仿真组件?稍后重新选择该 URDF 时仍会再次询问。
|
||||
<p className="break-words text-sm text-text-secondary">
|
||||
导入 <strong className="text-text-primary">{path}</strong> 的仿真组件
|
||||
</p>
|
||||
<div className="mt-4 space-y-3">
|
||||
<OptionCard
|
||||
@@ -94,7 +93,7 @@ export function UrdfImportOptionsDialog({
|
||||
onChange={(addActuators) => setOptions((value) => ({ ...value, addActuators }))}
|
||||
icon={<Settings2 className="h-4 w-4" />}
|
||||
title="为关节添加驱动器"
|
||||
description="为每个 hinge/slide 关节生成控制输入不限幅的 motor 驱动器;hinge 使用 N·m、slide 使用 N。kp/kv 用于调整对应 MJCF 关节的刚度和阻尼,已有驱动器不会重复添加。"
|
||||
description="生成不限幅 motor 控制输入:hinge 为 N·m,slide 为 N。已有驱动器不重复添加。"
|
||||
/>
|
||||
<OptionCard
|
||||
checked={options.addSensors}
|
||||
@@ -104,9 +103,9 @@ export function UrdfImportOptionsDialog({
|
||||
description="在浮动基座添加三轴陀螺仪和三轴加速度计(6轴 IMU),并添加一台 640×480 固定摄像头。"
|
||||
/>
|
||||
{options.addSensors && (
|
||||
<div className="rounded-lg border border-border bg-surface p-3">
|
||||
<div className="border-t border-border pt-3">
|
||||
<div className="mb-2 text-xs font-medium text-text-primary">摄像头安装参数</div>
|
||||
<label className="block text-[11px] text-text-secondary">
|
||||
<label className="block text-xs text-text-secondary">
|
||||
<span className="mb-1 block">固连 Body</span>
|
||||
<Select
|
||||
aria-label="摄像头固连 Body"
|
||||
@@ -132,11 +131,14 @@ export function UrdfImportOptionsDialog({
|
||||
</label>
|
||||
<div className="mt-3 grid grid-cols-3 gap-2">
|
||||
{(['X', 'Y', 'Z'] as const).map((axis, index) => (
|
||||
<label key={axis} className="text-[11px] text-text-secondary">
|
||||
<span className="mb-1 block">位置 {axis}(m)</span>
|
||||
<label key={axis} className="text-xs text-text-secondary">
|
||||
<span className="mb-0.5 flex items-center gap-1">
|
||||
<span className={`axis-badge axis-${axis.toLowerCase()}`}>{axis}</span>
|
||||
位置(m)
|
||||
</span>
|
||||
<input
|
||||
aria-label={`摄像头位置 ${axis}`}
|
||||
className="field h-8 w-full px-2 text-xs"
|
||||
className="field technical-value h-8 w-full px-1.5 text-xs"
|
||||
type="number"
|
||||
step="0.01"
|
||||
value={(options.cameraPosition ?? [0.1, 0, 0.05])[index]}
|
||||
@@ -145,7 +147,7 @@ export function UrdfImportOptionsDialog({
|
||||
</label>
|
||||
))}
|
||||
</div>
|
||||
<label className="mt-3 block text-[11px] text-text-secondary">
|
||||
<label className="mt-3 block text-xs text-text-secondary">
|
||||
<span className="mb-1 block">镜头朝向(Body 局部轴)</span>
|
||||
<Select
|
||||
aria-label="摄像头朝向"
|
||||
@@ -163,9 +165,11 @@ export function UrdfImportOptionsDialog({
|
||||
))}
|
||||
</Select>
|
||||
</label>
|
||||
<p className="mt-2 text-[10px] leading-4 text-text-tertiary">
|
||||
位置和朝向均相对于所选 Body;常见 ROS 头部摄像头使用 +X 朝前、+Z 朝上。
|
||||
</p>
|
||||
<Tooltip content="位置和朝向相对于所选 Body;ROS 摄像头通常 +X 朝前、+Z 朝上">
|
||||
<span tabIndex={0} className="mt-2 inline-block text-xs text-text-tertiary">
|
||||
安装坐标说明
|
||||
</span>
|
||||
</Tooltip>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
|
||||
@@ -1,8 +1,6 @@
|
||||
import { Crosshair, Hand, MapPinned, MousePointer2, RotateCcw } from 'lucide-react';
|
||||
import { Crosshair, Hand, MousePointer2 } from 'lucide-react';
|
||||
import type { InteractionMode } from '../../viewer/MuJoCoViewer';
|
||||
import type { ViewerDisplayOptions } from '../../viewer/displayOptions';
|
||||
import { IconButton, ToolbarToggleGroup, type ToolbarItem } from '../../components/ui';
|
||||
import { ViewerDisplayPopover } from './ViewerDisplayPopover';
|
||||
import { ToolbarToggleGroup, type ToolbarItem } from '../../components/ui';
|
||||
const tools: ToolbarItem<InteractionMode>[] = [
|
||||
{ value: 'select', label: '选择', icon: MousePointer2 },
|
||||
{ value: 'joint', label: '关节拖动', icon: Hand },
|
||||
@@ -10,42 +8,12 @@ const tools: ToolbarItem<InteractionMode>[] = [
|
||||
];
|
||||
export function ViewerToolDock({
|
||||
mode,
|
||||
display,
|
||||
mapEditContext,
|
||||
onModeChange,
|
||||
onDisplayChange,
|
||||
onResetCamera,
|
||||
}: {
|
||||
mode: InteractionMode;
|
||||
display: ViewerDisplayOptions;
|
||||
mapEditContext?: {
|
||||
active: boolean;
|
||||
label: string;
|
||||
onActivate: () => void;
|
||||
};
|
||||
onModeChange: (mode: InteractionMode) => void;
|
||||
onDisplayChange: (next: ViewerDisplayOptions) => void;
|
||||
onResetCamera: () => void;
|
||||
}) {
|
||||
return (
|
||||
<div className="flex items-center gap-1">
|
||||
<ToolbarToggleGroup items={tools} value={mode} onChange={onModeChange} label="视口交互模式" />
|
||||
{mapEditContext && (
|
||||
<button
|
||||
type="button"
|
||||
aria-label={`地图编辑联动:${mapEditContext.label}`}
|
||||
aria-pressed={mapEditContext.active}
|
||||
onClick={mapEditContext.onActivate}
|
||||
className={`hidden h-7 items-center gap-1.5 rounded-lg border px-2 text-[10px] font-semibold transition-colors md:flex ${mapEditContext.active ? 'border-accent/40 bg-accent-soft text-accent' : 'border-border bg-surface/80 text-text-tertiary hover:bg-element-hover hover:text-text-primary'}`}
|
||||
>
|
||||
<MapPinned className="h-3.5 w-3.5" aria-hidden="true" />
|
||||
{mapEditContext.label}
|
||||
</button>
|
||||
)}
|
||||
<ViewerDisplayPopover value={display} onChange={onDisplayChange} />
|
||||
<IconButton tooltip="相机复位" aria-label="相机复位" onClick={onResetCamera}>
|
||||
<RotateCcw className="h-3.5 w-3.5" />
|
||||
</IconButton>
|
||||
</div>
|
||||
<ToolbarToggleGroup items={tools} value={mode} onChange={onModeChange} label="视口交互模式" />
|
||||
);
|
||||
}
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
import { act, fireEvent, render, screen } from '@testing-library/react';
|
||||
import { ViewportHUD } from './ViewportHUD';
|
||||
import { useAppStore } from '../../stores/useAppStore';
|
||||
describe('ViewportHUD', () => {
|
||||
beforeEach(() => useAppStore.getState().clearProject());
|
||||
it('独立更新性能,不带动父布局;细节按需展示,无手势长串', () => {
|
||||
const parent = vi.fn();
|
||||
function Layout() {
|
||||
parent();
|
||||
return <ViewportHUD />;
|
||||
}
|
||||
render(<Layout />);
|
||||
expect(screen.getByLabelText('视口状态')).toHaveTextContent('待导入');
|
||||
act(() => useAppStore.getState().setMetrics(60, 1.25, 32, false));
|
||||
expect(screen.getByRole('button', { name: /FPS 60/ })).toBeVisible();
|
||||
expect(screen.queryByText(/物理 1.25/)).not.toBeInTheDocument();
|
||||
expect(screen.queryByLabelText('视口操作提示')).not.toBeInTheDocument();
|
||||
expect(parent).toHaveBeenCalledOnce();
|
||||
fireEvent.click(screen.getByRole('button', { name: /FPS/ }));
|
||||
expect(screen.getByRole('dialog', { name: '性能详情' })).toHaveTextContent('1.25 ms');
|
||||
expect(screen.getByRole('dialog', { name: '性能详情' })).toHaveTextContent('WASM未加载');
|
||||
act(() => useAppStore.getState().setMetrics(30, 40, 32, true));
|
||||
expect(screen.getByLabelText('视口状态')).toHaveTextContent('步进预算超限');
|
||||
expect(parent).toHaveBeenCalledOnce();
|
||||
});
|
||||
});
|
||||
@@ -1,88 +1,26 @@
|
||||
import { CirclePause, CirclePlay, Mouse, MousePointer2 } from 'lucide-react';
|
||||
import type { InteractionMode, ViewerSelection } from '../../viewer/MuJoCoViewer';
|
||||
import { Badge, Kbd } from '../../components/ui';
|
||||
const labels: Record<InteractionMode, string> = {
|
||||
select: '选择',
|
||||
joint: '关节拖动',
|
||||
force: '外力施加',
|
||||
};
|
||||
const primaryGestures: Record<InteractionMode, string> = {
|
||||
select: '左键旋转',
|
||||
joint: '左键拖动关节',
|
||||
force: '左键拖动施力',
|
||||
};
|
||||
export function ViewportHUD({
|
||||
paused,
|
||||
mode,
|
||||
selection,
|
||||
ready,
|
||||
mapEditing = false,
|
||||
}: {
|
||||
paused: boolean;
|
||||
mode: InteractionMode;
|
||||
selection: ViewerSelection | null;
|
||||
ready: boolean;
|
||||
mapEditing?: boolean;
|
||||
}) {
|
||||
if (!ready) return null;
|
||||
import { CirclePause, CirclePlay } from 'lucide-react';
|
||||
import { useAppStore } from '../../stores/useAppStore';
|
||||
import { Badge } from '../../components/ui';
|
||||
import { StoreStatusBar } from './StatusBar';
|
||||
|
||||
export function ViewportHUD() {
|
||||
const paused = useAppStore((state) => state.paused);
|
||||
const ready = useAppStore((state) => Boolean(state.snapshot));
|
||||
return (
|
||||
<>
|
||||
<div
|
||||
aria-label="视口状态"
|
||||
className="pointer-events-none absolute left-3 top-3 z-10 flex max-w-[70%] flex-wrap items-center gap-1.5"
|
||||
>
|
||||
<Badge tone={paused ? 'neutral' : 'success'}>
|
||||
{paused ? <CirclePause className="h-3 w-3" /> : <CirclePlay className="h-3 w-3" />}
|
||||
{paused ? '已暂停' : '仿真中'}
|
||||
</Badge>
|
||||
<Badge tone="accent">
|
||||
<MousePointer2 className="h-3 w-3" />
|
||||
{labels[mode]}
|
||||
</Badge>
|
||||
{selection && (
|
||||
<Badge title={`body ${selection.bodyId} · geom ${selection.geomId}`}>
|
||||
{selection.bodyName}
|
||||
</Badge>
|
||||
)}
|
||||
</div>
|
||||
<div
|
||||
aria-label="视口操作提示"
|
||||
className="pointer-events-none absolute bottom-14 left-1/2 z-10 hidden -translate-x-1/2 items-center gap-2 whitespace-nowrap rounded-full border border-border-strong bg-panel/85 px-3 py-1.5 text-[10px] text-text-secondary shadow-xl backdrop-blur lg:flex"
|
||||
>
|
||||
<Mouse aria-hidden="true" className="h-3 w-3 text-text-secondary" />
|
||||
<span>{primaryGestures[mode]}</span>
|
||||
<span aria-hidden="true" className="text-border-strong">
|
||||
·
|
||||
</span>
|
||||
<span>右键平移</span>
|
||||
<span aria-hidden="true" className="text-border-strong">
|
||||
·
|
||||
</span>
|
||||
<span>滚轮缩放</span>
|
||||
{mapEditing && mode === 'select' ? (
|
||||
<>
|
||||
<span aria-hidden="true" className="text-border-strong">
|
||||
·
|
||||
</span>
|
||||
<Kbd>W</Kbd>
|
||||
<Kbd>E</Kbd>
|
||||
<Kbd>R</Kbd>
|
||||
<span>变换</span>
|
||||
<Kbd>F</Kbd>
|
||||
<span>聚焦</span>
|
||||
</>
|
||||
<div
|
||||
aria-label="视口状态"
|
||||
className="engineering-glass flex flex-wrap items-center gap-2 rounded-lg border border-border px-2 py-1 text-xs"
|
||||
>
|
||||
<Badge tone={!ready || paused ? 'neutral' : 'success'}>
|
||||
<span aria-hidden="true" className="hud-signal" data-running={ready && !paused} />
|
||||
{!ready || paused ? (
|
||||
<CirclePause className="h-3.5 w-3.5" />
|
||||
) : (
|
||||
mode !== 'select' && (
|
||||
<>
|
||||
<span aria-hidden="true" className="text-border-strong">
|
||||
·
|
||||
</span>
|
||||
<Kbd>1</Kbd>
|
||||
<span>旋转视角</span>
|
||||
</>
|
||||
)
|
||||
<CirclePlay className="h-3.5 w-3.5" />
|
||||
)}
|
||||
</div>
|
||||
</>
|
||||
{!ready ? '待导入' : paused ? '已暂停' : '仿真中'}
|
||||
</Badge>
|
||||
<StoreStatusBar />
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
import type { ReactNode } from 'react';
|
||||
|
||||
/** 仅分配空间;槽位内容自行订阅状态,禁止在此接入仿真业务。 */
|
||||
export function ViewportOverlayLayout({
|
||||
status,
|
||||
view,
|
||||
notices,
|
||||
context,
|
||||
camera,
|
||||
orientation,
|
||||
controls,
|
||||
draft,
|
||||
}: {
|
||||
status: ReactNode;
|
||||
view: ReactNode;
|
||||
notices?: ReactNode;
|
||||
context?: ReactNode;
|
||||
camera?: ReactNode;
|
||||
orientation?: ReactNode;
|
||||
controls: ReactNode;
|
||||
draft?: ReactNode;
|
||||
}) {
|
||||
return (
|
||||
<div className="viewport-overlays pointer-events-none absolute inset-0 z-20 p-3">
|
||||
<div className="viewport-overlay-top">
|
||||
<div data-overlay-slot="status">{status}</div>
|
||||
<div data-overlay-slot="view">{view}</div>
|
||||
</div>
|
||||
<div className="viewport-overlay-notices">
|
||||
<div data-overlay-slot="context">{context}</div>
|
||||
<div data-overlay-slot="notices">{notices}</div>
|
||||
</div>
|
||||
<div className="min-h-0 flex-1" />
|
||||
<div className="viewport-overlay-bottom">
|
||||
<div data-overlay-slot="camera">{camera}</div>
|
||||
<div data-overlay-slot="orientation">{orientation}</div>
|
||||
<div className="viewport-overlay-actions">
|
||||
<div data-overlay-slot="controls">{controls}</div>
|
||||
<div data-overlay-slot="draft">{draft}</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
@@ -2,57 +2,41 @@ import { fireEvent, render, screen } from '@testing-library/react';
|
||||
import { WorkbenchHeader } from './WorkbenchHeader';
|
||||
const fn = () => {};
|
||||
describe('WorkbenchHeader', () => {
|
||||
it('透传仿真动作且保留可访问名称', () => {
|
||||
const pause = vi.fn(),
|
||||
step = vi.fn(),
|
||||
reset = vi.fn(),
|
||||
speed = vi.fn(),
|
||||
openSource = vi.fn();
|
||||
it('保留导入、源代码直达、命令与侧栏入口,不再重复仿真控制', () => {
|
||||
const source = vi.fn(),
|
||||
files = vi.fn(),
|
||||
left = vi.fn();
|
||||
render(
|
||||
<WorkbenchHeader
|
||||
paused
|
||||
ready
|
||||
speed={1}
|
||||
theme="dark"
|
||||
loading={false}
|
||||
hasProject
|
||||
leftOpen
|
||||
rightOpen
|
||||
fullscreen={false}
|
||||
center={<span>工具</span>}
|
||||
onFiles={fn}
|
||||
onFiles={files}
|
||||
onFolder={fn}
|
||||
onOpenSource={openSource}
|
||||
onTogglePause={pause}
|
||||
onStep={step}
|
||||
onReset={reset}
|
||||
onSpeed={speed}
|
||||
onToggleLeft={fn}
|
||||
onOpenSource={source}
|
||||
onToggleLeft={left}
|
||||
onToggleRight={fn}
|
||||
onToggleTheme={fn}
|
||||
onHelp={fn}
|
||||
onCommands={fn}
|
||||
onToggleFullscreen={fn}
|
||||
/>,
|
||||
);
|
||||
const sourceButton = screen.getByRole('button', { name: '源代码' });
|
||||
expect(sourceButton).toHaveTextContent('源代码');
|
||||
fireEvent.click(sourceButton);
|
||||
fireEvent.click(screen.getByRole('button', { name: '▶ 播放' }));
|
||||
fireEvent.click(screen.getByRole('button', { name: '单步' }));
|
||||
fireEvent.click(screen.getByRole('button', { name: '重置' }));
|
||||
fireEvent.change(screen.getByLabelText('仿真速度'), { target: { value: '2' } });
|
||||
expect(openSource).toHaveBeenCalledTimes(1);
|
||||
expect(pause).toHaveBeenCalledTimes(1);
|
||||
expect(step).toHaveBeenCalledTimes(1);
|
||||
expect(reset).toHaveBeenCalledTimes(1);
|
||||
expect(speed).toHaveBeenCalledWith(2);
|
||||
expect(screen.getByRole('button', { name: '切换到白天主题' })).toBeInTheDocument();
|
||||
expect(screen.getByRole('heading', { name: 'MuJoCo' })).toBeVisible();
|
||||
fireEvent.change(screen.getByLabelText('打开文件'), {
|
||||
target: { files: [new File(['x'], 'model.xml')] },
|
||||
});
|
||||
expect(files).toHaveBeenCalledOnce();
|
||||
expect(screen.getByLabelText('打开文件夹')).toHaveAttribute('webkitdirectory');
|
||||
expect(screen.getByRole('button', { name: '源代码' })).toBeVisible();
|
||||
expect(screen.queryByRole('button', { name: '工程' })).not.toBeInTheDocument();
|
||||
fireEvent.click(screen.getByRole('button', { name: '源代码' }));
|
||||
expect(source).toHaveBeenCalledOnce();
|
||||
fireEvent.click(screen.getByRole('button', { name: '隐藏工程面板' }));
|
||||
expect(left).toHaveBeenCalledOnce();
|
||||
expect(screen.getByRole('button', { name: '隐藏工程面板' })).toHaveAttribute(
|
||||
'aria-expanded',
|
||||
'true',
|
||||
);
|
||||
expect(screen.getByRole('button', { name: '打开命令面板' })).toBeInTheDocument();
|
||||
expect(screen.getByRole('button', { name: '进入全屏' })).toBeInTheDocument();
|
||||
expect(screen.getByRole('button', { name: '打开命令面板' })).toBeVisible();
|
||||
expect(screen.queryByRole('button', { name: '▶ 播放' })).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user