feat(training): release V0.9.1 避障训练与基础策略迁移
web-platform-ci / TypeScript, lint, unit, build (push) Has been cancelled
web-platform-ci / Playwright E2E (push) Has been cancelled

This commit is contained in:
2026-09-08 10:50:13 +08:00
parent fa5485049a
commit 438e56bcc8
113 changed files with 15027 additions and 539 deletions
+2 -1
View File
@@ -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文件
+96
View File
@@ -2,6 +2,82 @@
本项目的重要变更记录在此文件中,版本标签沿用仓库现有的 `V主版本.次版本[.修订版本]` 格式。
## [0.9.1] - 2026-09-08
- Go2避障训练改为每个episode从同一连通自由区域随机采样起终点并随机化初始朝向,保证障碍/边界净空和最小目标距离;固定三seed评估及浏览器部署仍使用声明的参考起终点,兼顾泛化与可比性。
- 修正基础策略上传存储错误分类:文件打开/写入/关闭失败返回503并提示检查磁盘空间/权限;流读取超时或连接截断仍返回400,均保留临时清理与socket超时复位。
- 完成普通训练与自调参共用单文件.pt/.onnx上传入口:默认服务直接连接、Go2 legacy47模板确认、自动选择内容ID/文件SHA、取消及同文件重试;任务/连接epoch变化保留旧选择并隔离晚到响应,上传中禁止启动。补双面板真实HTTP浏览器验收与4环境×1iteration CPU ONNX-derived PPO/导出/统计恢复检查(仅链路验证)。
- 新增Go2 legacy47单文件.pt/受限MLP ONNX上传后端,无需管理员注册/邻接配置;认证有界二进制接收、CPU限额安全校验、持久化内容SHA来源,复用普通训练/自调参warm-start与rung自身resume。ONNX只继承确定性网络,明确模板假设、合成normalizer count及fresh训练状态;双UI文件选择待接入。
- 修复跨任务奖励preset混入Flat:以持久化来源session恢复权威taskId,读写与训练入口完整验证任务schema;跨任务或损坏来源在创建作业前拒绝,Flat菜单仅展示已确认Flat的preset。
- 修复导航设定拖回起点误触:超过5px后持续记为拖动,取消/Esc/清理/卸载清空整个手势,下一次正常点击不受影响。
- 新增显式multi_ring_raycast三层48射线/97维部署,旧缺省32射线/81维不变;task/ONNX严格模式、pitch/yaw角度及顺序白名单,浏览器观测、真实ORT shape、导航面板/PiP动态线数贯通。
- multi奖励仅按已验证标准底板顶面几何分类过滤地面,观测保留实际地板距离;缓存静态box/方向并优化解析slab,补真实CPU/WASM/GPU短smoke及48ray整帧基准。仍存在层间/侧后/坑盲区,未实现局部高程图。
- 修复训练CLI未导入仓库任务注册模块导致直接启动Flat/Rough/Obstacle失败;保留独立解释器三seed评估与调参导航速度。事务导入失败面板改为保留完整diagnostic detail,不再用“模型编译失败”摘要掩盖ORT shape/初始化错误。
- 新增已应用场景静态碰撞多实例→权威custom_boxes导出,保留世界坐标并诚实标记AABB/标准底板近似;前后端严格校验布局、摩擦、数量及起终点圆形安全区,贯通任务JSON/ONNX与最小坐标配置。
- 新增 Go2 避障目标贴地呼吸信标、一次性点击设定目标模式、实时目标坐标/距离和目标复位;保持81维观测与异步held-action,输入与地图操纵器互斥。
- 本地训练卡片新增可折叠Loss/Reward趋势,按迭代去重补全并有界保留500点;复用共享Canvas图表,综合指标独立缩放,曲线平滑但保留原值。
- 新增 Go2 前视32射线避障训练任务、自定义地形参数、配套碰撞地图与81维ONNX浏览器部署闭环,支持PiP和射线显示。
- 自定义地图与策略采用候选会话事务,真实ONNX初始化/graph校验或地图绑定失败时保留原场景、策略和物理状态;默认Flat兼容旧无metadata的47维导出。
- 对齐原Go2训练模型的关节armature、足端接触参数和自碰撞mask,补充真实WASM/CPU射线、动力学参数及浏览器回归测试。短训练仅验证链路,不代表避障收敛。
## [0.8.3] - 2026-09-04
### 新增
- 增加视口浮动地图工具栏(`MapViewportTools`),在视口上方提供变换模式切换、坐标空间切换和快捷对齐工具。
- 增加三维场景大纲树(`SceneOutliner`),直观查看场景中机器人、地图实例及层级结构。
- 增加视口快捷键挂载(`useMapEditorShortcuts`),支持快速切换操纵器模式、聚焦对象和撤销操作。
- 增加属性检查器面板(`MapObjectInspector` 与 `RobotInspector`),支持精确调节地图对象位姿与机器人关节/执行器。
- 增加右侧栏自适应 Tab 与多工具工作区容器(`RightSidebarTabs`、`WorkspaceToolsPanel`)。
- 增加拖拽式连续数值调节输入组件(`ScrubbableNumberInput`)与通用垂直双栏拆分面板(`VerticalSplitPane`)。
### 变更
- 重构并收敛模型控制侧栏(`ModelControlsSidebar`)与工程侧栏(`ProjectSidebar`),移除冗余的行内变换逻辑。
- 优化地图编辑面板(`MapEditorPanel`)与物理地图面板(`PhysicalMapPanel`),与统一检查器架构对齐。
- 清理已落地的历史设计草案与计划文档。
## [0.8.2] - 2026-09-03
### 新增
- 将自调参工作台重构为高内聚的组件群:Session 列表导航轨(`TuningSessionRail`)、控制工具栏(`TuningControlToolbar`)、Agent 决策时间线(`AgentDecisionTimeline`)、指标对比看板(`MetricsComparisonBoard`)、排行榜(`TuningLeaderboard`)、日志控制台(`TuningConsole`)与参数差异对比器(`RewardConfigDiffEditor`)。
- 建立自调参集中式状态管理(`tuningStore.ts`)与自适应轮询逻辑(`useTuningPolling.ts`)。
- 引入标量环形缓冲区(`ScalarRingBuffer.ts`),优化长周期 TensorBoard 曲线的高频更新与渲染性能。
- 服务端增加单步执行令牌(Step Token)调度,支持逐轮审批模式下的受控单步推进。
- 服务端增加参数约束 CAS 校验护栏与会话级别约束同步。
- 支持同会话内安全 Trial 的一键回滚与基准重设(Rollback)。
## [0.8.1] - 2026-09-03
### 新增
- 统一地图资产库与多实例堆叠系统(`MapAssetLibrary`、`MapStackComposer`),支持同源地图资产多次放置与独立位姿变换。
- 新增参数化地图实时预览图层(`ParametricMapPreviewLayer`),在视口中实时预览程序化地形与碰撞体变换。
- 新增程序地形贴合与落位计算(`sceneSurface.ts`),支持认证资产在程序化地形上自动重力落位。
- 引入地图场景草稿事务化编译(`mapSceneDraft`、`EditableMapDraftCommit`),保障物理、视觉与出生点数据同步更新。
- 增强地图对象拾取与选择机制(`compiledMapPick`),支持点击穿透、空白取消与实时变换操纵。
## [0.8.0] - 2026-09-02
### 新增
- 增加基于 DeepSeek Agent 的强化学习奖励函数自调参系统(`training_server/tuning/`),结合训练曲线与固定评估指标提出受限参数 patch。
- 支持全自动(Automatic)与逐轮审批(Approval)双调度模式,集成多阶段晋级(Successive-Halving)资源淘汰机制。
- 新增固定多场景评估逻辑(`evaluate.py`),通过固定 seed 和组合运动指令产出与奖励权重无关的客观综合评分。
- 新增 SQLite/WAL 存储管理,持久化保留调参会话、Trial 状态、参数护栏、审计记录与 TensorBoard 标量数据。
- 新增独立的自调参 Web 工作台(`web_platform/tuning.html`,`web_platform/src/tuning/`),包含实时曲线可视化组件(`ScalarChart.tsx`)。
- 训练服务支持将最佳配置保存为不可变预设并在普通训练中复用,支持导出最佳 `policy.onnx`。
- 本地训练面板(`LocalTrainingPanel.tsx`)与训练客户端集成自调参能力探测与快速跳转入口。
## [0.7.3] - 2026-09-01
### 新增
@@ -62,3 +138,23 @@
- 优化响应式工作区、可访问性、首屏加载、纹理兼容性和视口交互。
- 增加碰撞体、坐标系、关节轴、质心和惯量辅助可视化。
### 未发布 — Obstacle DeepSeek 自调参
- 新增Obstacle专属四标量schema与真实reward/导航command应用,保持Flat与护栏CAS语义;导出/浏览器执行配套目标速度。
- 增加固定3seed/1000步、首terminal前捕获的客观导航评价,固定权重和成功/跌倒硬门槛;custom保持权威地图并诚实标注;真实checkpoint/观测统计加载、逐seed进程隔离和失败关闭。
- 自调参界面支持任务选择、专属参数/评估展示、terrain/sensor/custom配置交接及最佳Obstacle策略事务导入;增加mock Agent、打分算例、CAS、argv配置与GPU评估验证。
### 未发布 — 已训练基础策略迁移与持续导航
- 普通训练与自调参增加已验证基础策略选择、checkpoint/SHA和兼容观测展示;管理员本地注册来源,内容寻址只读快照,拒绝客户端任意路径、symlink、unsafe反序列化及失效来源。
- 严格迁移47维actor到47/81/97:新增列零初始化,保留源归一化统计/count;新trial重置critic/optimizer/iteration,同trial晋级只resume自己的checkpoint,导出保留无路径来源metadata。
- 服务重启后的非终态session等待显式恢复,审计原状态并使旧pending proposal失效重提;保留来源、约束和Approval模式,防重复resume worker。
- 移除浏览器Obstacle交互导航20秒截止,保留训练/固定三seed评估协议及跌倒/越界安全停止;设点不自动启用策略,明确显示暂停/未启用/错误状态。
- 新增来源安全、SHA快照、trial续训/重启、双面板和真实WASM/ORT持续换目标回归;4env×1iteration仅证明初始化与优化,空旷地图实测不代表复杂障碍收敛。
### 未发布 — 基础策略目录刷新安全修复
- 修复普通训练在刷新服务目录后清掉失效来源、静默降级随机初始化的问题;保留已选内容ID,不按同别名自动替换新内容。
- 两种训练面板统一对目录缺项、验证失败及任务不兼容来源显示失效提示,禁用启动并在handler再次拦截;必须用户明确从头训练或选择有效来源。显式切换任务仍清空选择。
- 增加两面板A→B、not-ready、缺项及兼容性变更的真实组件刷新回归,验证失效时零启动请求及用户明确选择后的请求体。
+2 -2
View File
@@ -1,12 +1,12 @@
{
"name": "mujoco-web-platform",
"version": "0.8.3",
"version": "0.9.1",
"lockfileVersion": 3,
"requires": true,
"packages": {
"": {
"name": "mujoco-web-platform",
"version": "0.8.3",
"version": "0.9.1",
"license": "Apache-2.0",
"dependencies": {
"@monaco-editor/react": "^4.7.0",
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "mujoco-web-platform",
"version": "0.8.3",
"version": "0.9.1",
"description": "基于 MuJoCo WebAssembly 的本地机器人仿真与控制平台",
"private": true,
"type": "module",
+188
View File
@@ -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必须标准。出生z固定.32,默认XY/四元数取成功加载时真实初始浮动基座姿态,不取跑动后的pose;目标建议为初始朝向前方3m,不自动寻找安全点。面板提供最小出生/目标XY配置,修改后需重新同步。
- 出生及目标中心必须距边界>=.5m;与每个非底板障碍XY AABB的最近距离必须严格>.5m(圆形安全区,包括切触拒绝),不能删除或移动障碍以修复。四元数必须有限且归一化。后端env origin、初始化姿态、目标偏移和越界中心按自定义spawn计算。
- 同步/提交均拒绝未应用草稿、训练部署替换后的场景、加载中或已过时场景;提交前再次比较当前编译布局。HTTP `MAX_REQUEST_BYTES` 与训练JSON上限均128KiB,足够257个double-precision boxes且有界;ONNX部署metadata上限仍100KB。
## 81维观测/12维动作
Actor顺序(float32):
| slice | 内容 |
| ----- | --------------------------------------------------------- |
| 0:3 | `robot/imu_ang_vel`,本体坐标角速度rad/s |
| 3:6 | 本体坐标重力单位向量(静止 `[0,0,-1]`) |
| 6:9 | 下述目标导航速度命令 `[vx,0,wz]` |
| 9:11 | sin/cos(2π×episodeTime/0.6);命令范数<0.1时均为0 |
| 11:23 | 相对默认关节位置 |
| 23:35 | 关节速度,无缩放 |
| 35:47 | 上一原始12维动作,不是力矩或位置目标 |
| 47:79 | 32条前向测距,miss→1,否则clamp(distance/maxDistance,0,1) |
| 79:81 | `[headingError/π, clamp(goalDistance/size,0,1)]` |
关节顺序FL/FR/RL/RR,每腿hip/thigh/calf;默认角 `[-.1,.9,-1.8, .1,.9,-1.8, -.1,.9,-1.8, .1,.9,-1.8]`。
位置目标 `defaultJointPosition+0.25*rawAction`;kp每腿 `[20,20,40]`,kd `[1,1,2]`,力矩限幅 `[23.5,23.5,45]`。后端物理dt=.005,decimation=4,策略50Hz;浏览器仍可更小物理dt并采用single-in-flight/held-action。Go2-W轮式模型不是此任务的同构训练机器人,不应声称动力学严格等同。
### 射线
安装在 `base_link`。对 i=0..31:
- `angle=(-fov/2+i*fov/31)*π/180`,局部direction=`[cos(angle),sin(angle),0]`。
- 局部offset=`[.3,0,.05]`,origin=`baseWorldPosition+Rbase*offset`,directionWorld=`Rbase*direction`。采用**完整body姿态**,不可只用yaw。采样从右到左且包括两端;无恰好0°的中间射线。
- 后端地形group0、机器人visual group2/collision group3,include_geom_groups=(0,)排除整个机器人。浏览器必须语义等价地仅射向地形(如其group2),不能照搬group0或只排除base_link。包含地面;倾斜时地面命中也构成感知数据。miss或超过range归一化为1。
- 前向扇形只能在一个高度切片感知;矮障碍物/跌落/盲区并不保证可检测,不能称安全导航保证。
### 导航及episode
`delta=target.xy-basePosition.xy`;distance=hypot(delta);yaw由base wxyz计算;headingError=`atan2(sin(atan2(dy,dx)-yaw),cos(...))`。未到达时 `vx=navigation.speed*max(0,cos(headingError))`(缺省.6,调参可.3–1.2)、vy=0、wz=clamp(headingError,-1,1)。distance<.5m到达后command全0、headingError=0,distance observation仍保留实际距离;不重采样新目标,不因到达施加失败惩罚。移动机器人离开半径后自然恢复导航。
训练episode20s,超时重置;跌倒、非足端接触force>10N、机器人中心越过地图边界内0.3m终止并重置。训练复位随机选择同一连通自由区域的起终点和初始yaw,恢复初始关节姿态及零速度,并清phase/上一动作;固定评估仍使用deployment声明的参考起终点以保持版本间可比。浏览器episode若不自动重置必须明确自己的评测行为;到达至少要停命令而非无限推进。奖励保留速度跟踪/平滑/姿态/非法接触终止惩罚,并增加朝目标世界速度投影及归一化近障平方惩罚。
## 验证
```bash
.venv/bin/python -m unittest discover -s training_server/tests
.venv/bin/python -m ruff check training_server
GO2_RUN_MJLAB_SMOKE=1 MUJOCO_GL=egl .venv/bin/python -m unittest discover -s training_server/tests -p test_obstacle_env.py
WANDB_MODE=disabled .venv/bin/python training_server/rl/scripts/train.py Unitree-Go2-ObstacleAvoidance --env.scene.num-envs=4 --agent.max-iterations=1 --gpu-ids '[0]' --output-dir /tmp/go2-obstacle-train-smoke
```
GPU smoke对比真实mjlab/Warp与MuJoCo `mj_ray`距离、检查多env出生点/关节顺序/81维有限obs。1iteration只验训练/导出管线,**不证明策略已学会避障或Sim2Sim行为收敛**。
## Obstacle DeepSeek 自调参(obstacle-v1)
`POST /api/tuning/sessions` 新支持 `taskId=Unitree-Go2-ObstacleAvoidance`。
沿用 automatic/approval、运行时双向切换、稀疏patch、最多4项/单轮0.5–2倍、工程护栏与revision CAS;Flat完整schema、奖励应用与相对基线六维评分不变。
可携带与训练相同的顶层terrain/sensor字段,或单独`taskConfig`对象(不能混用);场景/seed仅创建时确定。所有字段/范围、NaN/Inf、跨任务patch/护栏拒绝;旧revision返回409且不修改配置/审计/待审批项。
| JSON路径 | 范围;默认 | 真实环境映射 |
| --------------------------- | -------------------------------- | -------------------------------------------------------------------------- |
| `weights.avoidance_weight` | .5–5;2(创建时可取sensorCfg值) | `obstacle_proximity.weight=-value` |
| `params.target_velocity` | .3–1.2;.6 m/s | `NavigationCommandCfg.speed`,**不是奖励权重** |
| `weights.collision_penalty` | -10–-.5;-5 | `obstacle_collision`:同illegal_contact函数/参数,非足端地形接触>10N指示量 |
| `weights.action_smoothness` | -.05–-.001;-.05 | `action_rate_l2.weight` |
保留原is_terminated等非白名单项,collision_penalty额外独立惩罚非法接触,不用它替代客观碰撞统计。服务写每trial独立reward/task JSON,以`--reward-config/--task-config` argv传入真实训练器和评估器;训练入口先应用场景再应用调参。导出ONNX/deployment的navigation.speed与sensorCfg.avoidanceWeight反映实际配置,浏览器仍81→12且使用配套速度。旧无调参模型保持.6。奖励Preset在普通训练面板仍仅Flat可选;Obstacle可从调参工作台直接导入最佳带metadata策略,不把Obstacle preset混入Flat训练。
### 固定客观评估与安全门槛
- 固定seed **101/202/303**、每seed **1000策略步×.02s=20s**;每env仅统计reset后的**第一个episode**。提前终止后仍运行完整horizon,但不将自动reset后新episode混入首episode;剩余步数不能带来平滑/净距奖励。各seed先按env等权均值,再三seed等权。evalNumEnvs创建时固定(所有trial相同),Agent不能改seed、horizon、权重、场景、传感器、起终点。
- 固定评估按三个seed分别构造布局,非随机预设不保证三张不同地图;评估关闭训练期随机起终点。custom严格保持上传的同一地图及验证过的参考spawn/target,标记`fixed-custom-map`,**不是三个随机地形**。保留任务原有reset_joint/startup随机化并用seed复现;若轨迹相同仅是确定性重复试验,不能当作泛化证据。
- 读取真实checkpoint Actor网络及其`obs_normalizer`统计,strict加载;没有零动作/假策略fallback,不读取训练reward作为指标。不同地图CUDA编译隔离:三个seed顺序启动同一Python解释器子进程(非fork/非并行),共享不可变checkpoint快照并核对SHA256;父进程使用独立临时目录,worker继承父进程组,现有session取消会终止worker,每seed超时1200s即整轮失败。任意seed失败、缺失、错身份、样本不足、非有限或协议不一致整轮fail closed,不用部分结果凑均值。
- hook在termination manager计算后、`_reset_idx`之前捕获包括**首terminal**的物理量。与mjlab终止判断一致,派生pose最多滞后一个.005s物理子步;接触force history覆盖全部4子步。每步读取目标水平距离、前视射线命中比例、非底板障碍净距、实际action差分、接触/跌倒标志。前视ray_hit可包含地板,仅作为感知诊断,不直接评分。
- **擦碰/碰撞**:任一机器人(含足端)与非底板障碍的实际接触力幅值>1N,或非足端与任意地形(含地面)>10N。使用编译terrain_1…几何的独立contact sensor(maxforce+4子步history),排除terrain_0标准底板;不是用ray命中推测碰撞。1N以下触碰不计,不能称零接触安全保证。
- **跌倒**:base高度<.12m或姿态倾角>70°(projected_gravity.z > -cos70°)。**到达**:水平距离<.5m。**成功**:episode内到达,且整段首episode无碰撞、无跌倒;先到达后擦碰/跌倒也失败。未成功时间项必为0,不因提前跌倒取短时间高分。
固定五项0–1分量:`success`为成功指示;`time`成功时`1-firstArrivalStep/1000`,否则0;`clearance`每步`clamp((baseXY到非底板box XY AABB距离-.3m)/.5m,0,1)`,按1000步求均值(终止后补0,任何跌倒整项0,无障碍时每有效步1);`smooth`每步`1-clamp(mean((action_t-action_t-1)^2),0,1)`,初始action=0,按1000步求均值(终止后补0);`no_fall`为无跌倒指示。净距是保守圆形机身足迹代理,并非全机身mesh最近距离,地板不作为障碍、躺地不当安全。
绝对总分=`.4*success+.2*time+.2*clearance+.1*smooth+.1*no_fall`。例如两步简化fixture,第二步无碰撞到达,净距均.25m、action差分MSE均.5,得分`.4+0+.1+.05+.1=.65`。固定1000步实际协议不接受fixture horizon。
相对初始baseline的硬门槛:`fall_rate<=baseline+.02`且`success>=baseline-.02`(边界包含;1e-12浮点容差)。不通过eligible=false且score=-1,不能用其他高分绕过。baseline也用完整客观绝对评分,不默认指标为0。result记录完整协议/场景布局、seed结果、样本数、checkpoint SHA256、各项均值,candidate必须与baseline协议完全一致。
## 显式多层射线(阶段5;不是局部高程图)
`sensorCfg.sensorMode` 是唯一模式字段:缺省/`single_ring_raycast` 保持旧32ray/81obs及旧proximity奖励行为;显式 `multi_ring_raycast` 为3×16=48ray/97obs。`sensorType/type`仍为`raycast`。health新增`sensorModes`,UI可选择;不是任意自定义扫描图案。没有局部高程图网格、用途或输入定义,本阶段**未实现高程图**。
任务JSON、deployment JSON及ONNX metadata规范化导出以下字段(FOV仍叫`fov`):
```json
{
"sensorMode": "multi_ring_raycast",
"rayCount": 48,
"pitchAngles": [0, -20, -45],
"yawCount": 16,
"yawAngles": [-45, -39, -33, -27, -21, -15, -9, -3, 3, 9, 15, 21, 27, 33, 39, 45],
"fov": 90,
"angleUnit": "deg",
"rayOrder": "layer-major"
}
```
旧single规范化为pitch `[0]`、yawCount/rayCount=32。所有yaw为`-fov/2+i*fov/(yawCount-1)`含两端,从右到左,无0度中心ray。只允许两种固定组合;显式矛盾字段/未知字段/不支持的模式拒绝,不静默覆盖。旧metadata缺新增字段可规范化后与新job语义匹配;旧81 graph绝不冒充97。
局部direction=`[cos(pitch)*cos(yaw),cos(pitch)*sin(yaw),sin(pitch)]`,先一层16yaw再下一pitch层;origin仍为`basePosition+Rbase*[.3,0,.05]`,direction乘**完整base姿态**。后端实际RayCastSensorCfg使用同一pattern;浏览器预计算局部direction并缓存每个已编译静态box世界变换,解析slab求交;只对精确单位旋转用轴对齐快速路径,无变换近似。
97维slice:基础`0:47`完全不变,真实测距`47:95`,目标误差`95:97`仍为原有符号的heading/π与距离归一化。miss/超range→1,地板命中保留真实距离;调参`target_velocity`继续传入真实command与导出navigation.speed。原single-in-flight、held-action、事务导入和资源释放协议不变。三seed评估仍是顺序独立解释器;advisor按实际32/48模式说明盲区。
### 地板与奖励
仅multi的`obstacle_proximity`排除标准底板顶面:首次计算用CPU编译模型验证**唯一**与首box相符的静态world-weld/group0 box(世界中心`[0,0,-.1]`、半尺寸`[size/2,size/2,.1]`、单位世界旋转),再核实生成器名称`terrain_0`,失败关闭;不是只根据名称。真实编译terrain为固定body而非bodyid=0,使用weld身份。
当前mjlab RayCastData没有hit geom ID,因此是**几何分类过滤**:必须finite实际命中,hit世界XY在floor范围(容差1e-5m),`abs(hit.z)<=1e-5m`,normal与+Z逐分量误差<=1e-5。每个Warp环境是独立同坐标地图,hit不额外加env_origin。容差为float32的10微米,不把5cm低障碍抹除。低box顶面、侧面不滤;全floor/miss惩罚0且finite,其余仍用最近有效距离的原平方归一化。观测数据不原地修改。与底板顶面共面的几何无法凭此分类区分,不能声称等价命中geom ID;台阶/非零高度表面不是标准地板。
下倾层改善部分低障碍/有界地板边缘感知,但层间、侧后方、遮挡和有限range仍有盲区;标准boxes-v1底板本身会填补内部空洞,不能声称已解决坑探测或安全导航。
### 阶段5复现
```bash
.venv/bin/python training_server/tests/generate_multi_ring_golden.py
GO2_RUN_MULTI_SMOKE=1 MUJOCO_GL=egl .venv/bin/python -m unittest discover -s training_server/tests -p test_multi_ring.py
WANDB_MODE=disabled MUJOCO_GL=egl .venv/bin/python training_server/rl/scripts/train.py Unitree-Go2-ObstacleAvoidance --env.scene.num-envs=4 --agent.max-iterations=1 --gpu-ids '[0]' --task-config /tmp/go2-multi-ring-stage5/task.json --output-dir /tmp/go2-multi-ring-stage5/train
```
真实4env/1iteration新97策略仅验训练、归一化Actor导出和浏览器ORT链路,不是旧81 checkpoint改metadata,不证明导航成功。CLI子进程`--help`回归分别验证Flat/Rough/Obstacle的任务注册;修复前直接CLI仅注册mjlab内建任务,原错误日志保留。
事务重导入回归确认97→97、97→81→97合法graph/metadata成功;伪97metadata+81graph与缺失ORT算子均在预期阶段报具体错误并保留旧会话。曾看到“模型编译失败”是App只转抛diagnostic.summary遮蔽了原detail,并非实际MJCF编译失败;已改为detail优先,未改事务资源/rollback。
### 复审修复:奖励preset任务归属
公共preset仍可保存Obstacle调参结果,但`save/get/list`均从来源session的`config.taskId`恢复权威任务归属并验证完整reward schema;列表和单项返回`taskId`。历史session仅缺`taskId`时按既有Flat默认解释,且必须通过完整Flat schema。来源缺失/损坏、显式未知或空taskId、损坏preset均失败关闭(列表不返回部分可信结果);不相信客户端声明,也不凭奖励字段猜任务。
普通本地训练的`rewardPresetId`入口仍仅支持Flat;resolver核对来源任务与请求任务,服务再次验证完整schema,跨任务/损坏preset返回400且不创建job。UI仅展示服务明确标记Flat的preset,旧服务无身份字段时不展示,而不是默认为Flat。Obstacle preset UI未扩展;Obstacle调参最佳策略配套导入、Approval/Automatic均保持原协议。
+178
View File
@@ -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产物。
+20 -4
View File
@@ -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不匹配或快照失效均明确拒绝,绝不静默随机初始化。
+502
View File
@@ -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)
+447
View File
@@ -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"],
]
+336
View File
@@ -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)
+4
View File
@@ -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
+36 -26
View File
@@ -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))
+117 -10
View File
@@ -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
View File
@@ -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
+479
View File
@@ -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
]
+14
View File
@@ -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()
+254
View File
@@ -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()
+214
View File
@@ -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()
+223
View File
@@ -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()
+475
View File
@@ -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()
+72 -2
View File
@@ -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")
+83
View File
@@ -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()
+5 -2
View File
@@ -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()
+257 -23
View File
@@ -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:
+130
View File
@@ -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,
}
+75 -32
View File
@@ -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}")
+56 -29
View File
@@ -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]
+80
View File
@@ -0,0 +1,80 @@
# 自定义训练与前视射线避障
1. 启动仓库本地训练服务,在「控制台 → 强化学习任务」输入令牌并连接。
2. 选择「前视射线避障导航」,设置地形、种子、障碍物参数、FOV、探测/安全距离及避障奖励权重,发起训练。界面显示进度、最近价值/策略/熵损失和原始日志。
3. 「同步当前场景地图」从全部**已应用**实例的已编译MuJoCo静态碰撞几何导出权威 `custom_boxes/boxes-v1`,支持多实例、平移/旋转、工程地图路径及嵌套body。成功显示“已将视口中 N 个自定义障碍物编译为训练地图布局”。不重新运行随机预设,不同步草稿/过时场景。初始姿态取场景加载时机器人位姿,出生高度标准化为.32m;目标建议为初始朝向前方3m,不自动寻找安全点。出生/目标XY可修改,必须重新同步,并保留与所有障碍XY AABB不相交的.5m圆形安全区;绝不删除或移动障碍来避开安全区。
4. 训练完成后,先导入并加载有对应12个关节的Go2模型(Go2-W仅为非同构实验)。点击「导入策略」会校验作业配置与ONNX内嵌部署契约,在候选场景中完成真实ORT初始化、graph与机器人绑定检查后才事务替换当前**物理**地形、复位到出生点,启动策略与仿真。当前编辑器地图保留,重新加载模型恢复它。失败保留原场景、策略、物理状态及原暂停/播放状态。旧Flat无配套地形时仍按原方式加载后手动启用;新服务的默认Flat作业允许旧外部训练器无metadata导出,但真实graph必须固定47→12;显式自定义地图/避障不允许此降级。
5. 避障任务自动导航到配套目标,目标半径0.5m内停止速度指令。摄像头PiP自动打开,可单独隐藏;「显示避障射线」可切换红色命中/绿色未命中线段。
## 交互式导航目标与训练趋势
- 避障策略加载后,主视口显示青色呼吸目标信标,按当前训练碰撞地形的顶部高度落位。ONNX 面板显示当前目标坐标及水平剩余距离。
- 点击「设定目标」,再在主视口地形单击:目标限制在地图边缘内缩0.5m的区域,并自动退出设定模式。Esc/取消可退出;空白、拖动及摄像头PiP不会更新目标。模式使用捕获阶段拦截输入,不选择刚体或操纵地图Gizmo。
- 暂停或停用策略时也可更改目标,但不会自动启动仿真或策略。「复位目标点」只恢复训练地图默认目标;「重置仿真」同时恢复默认目标与出生状态。换目标不修改部署JSON、不重建ONNX会话,下一次观测仍保持该策略的81或97维。信标仅表示目标位置,不保证该点可到达或路径无障碍。
- 训练卡片中的「训练指标趋势」可折叠,提供价值/策略损失及综合页;综合页的熵、平均奖励、平均回合长度各用独立纵轴。复用自调参工作台的轻量Canvas图表,EMA系数0.4仅用于曲线,悬停图例和摘要保留原始值,支持缩放。
- 按日志中的 `Learning iteration` 合并指标,重复轮询不重复入点,同迭代分批日志可补全;内存最多保留最近500个有指标的迭代。换作业/清除作业会清空。服务器仅保留日志尾部,断线重连不能恢复已丢失历史;没有迭代标题或已知重叠上下文的孤立指标会跳过,不猜测step。缺失指标不显示,原始日志仍可展开。
## 自定义场景导出边界
- 使用实际碰撞几何而非Three装饰包围盒;box的世界AABB半尺寸为abs(R)×halfsize,sphere/capsule/ellipsoid/cylinder有精确解析bounds。mesh/hfield暂明确拒绝,其他未声明静态碰撞也报错,不能漏掉。既有工程地图include导入限制不变。
- **AABB会膨胀旋转/非box形状,底板标准化可能填补支撑面之间空白**;custom总是approximation=true,仅保证训练和浏览器部署相同boxes,不保证原OBB/原场景几何同构。仅明确命名且顶面z≈0的水平支撑归并为floor z=[-.2,0];地下/坑底、倾斜或非零高度plane拒绝,不静默填坑。
- 固定floor加最多256障碍,超量报错不截断。保持世界坐标、覆盖范围8–24m,超出时明确拒绝,不clamp/平移障碍。所有半尺寸须严格正值;统一friction=.2–2且其余摩擦分量为.005/.0001,混合摩擦需先在地图中显式统一。
- 服务和训练入口都验证完整JSON白名单、尺寸/位置、标准floor、count/approximation一致性及起终点安全区;`terrainPreset=custom_boxes`配套`customTerrainBoxes`布局全程传入任务JSON/deployment/ONNX,不降级预设。旧服务无custom_boxes能力时同步按钮报升级提示。
## 明确边界
- 默认32条水平前向物理射线;显式multi模式为3×16=48条(pitch 0/-20/-45度),**不是真实camera_depth或RGB视觉策略**。PiP是观察窗口,不作为网络输入。低障碍、跌落与切片外障碍存在感知盲区。
- 配套`boxes-v1`布局直接来自训练器,使用相同世界坐标、半尺寸、摩擦、出生点和目标;`rough/wave/pyramid_stairs`明确为训练专用box近似,不等于编辑器同名高度场。
- 浏览器对**编译后的**静态训练box `geom_xpos/geom_xmat/geom_size`做解析slab求交,包括地板。它与MuJoCo box几何求交等价;完整机身姿态旋转射线,按命名与group2仅选训练地形,排除整个机器人。不接受额外无限plane、未声明静态几何或非box训练地形。地图组合先用`mj_saveLastXML`展开include,避免原地面残留。
- 使用原47维本体观测+32归一化距离+2目标误差,共81维;multi改为47+48+2=97维。保持单in-flight、50Hz异步held-action,不阻塞WASM子步。不另加动态观测归一化。导出器Actor内已有归一化。
- 原Go2自定义地图部署按`go2_constants.py`补齐hip/thigh/calf armature=0.01/0.01/0.02、零关节阻尼/摩擦损失、非足端condim1/priority0、足端condim3/priority1/solimp宽度0.023,碰撞contype1/conaffinity0关闭自碰撞;足端摩擦使用配套地形。旧无自定义地图的Flat路径不覆盖这些参数。
- 初始关节姿态写入`MjData.qpos`,不改变模型`qpos0`的关节参考角;重置会恢复出生位姿、初始关节、零速度/历史动作及步态时间。
- 浏览器评测在20秒、机身高度低于0.12m或越界时停止并报错,需重置后重新启用;**不是训练端的接触/姿态终止自动reset**。到达目标不会强制结束episode。
- 旧`Unitree-Go2-Rough`的234维观测未在浏览器实现,可训练但禁止一键导入。原Flat、Go2-W速度策略、Python控制器、Flat自调参工作台保持兼容。
- 模型拓扑及PD配置校验不等于动力学完全一致。短PPO和零动作fixture只能验证链路,不能证明避障收敛或Sim2Sim导航成功。
## 验证入口
```bash
npm run typecheck
npm run check
npm run test:e2e -- web_platform/e2e/obstacle.spec.ts
# 可选:加载后端1iteration实际导出的模型(无需复制模型进仓库)
GO2_SMOKE_POLICY=/tmp/go2-obstacle-train-smoke/policy.onnx npm run test:e2e -- web_platform/e2e/obstacle.spec.ts
```
`src/rl/fixtures/obstacleRayGolden.json`保存CPU MuJoCo `mj_ray`参考距离;单测覆盖yaw/pitch/roll、地板命中、box内部/表面/平行/边缘、range边界、miss;真实WASM单测编译Go2碰撞/惯量模型并逐元素验证32条射线。E2E默认模型为`fixtures/obstacle/zero-action.onnx`(测试专用Constant零动作,非训练策略),检查真实浏览器WASM+ORT、地图联动、PiP、射线开关、metadata不一致拒绝。
完整后端字段定义见[部署契约](../training_server/OBSTACLE_AVOIDANCE.md)。
## Obstacle 自调参
自调参工作台新建Session可选Flat或Obstacle,选择Obstacle自动绑定四个专属标量/范围和五项固定客观权重(success .4/time .2/clearance .2/smooth .1/no-fall .1)。目标速度位于`params.target_velocity`,是导航command而非奖励。Approval/Automatic及运行时切换、护栏revision保持兼容,Loss看板不变。
从训练面板打开工作台时传递当前任务、seed、terrain/sensor(包括已验证custom_boxes快照);未同步/过时场景不发送凭据或配置。独立打开时可直接选任务,在“避障场景配置JSON”输入与训练API同构的terrain/sensor配置;服务再次严格验证。自定义地图明确为同一权威地图/起终点上的三个固定seed重复评估,不冒充三地形。Agent上下文按32/48ray实际模式描述侧后、层间与高度盲区、目标导航、擦碰/绕行/动作抖动约束。
Obstacle最佳策略通过ONNX内嵌metadata事务导入,导航速度随部署契约应用(.3–1.2),旧81维.6m/s模型与Flat导入保持原行为。客观评分、成功/碰撞/跌倒定义及3seed串行隔离评估详见[部署契约](../training_server/OBSTACLE_AVOIDANCE.md#obstacle-deepseek-自调参obstacle-v1)。小GPU smoke只验证链路和指标,不代表学会导航。
## 多层射线与性能(阶段5)
避障高级设置的“传感器模式”可选默认水平32ray/81obs或显式三层48ray/97obs。旧服务未公布multi能力时选项禁用;加载策略以后按metadata动态分配PiP射线,不截断到32条。前47维及heading/π符号不变,Click-to-Navigate目标/面板、Loss曲线、custom_boxes、Obstacle调参速度与Flat保留。旧81 graph不能通过97部署shape校验。
**不是局部高程图**:尚无网格/用途定义,本阶段未实现。下倾仍可能漏掉层间矮物、遮挡及坑;标准地图底板会填内部空洞。floor真实命中进入观测,只有multi近障奖励使用已验证标准底板顶面的10微米几何分类过滤,不把floor测距伪装成miss。与floor顶面共面的其他几何无法区分。完整字段、公式和分类限制见[多层契约](../training_server/OBSTACLE_AVOIDANCE.md#显式多层射线阶段5不是局部高程图)。
`multiRingGolden.json`来自CPU mj_ray,覆盖全部48条完整姿态、5cm低障碍、有界地板跌落边缘和层序;真实WASM绑定逐元素比对。`multi-zero-action.onnx`为测试Constant零动作97维fixture,非训练策略。真实短训模型可用:
```bash
GO2_MULTI_SMOKE_POLICY=/tmp/go2-multi-ring-stage5/train/policy.onnx npm run test:e2e -- web_platform/e2e/multiRing.spec.ts
```
`src/rl/tasks/raycastBenchmark.ts`导出`benchmarkForwardRays()`:每帧真实调用生产sampleForwardRays与slab,预计算方向/静态box,包含完整48ray全部box求交、姿态旋转和depth/线段分配;不测单ray、不用平均值代替尾延迟。默认25box和257box的16×16最坏数量布局分别warmup3000帧、采样10000帧,逐帧改变姿态/位置,不缓存命中结果。仅报告数量最坏布局,不声称穷尽所有空间排布或整浏览器帧耗时。
本机Intel i7-14700F/Ubuntu24.04,Node24.19:25box p50/p95约.019/.020ms,257box约.169/.187ms;HeadlessChromium151(同机,cross-origin-isolated高精度计时)约.015/.020ms与.160/.165ms,测得均低于完整48ray .2ms目标。初版通用旋转slab的257box p95≈.219ms未达标;优化为严格单位旋转快速slab后达标,保留原日志。不隔离浏览器计时分辨率约.1ms,p95量化为.2ms,不能用该读数证明严格小于目标;高精度基准不改变生产隔离/时序配置。
结果随硬件、GC、浏览器调度变化,不是实时保证;只测解析感知,不含渲染/ORT/物理步。复现打包脚本、Node/浏览器原始分位数与硬件说明位于`/tmp/go2-multi-ring-stage5/run_benchmark.mjs`及`perf-*.json`;日志与独立阶段incremental.diff同目录,后续review可直接读取。
## 复审修复:preset与导航拖动
Flat奖励菜单只展示服务根据来源session确认`taskId=Unitree-Go2-Flat`的preset;Obstacle及未声明身份项不展示。历史Flat的任务身份由服务恢复并完整验证schema,跨任务/损坏preset在创建作业前返回400,不推迟到训练器失败。旧服务没有返回任务身份时需升级服务才能使用这些preset;不增加Obstacle普通训练preset入口。
设定导航目标时超过5px即永久标记本次手势为拖动,即使返回起点也不会设置目标;画布外的同pointer移动也记录。pointercancel、Esc、失焦、退出模式、clear/dispose会清空手势。取消后旧pointerup不生效,重新按下的正常点击仍可设定。
+153
View File
@@ -0,0 +1,153 @@
import { expect, test } from '@playwright/test';
import { readFileSync } from 'node:fs';
import { resolve } from 'node:path';
const fixture = JSON.parse(
readFileSync(resolve('web_platform/src/rl/fixtures/multiRingDeployment.json'), 'utf8'),
);
// Original Go2 collision/inertia model, visual meshes omitted to keep the smoke fixture lightweight.
const go2 = readFileSync(
resolve('training_server/rl/src/assets/robots/unitree_go2/xmls/go2.xml'),
'utf8',
)
.replace(/<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.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
await page.getByRole('tab', { name: '控制台' }).click();
const tools = page.getByRole('tabpanel', { name: '控制台' });
await tools.getByRole('button', { name: /强化学习任务/ }).click();
await tools.getByRole('button', { name: '连接', exact: true }).click();
await expect(tools.getByRole('button', { name: '导入策略' })).toBeEnabled();
await tools.getByRole('button', { name: '导入策略' }).click();
await expect(page.getByText('训练配套物理地图', { exact: false })).toBeVisible({
timeout: 30_000,
});
await expect(page.getByLabel('摄像头画面', { exact: true })).toBeVisible();
const rays = page.getByRole('checkbox', { name: '显示避障射线' });
await expect(rays).toBeChecked();
await rays.uncheck();
await rays.check();
const section = tools.getByRole('button', { name: /ONNX 策略运行/ });
if ((await section.getAttribute('aria-expanded')) === 'false') await section.click();
await expect(tools.getByText('97 / 12')).toBeVisible();
await expect(tools.getByText('Go2 前视射线避障导航', { exact: true })).toBeVisible();
const count = tools.getByText('推理次数', { exact: true }).locator('..');
await expect
.poll(async () => Number((await count.textContent())?.replace(/\D/g, '')))
.toBeGreaterThan(0);
const pause = page.getByRole('button', { name: '⏸ 暂停' });
if (await pause.isVisible()) await pause.click();
const stop = tools.getByRole('button', { name: '停止', exact: true });
if (await stop.isVisible()) await stop.click();
const targetRow = tools.getByText('当前目标 (X, Y)', { exact: true }).locator('..');
const initialTarget = await targetRow.textContent();
const selection = () => page.locator('[aria-selected="true"]').allTextContents();
const beforeSelection = await selection();
await tools.getByRole('button', { name: '设定目标', exact: true }).click();
await page.keyboard.press('Escape');
await expect(tools.getByRole('button', { name: '设定目标', exact: true })).toHaveAttribute(
'aria-pressed',
'false',
);
await expect(targetRow).toHaveText(initialTarget!);
await tools.getByRole('button', { name: '设定目标', exact: true }).click();
const canvas = page.locator('.viewport-shell canvas').first();
const bounds = (await canvas.boundingBox())!;
await canvas.click({ position: { x: bounds.width * 0.5, y: bounds.height * 0.55 } });
await expect(targetRow).not.toHaveText(initialTarget!);
await expect(tools.getByRole('button', { name: '设定目标', exact: true })).toHaveAttribute(
'aria-pressed',
'false',
);
expect(await selection()).toEqual(beforeSelection);
await expect(tools.getByRole('button', { name: '启用', exact: true })).toBeVisible();
await tools.getByRole('button', { name: '复位目标点' }).click();
await expect(targetRow).toHaveText(initialTarget!);
await tools.getByRole('button', { name: '导入策略' }).click();
await expect(tools.getByRole('alert')).toHaveCount(0);
// Switch real valid graph/metadata pairs in both directions, without relaxing transactions.
const multiModel = model;
job.deployment = JSON.parse(
readFileSync(resolve('web_platform/src/rl/fixtures/obstacleDeployment.json'), 'utf8'),
);
model = readFileSync(resolve('web_platform/fixtures/obstacle/zero-action.onnx'));
await tools.getByRole('button', { name: '连接', exact: true }).click();
await tools.getByRole('button', { name: '导入策略' }).click();
await expect(tools.getByText('81 / 12')).toBeVisible();
job.deployment = fixture;
model = multiModel;
await tools.getByRole('button', { name: '连接', exact: true }).click();
await tools.getByRole('button', { name: '导入策略' }).click();
await expect(tools.getByText('97 / 12')).toBeVisible();
const pauseAgain = page.getByRole('button', { name: '⏸ 暂停' });
if (await pauseAgain.isVisible()) await pauseAgain.click();
const stopAgain = tools.getByRole('button', { name: '停止', exact: true });
if (await stopAgain.isVisible()) await stopAgain.click();
// A real 81-feature ONNX graph with valid 97 metadata must fail transactionally.
model = readFileSync(resolve('web_platform/fixtures/obstacle/multi-wrong-graph.onnx'));
await tools.getByRole('button', { name: '导入策略' }).click();
await expect(tools.getByRole('alert')).toContainText('维度');
await expect(tools.getByText('97 / 12')).toBeVisible();
await expect(targetRow).toHaveText(initialTarget!);
await tools.getByRole('button', { name: '设定目标', exact: true }).click();
await tools.getByRole('button', { name: '卸载', exact: true }).click();
await expect(targetRow).toHaveCount(0);
await expect(tools.locator('.uplot')).toHaveCount(0);
await tools.getByRole('button', { name: /训练指标趋势/ }).click();
await expect(tools.locator('.uplot canvas')).toHaveCount(1);
await tools.getByRole('tab', { name: '综合', exact: true }).click();
await expect(tools.locator('.uplot canvas')).toHaveCount(5);
await tools.getByRole('button', { name: /训练指标趋势/ }).click();
await expect(tools.locator('.uplot')).toHaveCount(0);
});
+291
View File
@@ -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.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
await page.getByRole('tab', { name: '控制台' }).click();
const tools = page.getByRole('tabpanel', { name: '控制台' });
await tools.getByRole('button', { name: /强化学习任务/ }).click();
await tools.getByRole('button', { name: '连接', exact: true }).click();
await expect(tools.getByRole('button', { name: '导入策略' })).toBeEnabled();
await tools.getByRole('button', { name: '导入策略' }).click();
await expect(page.getByText('训练配套物理地图', { exact: false })).toBeVisible({
timeout: 30_000,
});
await expect(page.getByLabel('摄像头画面', { exact: true })).toBeVisible();
const rays = page.getByRole('checkbox', { name: '显示避障射线' });
await expect(rays).toBeChecked();
await rays.uncheck();
await rays.check();
const section = tools.getByRole('button', { name: /ONNX 策略运行/ });
if ((await section.getAttribute('aria-expanded')) === 'false') await section.click();
await expect(tools.getByText('81 / 12')).toBeVisible();
await expect(tools.getByText('Go2 前视射线避障导航', { exact: true })).toBeVisible();
const count = tools.getByText('推理次数', { exact: true }).locator('..');
await expect
.poll(async () => Number((await count.textContent())?.replace(/\D/g, '')))
.toBeGreaterThan(0);
const pause = page.getByRole('button', { name: '⏸ 暂停' });
if (await pause.isVisible()) await pause.click();
const stop = tools.getByRole('button', { name: '停止', exact: true });
if (await stop.isVisible()) await stop.click();
const targetRow = tools.getByText('当前目标 (X, Y)', { exact: true }).locator('..');
const initialTarget = await targetRow.textContent();
const selection = () => page.locator('[aria-selected="true"]').allTextContents();
const beforeSelection = await selection();
await tools.getByRole('button', { name: '设定目标', exact: true }).click();
await page.keyboard.press('Escape');
await expect(tools.getByRole('button', { name: '设定目标', exact: true })).toHaveAttribute(
'aria-pressed',
'false',
);
await expect(targetRow).toHaveText(initialTarget!);
await tools.getByRole('button', { name: '设定目标', exact: true }).click();
const canvas = page.locator('.viewport-shell canvas').first();
const bounds = (await canvas.boundingBox())!;
await canvas.click({ position: { x: bounds.width * 0.5, y: bounds.height * 0.55 } });
await expect(targetRow).not.toHaveText(initialTarget!);
await expect(tools.getByRole('button', { name: '设定目标', exact: true })).toHaveAttribute(
'aria-pressed',
'false',
);
expect(await selection()).toEqual(beforeSelection);
await expect(tools.getByRole('button', { name: '启用', exact: true })).toBeVisible();
await tools.getByRole('button', { name: '复位目标点' }).click();
await expect(targetRow).toHaveText(initialTarget!);
await tools.getByRole('button', { name: '设定目标', exact: true }).click();
await tools.getByRole('button', { name: '卸载', exact: true }).click();
await expect(targetRow).toHaveCount(0);
await expect(tools.locator('.uplot')).toHaveCount(0);
await tools.getByRole('button', { name: /训练指标趋势/ }).click();
await expect(tools.locator('.uplot canvas')).toHaveCount(1);
await tools.getByRole('tab', { name: '综合', exact: true }).click();
await expect(tools.locator('.uplot canvas')).toHaveCount(5);
await tools.getByRole('button', { name: /训练指标趋势/ }).click();
await expect(tools.locator('.uplot')).toHaveCount(0);
});
test('下载metadata与作业不匹配时拒绝,未更换地图或启用策略', async ({ page }) => {
await page.addInitScript(() => sessionStorage.setItem('mujoco-local-training-token', 'test'));
const job = {
id: 'b'.repeat(32),
taskId: fixture.taskId,
state: 'succeeded',
artifactReady: true,
progress: 1,
iteration: 1,
maxIterations: 1,
logs: [],
message: '',
deployment: { ...fixture, sensorCfg: { ...fixture.sensorCfg, fov: 60 } },
};
await page.route('http://127.0.0.1:8765/**', (route) => {
const url = route.request().url();
if (url.endsWith('/policy.onnx'))
return route.fulfill({ contentType: 'application/octet-stream', body: model });
if (url.endsWith('/health'))
return route.fulfill({ json: { ready: true, tasks: [job.taskId], activeJobId: job.id } });
if (url.includes('/presets')) return route.fulfill({ json: { presets: [] } });
return route.fulfill({ json: job });
});
await page.goto('/');
await page
.locator('input[type="file"]')
.first()
.setInputFiles({ name: 'go2.xml', mimeType: 'text/xml', buffer: Buffer.from(go2) });
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
await page.getByRole('tab', { name: '控制台' }).click();
const tools = page.getByRole('tabpanel', { name: '控制台' });
await tools.getByRole('button', { name: /强化学习任务/ }).click();
await tools.getByRole('button', { name: '连接', exact: true }).click();
await tools.getByRole('button', { name: '导入策略' }).click();
await expect(page.getByText('下载的策略与训练作业部署配置不一致').first()).toBeVisible();
await expect(page.getByText('训练配套物理地图', { exact: false })).toHaveCount(0);
await expect(page.getByLabel('摄像头画面', { exact: true })).toHaveCount(0);
});
for (const failure of ['wrong-graph', 'ort-init-failure']) {
test(`事务导入${failure}时保留旧地图、策略和暂停状态`, async ({ page }) => {
await page.addInitScript(() => sessionStorage.setItem('mujoco-local-training-token', 'test'));
let bytes = model;
const job = {
id: 'c'.repeat(32),
taskId: fixture.taskId,
state: 'succeeded',
artifactReady: true,
progress: 1,
iteration: 1,
maxIterations: 1,
logs: [],
message: '',
deployment: fixture,
};
await page.route('http://127.0.0.1:8765/**', (route) => {
const url = route.request().url();
if (url.endsWith('/policy.onnx'))
return route.fulfill({ contentType: 'application/octet-stream', body: bytes });
if (url.endsWith('/health'))
return route.fulfill({ json: { ready: true, tasks: [job.taskId], activeJobId: job.id } });
if (url.includes('/presets')) return route.fulfill({ json: { presets: [] } });
return route.fulfill({ json: job });
});
await page.goto('/');
await page
.locator('input[type="file"]')
.first()
.setInputFiles({ name: 'go2.xml', mimeType: 'text/xml', buffer: Buffer.from(go2) });
await expect(page.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
await page.getByRole('tab', { name: '控制台' }).click();
const tools = page.getByRole('tabpanel', { name: '控制台' });
await tools.getByRole('button', { name: /强化学习任务/ }).click();
await tools.getByRole('button', { name: '连接', exact: true }).click();
await tools.getByRole('button', { name: '导入策略' }).click();
await expect(page.getByLabel('摄像头画面', { exact: true })).toBeVisible();
const pause = page.getByRole('button', { name: '⏸ 暂停' });
if (await pause.isVisible()) await pause.click();
const section = tools.getByRole('button', { name: /ONNX 策略运行/ });
if ((await section.getAttribute('aria-expanded')) === 'false') await section.click();
await expect(tools.getByText('81 / 12')).toBeVisible();
const inference = tools.getByText('推理次数', { exact: true }).locator('..');
// Stop policy as well to invalidate any pending inference, leaving a stable baseline.
const stop = tools.getByRole('button', { name: '停止', exact: true });
if (await stop.isVisible()) await stop.click();
const before = await inference.textContent();
bytes = readFileSync(resolve(`web_platform/fixtures/obstacle/${failure}.onnx`));
await tools.getByRole('button', { name: '导入策略' }).click();
await expect(tools.getByRole('alert')).toContainText(
failure === 'wrong-graph' ? '维度' : 'MissingOperatorForTransactionTest',
);
await expect(tools.getByText('81 / 12')).toBeVisible();
await expect(inference).toHaveText(before!);
await expect(page.getByLabel('摄像头画面', { exact: true })).toBeVisible();
await expect(page.getByText('训练配套物理地图', { exact: false })).toBeVisible();
await expect(page.getByRole('button', { name: '▶ 播放' })).toBeVisible();
});
}
test('新server默认Flat作业兼容旧无metadata47维ONNX,错误graph拒绝且保留策略', async ({ page }) => {
await page.addInitScript(() => sessionStorage.setItem('mujoco-local-training-token', 'test'));
const flat = {
...fixture,
taskId: 'Unitree-Go2-Flat',
observationSize: 47,
observationTerms: fixture.observationTerms.slice(0, 7),
terrain: undefined,
terrainPreset: undefined,
terrainParams: undefined,
sensorCfg: undefined,
navigation: undefined,
};
let bytes = readFileSync(resolve('web_platform/fixtures/obstacle/legacy-flat.onnx'));
const job = {
id: 'd'.repeat(32),
taskId: flat.taskId,
state: 'succeeded',
artifactReady: true,
progress: 1,
iteration: 1,
maxIterations: 1,
logs: [],
message: '',
deployment: flat,
};
await page.route('http://127.0.0.1:8765/**', (route) => {
const url = route.request().url();
if (url.endsWith('/policy.onnx'))
return route.fulfill({ contentType: 'application/octet-stream', body: bytes });
if (url.endsWith('/health'))
return route.fulfill({ json: { ready: true, tasks: [job.taskId], activeJobId: job.id } });
if (url.includes('/presets')) return route.fulfill({ json: { presets: [] } });
return route.fulfill({ json: job });
});
const actuators =
'<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.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
await page.getByRole('tab', { name: '控制台' }).click();
const tools = page.getByRole('tabpanel', { name: '控制台' });
await tools.getByRole('button', { name: /强化学习任务/ }).click();
await tools.getByRole('button', { name: '连接', exact: true }).click();
await tools.getByRole('button', { name: '导入策略' }).click();
const section = tools.getByRole('button', { name: /ONNX 策略运行/ });
if ((await section.getAttribute('aria-expanded')) === 'false') await section.click();
await expect(tools.getByText('47 / 12')).toBeVisible();
bytes = readFileSync(resolve('web_platform/fixtures/obstacle/legacy-wrong-shape.onnx'));
await tools.getByRole('button', { name: '导入策略' }).click();
await expect(tools.getByRole('alert')).toContainText('维度');
await expect(tools.getByText('47 / 12')).toBeVisible();
await expect(page.getByText('训练配套物理地图', { exact: false })).toHaveCount(0);
});
@@ -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.getByText('WASM 已加载')).toBeVisible({ timeout: 30_000 });
await page.getByRole('tab', { name: '控制台' }).click();
const tools = page.getByRole('tabpanel', { name: '控制台' });
const section = tools.getByRole('button', { name: /ONNX 策略运行/ });
if ((await section.getAttribute('aria-expanded')) === 'false') await section.click();
const chooser = page.waitForEvent('filechooser');
await tools.getByRole('button', { name: '导入 ONNX' }).click();
await (await chooser).setFiles(source!);
await expect
.poll(
async () =>
(await tools.getByText('47 / 12').count()) +
(await page.getByRole('button', { name: '技术详情', exact: true }).count()),
)
.toBeGreaterThan(0);
const details = page.getByRole('button', { name: '技术详情', exact: true });
if (await details.isVisible()) {
await details.click();
throw new Error((await page.getByRole('alert').allTextContents()).join('\n'));
}
await expect(tools.getByText('47 / 12')).toBeVisible({ timeout: 30_000 });
const enable = tools.getByRole('button', { name: '启用', exact: true });
if (await enable.isVisible()) await enable.click();
const play = page.getByRole('button', { name: '▶ 播放', exact: true });
if (await play.isVisible()) await play.click();
await expect
.poll(async () =>
Number(
(await tools.getByText('推理次数', { exact: true }).locator('..').textContent())?.replace(
/\D/g,
'',
),
),
)
.toBeGreaterThan(2);
await expect(tools.getByRole('button', { name: '设定目标', exact: true })).toHaveCount(0);
});
+177
View File
@@ -0,0 +1,177 @@
import { expect, test } from '@playwright/test';
import { spawn } from 'node:child_process';
import { createHash } from 'node:crypto';
import { mkdtempSync, readFileSync, writeFileSync } from 'node:fs';
import { createServer } from 'node:net';
import { tmpdir } from 'node:os';
import { join, resolve } from 'node:path';
const realRoot = process.env.GO2_UPLOAD_REAL_DIR;
const token = 'local-upload-browser-test-token';
async function unusedPort() {
const socket = createServer();
await new Promise<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 expect(ui.getByLabel('确认Go2 legacy47模板')).toBeEnabled();
await expect(ui.getByLabel('基础策略', { exact: true })).toHaveValue('');
await ui.getByLabel('确认Go2 legacy47模板').check();
const path = join(realRoot!, format === 'pt' ? 'model_10000.pt' : 'policy.onnx');
const sha = createHash('sha256').update(readFileSync(path)).digest('hex');
const responsePromise = page.waitForResponse((response) =>
response.url().startsWith(`${endpoint}/api/training/pretrained-sources/upload?`),
);
await ui.getByLabel('选择基础策略文件').setInputFiles(path);
const response = await responsePromise;
expect(response.status()).toBe(201);
const catalog = await (
await fetch(`${endpoint}/api/training/health`, {
headers: { Authorization: `Bearer ${token}` },
})
).json();
const record = catalog.pretrainedSources[0];
await expect(ui.getByLabel('基础策略', { exact: true })).toHaveValue(record.id);
await expect(ui.getByText(/原文件 SHA256/)).toContainText(sha);
await expect(ui.getByText(/上传格式/)).toContainText(format);
await expect(ui.getByText(/已选择.*仅继承策略权重/)).toBeVisible();
if (format === 'onnx')
await expect(ui.getByText(/ONNX统计count合成/)).toContainText('1000000');
// Intercept only train dispatch; uploads and catalog are real authenticated HTTP.
// Never dispatch 4096-env training or an Agent/three-seed evaluation from this test.
let submitted: Record<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);
}
});
}
}
+9
View File
@@ -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.
+155 -52
View File
@@ -1,3 +1,4 @@
import { resolvePolicyDeployment, type PolicyDeployment } from '../rl/deployment';
/* Zustand 的 action 引用稳定;初始化 viewer 与导入回调有意只创建一次。 */
/* eslint-disable react-hooks/exhaustive-deps */
import {
@@ -267,6 +268,7 @@ export function App() {
const manifest = useRef<ProjectManifest | null>(null),
notificationId = useRef(0),
loadInFlight = useRef(false),
policyLoadInFlight = useRef(false),
importInFlight = useRef(false),
adapter = useRef(new MainThreadPhysicsAdapter()),
root = useRef<HTMLDivElement>(null),
@@ -319,6 +321,9 @@ export function App() {
[controllerStatus, setControllerStatus] = useState<ControllerStatus>(),
[selectedPolicyPath, setSelectedPolicyPath] = useState<string>(),
[policyStatus, setPolicyStatus] = useState<RLPolicyStatus>(),
[navigationTargetMode, setNavigationTargetMode] = useState(false),
[trainingDeployment, setTrainingDeployment] = useState<PolicyDeployment>(),
[showPerceptionRays, setShowPerceptionRays] = useState(true),
[projectMaps, setProjectMaps] = useState<MapEntry[]>([]),
[editorDocument, setEditorDocument] = useState<EditableMapDocument | null>(null),
[editorDrafts, setEditorDrafts] = useState<Map<string, EditableMapDocument>>(() => new Map()),
@@ -444,6 +449,13 @@ export function App() {
setEditorSelection(selection ? { kind: 'body', bodyId: selection.bodyId } : null);
if (selection) setRightOpen(true);
},
onNavigationMode: setNavigationTargetMode,
onNavigationTarget: (target) => {
adapter.current.setNavigationTarget(target);
const snapshot = adapter.current.snapshot();
state.setSnapshot(snapshot ?? undefined);
setPolicyStatus(snapshot?.rlPolicy);
},
onFrame: (frame, fps, snapshot) => {
const memory = (performance as Performance & { memory?: { usedJSHeapSize: number } })
.memory?.usedJSHeapSize;
@@ -550,8 +562,14 @@ export function App() {
viewer.current?.setMapDisplay(showVisualMap, showMapCollision);
}, [showVisualMap, showMapCollision]);
useEffect(() => {
viewer.current?.setParametricMapAssets(placedMapAssets, mapSceneDraft.changedIds);
}, [placedMapAssets, mapSceneDraft.changedIds]);
viewer.current?.setParametricMapAssets(
trainingDeployment ? [] : placedMapAssets,
mapSceneDraft.changedIds,
);
}, [placedMapAssets, mapSceneDraft.changedIds, trainingDeployment]);
useEffect(() => {
viewer.current?.setShowPerceptionRays(showPerceptionRays);
}, [showPerceptionRays]);
useEffect(() => {
if (window.innerWidth < 900) return;
try {
@@ -585,6 +603,8 @@ export function App() {
path: string,
requestedMode?: UrdfLoadMode,
requestedSceneAssets?: readonly PlacedMapAsset[],
requestedDeployment?: PolicyDeployment,
requestedPolicy?: { data: Uint8Array; path: string },
) => {
if (!manifest.current || loadInFlight.current) return false;
// 普通模型重载始终使用上次成功应用的地图基线。只有场景提交入口可以显式
@@ -623,6 +643,8 @@ export function App() {
baseMode: baseModeRef.current,
enhancements: urdfEnhancementsRef.current,
mapAssets: sceneAssets,
trainingDeployment: requestedDeployment,
trainingPolicy: requestedPolicy,
onProgress: ({ value, label }) =>
setImportProgress({
title: '正在准备仿真',
@@ -670,9 +692,10 @@ export function App() {
await activeViewer.setVisualMaps([]);
let visualMapWarning: string | undefined;
try {
const assets = manifest.current
? visualMapAssets(manifest.current, placedMapAssetsRef.current)
: [];
const assets =
manifest.current && !requestedDeployment
? visualMapAssets(manifest.current, placedMapAssetsRef.current)
: [];
await activeViewer.setVisualMaps(assets);
} catch (error) {
visualMapWarning = `视觉地图加载失败:${error instanceof Error ? error.message : String(error)}`;
@@ -704,11 +727,12 @@ export function App() {
setAppliedMapAssets(committedMapAssets);
}
activeViewer.setParametricMapAssets(
placedMapAssetsRef.current,
requestedDeployment ? [] : placedMapAssetsRef.current,
summarizeMapSceneDraft(placedMapAssetsRef.current, appliedMapAssetsRef.current)
.changedIds,
);
adapter.current.releaseRetired();
setTrainingDeployment(requestedDeployment);
sessionSwapped = false;
return true;
} catch (error) {
@@ -734,6 +758,7 @@ export function App() {
if (retained && previousEntry) state.setEntry(previousEntry);
adapter.current.setPaused(previousPaused);
state.setPaused(previousPaused);
state.setSnapshot(retained ?? undefined);
setControllerStatus(retained?.controller);
setPolicyStatus(retained?.rlPolicy);
return false;
@@ -874,6 +899,7 @@ export function App() {
setSelectedControllerPath(undefined);
setControllerStatus(undefined);
setSelectedPolicyPath(undefined);
setTrainingDeployment(undefined);
setPolicyStatus(undefined);
setProjectMaps([]);
setCommittedEditorDocuments(new Map());
@@ -1962,26 +1988,62 @@ export function App() {
setControllerStatus(undefined);
state.setSnapshot(adapter.current.snapshot() ?? undefined);
};
const loadPolicyBytes = async (data: Uint8Array, path: string) => {
const loadPolicyBytes = async (data: Uint8Array, path: string, expected?: PolicyDeployment) => {
if (policyLoadInFlight.current || loadInFlight.current)
throw new Error('模型/策略正在加载,请稍后重试');
viewer.current?.setNavigationTargetMode(false);
const previousPaused = useAppStore.getState().paused;
const previousSession = adapter.current.session;
policyLoadInFlight.current = true;
state.setLoading(true);
setImportProgress({
title: '正在加载强化学习策略',
label: '初始化 ONNX Runtime',
detail: path,
value: 0.55,
});
state.setDiagnostic(undefined);
adapter.current.setPaused(true);
state.setPaused(true);
try {
const status = await adapter.current.loadRLPolicy(data, path);
setPolicyStatus(status);
const deployment = resolvePolicyDeployment(data, expected);
if (deployment?.terrain) {
const entry = useAppStore.getState().selectedEntry;
if (!entry || !manifest.current) throw new Error('请先导入并加载Go2机器人');
if (!(await loadEntry(entry, 'mjcf', undefined, deployment, { data, path })))
throw new Error(
useAppStore.getState().diagnostic?.detail ?? '配套训练地图加载失败,策略未启用',
);
}
setImportProgress({
title: '正在加载强化学习策略',
label: '初始化 ONNX Runtime',
detail: path,
value: 0.55,
});
state.setLoading(true);
const session = adapter.current.session;
const status = deployment?.terrain
? adapter.current.snapshot()!.rlPolicy!
: await adapter.current.loadRLPolicy(data, path, deployment);
if (session !== adapter.current.session) throw new Error('模型已切换,策略加载取消');
if (deployment?.terrain) {
adapter.current.setPaused(false);
state.setPaused(false);
setShowSensorCamera(true);
viewer.current?.setShowSensorCamera(true);
}
setPolicyStatus(adapter.current.snapshot()?.rlPolicy ?? status);
state.setSnapshot(adapter.current.snapshot() ?? undefined);
notify(
'ONNX 策略已加载',
`${status.taskName} · ${status.observationSize} → ${status.actionSize}`,
);
} catch (error) {
if (adapter.current.session === previousSession) {
adapter.current.setPaused(previousPaused);
state.setPaused(previousPaused);
state.setSnapshot(adapter.current.snapshot() ?? undefined);
setPolicyStatus(adapter.current.snapshot()?.rlPolicy);
}
state.setDiagnostic(diagnostic('仿真', error, path));
throw error;
} finally {
policyLoadInFlight.current = false;
setImportProgress(undefined);
state.setLoading(false);
}
@@ -1993,45 +2055,49 @@ export function App() {
return;
}
setSelectedPolicyPath(path);
void loadPolicyBytes(file.data, path);
void loadPolicyBytes(file.data, path).catch(() => {});
};
const importPolicy = (file: File) => {
void (async () => {
try {
if (!/\.onnx$/i.test(file.name)) throw new Error('请选择 .onnx 文件');
if (file.size > 64 * 1024 * 1024) throw new Error('ONNX 策略不能超过 64 MiB');
const path = normalizeProjectPath(file.name),
data = new Uint8Array(await file.arrayBuffer());
if (manifest.current) {
const index = manifest.current.files.findIndex((candidate) => candidate.path === path),
files = manifest.current.files.slice(),
entry = {
path,
data,
size: data.byteLength,
source: 'file' as const,
mimeType: file.type || 'application/octet-stream',
};
if (index >= 0) files[index] = entry;
else files.push(entry);
const totalBytes = files.reduce((total, item) => total + item.size, 0);
if (totalBytes > DEFAULT_IMPORT_LIMITS.maxTotalBytes)
throw new Error('加入 ONNX 后工程总大小超过 512 MiB');
manifest.current = { ...manifest.current, files, totalBytes };
state.setProject(
manifest.current.name,
files.map(({ path: filePath, size }) => ({ path: filePath, size })),
manifest.current.entries,
manifest.current.selectedEntry,
);
state.setSnapshot(adapter.current.snapshot() ?? undefined);
}
setSelectedPolicyPath(path);
await loadPolicyBytes(data, path);
} catch (error) {
state.setDiagnostic(diagnostic('仿真', error, file.name));
const importPolicy = async (file: File, expected?: PolicyDeployment) => {
try {
if (!/\.onnx$/i.test(file.name)) throw new Error('请选择 .onnx 文件');
if (file.size > 64 * 1024 * 1024) throw new Error('ONNX 策略不能超过 64 MiB');
const project = manifest.current;
const path = normalizeProjectPath(file.name),
data = new Uint8Array(await file.arrayBuffer());
if (project !== manifest.current) throw new Error('工程已切换,策略导入取消');
const projectedBytes =
(project?.files
.filter((item) => item.path !== path)
.reduce((total, item) => total + item.size, 0) ?? 0) + data.byteLength;
if (projectedBytes > DEFAULT_IMPORT_LIMITS.maxTotalBytes)
throw new Error('加入 ONNX 后工程总大小超过 512 MiB');
await loadPolicyBytes(data, path, expected);
if (project && manifest.current?.id === project.id) {
const files = manifest.current.files.filter((candidate) => candidate.path !== path);
files.push({
path,
data,
size: data.byteLength,
source: 'file',
mimeType: file.type || 'application/octet-stream',
});
const totalBytes = files.reduce((total, item) => total + item.size, 0);
if (totalBytes > DEFAULT_IMPORT_LIMITS.maxTotalBytes)
throw new Error('加入 ONNX 后工程总大小超过 512 MiB');
manifest.current = { ...manifest.current, files, totalBytes };
state.setProject(
manifest.current.name,
files.map(({ path, size }) => ({ path, size })),
manifest.current.entries,
useAppStore.getState().selectedEntry,
);
state.setSnapshot(adapter.current.snapshot() ?? undefined);
}
})();
setSelectedPolicyPath(path);
} catch (error) {
state.setDiagnostic(diagnostic('仿真', error, file.name));
throw error;
}
};
const togglePolicy = (enabled: boolean) => {
adapter.current.setRLPolicyEnabled(enabled);
@@ -2372,8 +2438,29 @@ export function App() {
onSelectPolicyPath={setSelectedPolicyPath}
onLoadPolicyPath={loadPolicyPath}
onImportPolicy={importPolicy}
compileTrainingScene={(coordinates) => {
if (
mapSceneDirty ||
trainingDeployment ||
useAppStore.getState().loading ||
loadInFlight.current
)
throw new Error('请先应用地图草稿;训练部署/加载中的场景不能同步');
return adapter.current.exportTrainingTerrain(appliedMapAssets, coordinates);
}}
trainingSceneMaps={appliedMapAssets}
trainingSceneDirty={mapSceneDirty || Boolean(trainingDeployment)}
onTogglePolicy={togglePolicy}
onPolicyCommand={setPolicyCommand}
navigationTargetMode={navigationTargetMode}
onNavigationTargetMode={(active) => viewer.current?.setNavigationTargetMode(active)}
onResetNavigationTarget={() => {
viewer.current?.setNavigationTargetMode(false);
adapter.current.resetNavigationTarget();
const snapshot = adapter.current.snapshot();
state.setSnapshot(snapshot ?? undefined);
setPolicyStatus(snapshot?.rlPolicy);
}}
onRemovePolicy={removePolicy}
onDataRecorderConfigure={configureDataRecorder}
onDataRecordingStart={startDataRecording}
@@ -2570,6 +2657,22 @@ export function App() {
progress={importProgress}
/>
<ToastViewport item={toast} onDismiss={() => setToast(undefined)} />
{trainingDeployment && (
<div className="absolute top-16 left-4 z-20 rounded bg-app p-2 text-xs">
训练配套物理地图(编辑器地图未更改,重载模型恢复)。
{trainingDeployment.terrain?.approximation && '训练专用离散近似。'}
{(policyStatus?.observationSize === 81 || policyStatus?.observationSize === 97) && (
<label>
<input
type="checkbox"
checked={showPerceptionRays}
onChange={(e) => setShowPerceptionRays(e.target.checked)}
/>
显示避障射线
</label>
)}
</div>
)}
{Boolean(state.snapshot?.model.ncam) &&
(showSensorCamera ? (
<div
@@ -1,3 +1,6 @@
import type { TrainingSceneCompiler } from '../../map/trainingMap';
import type { PolicyDeployment } from '../../rl/deployment';
import type { PlacedMapAsset } from '../../map/types';
import type { ControllerCommand, ControllerStatus } from '../../controller/types';
import { PythonControllerPanel } from '../../controller/PythonControllerPanel';
import type { RLCommand, RLPolicyStatus } from '../../rl/types';
@@ -63,8 +66,14 @@ export function WorkspaceToolsPanel({
onSelectPolicyPath,
onLoadPolicyPath,
onImportPolicy,
compileTrainingScene,
trainingSceneMaps,
trainingSceneDirty,
onTogglePolicy,
onPolicyCommand,
navigationTargetMode,
onNavigationTargetMode,
onResetNavigationTarget,
onRemovePolicy,
onDataRecorderConfigure,
onDataRecordingStart,
@@ -85,6 +94,9 @@ export function WorkspaceToolsPanel({
policyPaths: string[];
selectedPolicyPath?: string;
policyStatus?: RLPolicyStatus;
navigationTargetMode?: boolean;
onNavigationTargetMode?: (active: boolean) => void;
onResetNavigationTarget?: () => void;
onResetJoints: () => void;
onToggleJointLimits: () => void;
onToggleAdvanced: () => void;
@@ -101,7 +113,10 @@ export function WorkspaceToolsPanel({
onRemoveController: () => void;
onSelectPolicyPath: (path: string) => void;
onLoadPolicyPath: (path: string) => void;
onImportPolicy: (file: File) => void;
onImportPolicy: (file: File, deployment?: PolicyDeployment) => void | Promise<void>;
compileTrainingScene?: TrainingSceneCompiler;
trainingSceneMaps?: readonly PlacedMapAsset[];
trainingSceneDirty?: boolean;
onTogglePolicy: (enabled: boolean) => void;
onPolicyCommand: (command: RLCommand) => void;
onRemovePolicy: () => void;
@@ -251,9 +266,14 @@ export function WorkspaceToolsPanel({
loading={loading}
onSelectPath={onSelectPolicyPath}
onLoadPath={onLoadPolicyPath}
onImport={onImportPolicy}
onImport={(file) => {
void Promise.resolve(onImportPolicy(file)).catch(() => {});
}}
onToggle={onTogglePolicy}
onCommand={onPolicyCommand}
navigationTargetMode={navigationTargetMode}
onNavigationTargetMode={onNavigationTargetMode}
onResetNavigationTarget={onResetNavigationTarget}
onRemove={onRemovePolicy}
/>
</CollapsibleSection>
@@ -262,7 +282,12 @@ export function WorkspaceToolsPanel({
defaultOpen={false}
badge={<ConsoleSectionBadges category="训练" />}
>
<LocalTrainingPanel onPolicyReady={onImportPolicy} />
<LocalTrainingPanel
onPolicyReady={onImportPolicy}
compileScene={compileTrainingScene}
sceneMaps={trainingSceneMaps}
sceneDirty={trainingSceneDirty}
/>
</CollapsibleSection>
</>
) : (
@@ -0,0 +1,47 @@
import { fireEvent, render, screen } from '@testing-library/react';
import { ScalarChart } from './ScalarChart';
import { ScalarChart as TuningChart } from '../../tuning/ScalarChart';
import type uPlot from 'uplot';
const spies = vi.hoisted(() => ({
create: vi.fn(),
destroy: vi.fn(),
setData: vi.fn(),
setScale: vi.fn(),
}));
vi.mock('uplot', () => ({
default: class {
over = document.createElement('div');
scales = { x: { min: 0, max: 10 }, y: { min: 0, max: 10 } };
constructor(
options: unknown,
public data: unknown,
) {
spies.create(options, data);
}
setScale = spies.setScale;
setData = spies.setData;
setSize() {}
destroy = spies.destroy;
},
}));
it('共享原tuning API、EMA只修改曲线、悬停保留原值、缩放及卸载清理', () => {
expect(TuningChart).toBe(ScalarChart);
const points = [
{ step: 1, wallTime: 0, value: 0 },
{ step: 2, wallTime: 0, value: 10 },
];
const view = render(<ScalarChart series={[{ tag: 'loss', points }]} smoothing={0.4} />);
const [options, data] = spies.create.mock.lastCall!;
expect(data).toEqual([
[1, 2],
[0, 6],
]);
expect(options.series[1].value({} as uPlot, 6, 1, 1)).toBe('10');
expect(points[1].value).toBe(10);
fireEvent.click(screen.getByRole('button', { name: '训练与评估 Scalars 放大' }));
expect(spies.setScale).toHaveBeenCalled();
fireEvent.click(screen.getByRole('button', { name: '训练与评估 Scalars 重置缩放' }));
expect(spies.setData).toHaveBeenCalled();
view.unmount();
expect(spies.destroy).toHaveBeenCalled();
});
@@ -0,0 +1,216 @@
import { useEffect, useMemo, useRef } from 'react';
import { RotateCcw, ZoomIn, ZoomOut } from 'lucide-react';
import uPlot from 'uplot';
import 'uplot/dist/uPlot.min.css';
import type { ScalarSeries } from '../../training/types';
const COLORS = ['#38d39f', '#60a5fa', '#f59e0b', '#f472b6', '#a78bfa', '#fb7185'];
function smooth(values: (number | null)[], factor: number): (number | null)[] {
if (factor <= 0) return values;
let previous: number | undefined;
return values.map((value) => {
if (value === null) return null;
previous = previous === undefined ? value : factor * previous + (1 - factor) * value;
return previous;
});
}
export function ScalarChart({
series,
smoothing,
title = '训练与评估 Scalars',
xLabel = 'Step',
}: {
series: ScalarSeries[];
smoothing: number;
title?: string;
xLabel?: string;
}) {
const host = useRef<HTMLDivElement>(null);
const chartRef = useRef<uPlot | null>(null);
const trackZoom = useRef(false);
const zoomRanges = useRef<Partial<Record<'x' | 'y', { min: number; max: number }>>>({});
const prepared = useMemo(() => {
const steps = Array.from(
new Set(series.flatMap((item) => item.points.map((point) => point.step))),
).sort((a, b) => a - b);
const columns: uPlot.AlignedData = [steps];
for (const item of series) {
const byStep = new Map(item.points.map((point) => [point.step, point.value]));
columns.push(
smooth(
steps.map((step) => byStep.get(step) ?? null),
smoothing,
),
);
}
return columns;
}, [series, smoothing]);
useEffect(() => {
if (!host.current || series.length === 0 || prepared[0].length === 0) return;
const element = host.current;
const width = Math.max(280, Math.floor(element.getBoundingClientRect().width));
trackZoom.current = false;
const style = getComputedStyle(element);
const axisColor = style.getPropertyValue('--ui-text-tertiary').trim() || '#8fa0b5';
const gridColor = style.getPropertyValue('--ui-border').trim() || '#213044';
const chart = new uPlot(
{
width,
height: 280,
cursor: { drag: { x: true, y: true, setScale: true } },
scales: { x: { time: false } },
hooks: {
setScale: [
(instance, key) => {
if (!trackZoom.current || (key !== 'x' && key !== 'y')) return;
const scale = instance.scales[key];
if (typeof scale.min === 'number' && typeof scale.max === 'number')
zoomRanges.current[key] = { min: scale.min, max: scale.max };
},
],
},
axes: [
{ stroke: axisColor, grid: { stroke: gridColor } },
{ stroke: axisColor, grid: { stroke: gridColor } },
],
series: [
{ label: xLabel },
...series.map((item, index) => ({
label: item.tag,
stroke: COLORS[index % COLORS.length],
width: 2,
spanGaps: true,
value: (_chart: uPlot, _value: number, _seriesIndex: number, index: number | null) => {
if (index === null) return '—';
const raw = item.points.find((point) => point.step === prepared[0][index]);
return raw ? String(raw.value) : '—';
},
})),
],
},
prepared,
element,
);
chartRef.current = chart;
trackZoom.current = true;
for (const key of ['x', 'y'] as const) {
const range = zoomRanges.current[key];
if (range) chart.setScale(key, range);
}
const wheelZoom = (event: WheelEvent) => {
if (!event.deltaY) return;
event.preventDefault();
event.stopPropagation();
const bounds = chart.over.getBoundingClientRect();
if (!bounds.width || !bounds.height) return;
const xRatio = Math.min(1, Math.max(0, (event.clientX - bounds.left) / bounds.width));
const yRatio = Math.min(1, Math.max(0, (event.clientY - bounds.top) / bounds.height));
const unit = event.deltaMode === 1 ? 16 : event.deltaMode === 2 ? window.innerHeight : 1;
const factor = Math.min(2, Math.max(0.5, Math.exp(event.deltaY * unit * 0.002)));
for (const key of ['x', 'y'] as const) {
const scale = chart.scales[key];
if (typeof scale.min !== 'number' || typeof scale.max !== 'number') continue;
const ratio = key === 'x' ? xRatio : 1 - yRatio;
const anchor = scale.min + (scale.max - scale.min) * ratio;
chart.setScale(key, {
min: anchor - (anchor - scale.min) * factor,
max: anchor + (scale.max - anchor) * factor,
});
}
};
chart.over.addEventListener('wheel', wheelZoom, { passive: false });
let frame = 0;
let lastWidth = width;
const observer = new ResizeObserver(() => {
window.cancelAnimationFrame(frame);
frame = window.requestAnimationFrame(() => {
const nextWidth = Math.max(280, Math.floor(element.getBoundingClientRect().width));
if (nextWidth !== lastWidth) {
lastWidth = nextWidth;
chart.setSize({ width: nextWidth, height: 280 });
}
});
});
observer.observe(element);
return () => {
observer.disconnect();
chart.over.removeEventListener('wheel', wheelZoom);
window.cancelAnimationFrame(frame);
trackZoom.current = false;
chartRef.current = null;
chart.destroy();
};
}, [prepared, series, xLabel]);
const zoom = (factor: number) => {
const chart = chartRef.current;
if (!chart) return;
for (const key of ['x', 'y']) {
const scale = chart.scales[key];
if (typeof scale.min !== 'number' || typeof scale.max !== 'number') continue;
const center = (scale.min + scale.max) / 2;
const radius = ((scale.max - scale.min) * factor) / 2 || 1;
chart.setScale(key, { min: center - radius, max: center + radius });
}
};
const resetZoom = () => {
const chart = chartRef.current;
if (!chart) return;
zoomRanges.current = {};
trackZoom.current = false;
chart.setData(chart.data, true);
trackZoom.current = true;
};
if (series.length === 0)
return (
<div className="grid h-[280px] place-items-center rounded-lg border border-border bg-app text-xs text-text-tertiary">
当前 trial 尚无 scalar 数据
</div>
);
return (
<section className="min-w-0 overflow-hidden rounded-lg border border-border bg-app">
<header className="flex items-center justify-between gap-2 border-b border-border px-3 py-2">
<h3 className="min-w-0 truncate text-[10px] font-semibold" title={title}>
{title}
</h3>
<div className="flex shrink-0 items-center gap-1">
<button
type="button"
aria-label={`${title} 放大`}
title="放大"
className="rounded p-1 text-text-tertiary hover:bg-element-hover hover:text-text-primary"
onClick={() => zoom(0.7)}
>
<ZoomIn className="h-3.5 w-3.5" />
</button>
<button
type="button"
aria-label={`${title} 缩小`}
title="缩小"
className="rounded p-1 text-text-tertiary hover:bg-element-hover hover:text-text-primary"
onClick={() => zoom(1.4)}
>
<ZoomOut className="h-3.5 w-3.5" />
</button>
<button
type="button"
aria-label={`${title} 重置缩放`}
title="重置缩放"
className="rounded p-1 text-text-tertiary hover:bg-element-hover hover:text-text-primary"
onClick={resetZoom}
>
<RotateCcw className="h-3.5 w-3.5" />
</button>
</div>
</header>
<div ref={host} className="min-w-0 w-full overflow-hidden" />
<p className="border-t border-border px-3 py-1.5 text-[9px] text-text-tertiary">
悬停图例显示原始值;曲线可平滑。滚轮或拖拽缩放,右上角复位。
</p>
</section>
);
}
@@ -0,0 +1,137 @@
import { describe, expect, it } from 'vitest';
import {
collisionHalfBounds,
trainingTerrainFromCompiledScene,
type CompiledMapGeometry,
} from './trainingMap';
import { validateCustomTerrain, validatePolicyDeployment } from '../rl/deployment';
import deployment from '../rl/fixtures/obstacleDeployment.json';
import shared from '../../../training_server/tests/fixtures/custom-boxes.json';
const identity = [1, 0, 0, 0, 1, 0, 0, 0, 1];
function geometry(patch: Partial<CompiledMapGeometry> = {}): CompiledMapGeometry {
return {
name: '__platform_map_one__obstacle',
type: 6,
position: [1, 2, 0.5],
rotation: identity,
size: [0.4, 0.3, 0.5],
friction: [0.8, 0.005, 0.0001],
collision: true,
static: true,
map: true,
...patch,
};
}
const pose = [...shared.spawn, ...shared.spawnQuaternion];
const coords = { spawn: [-2, -1] as [number, number], target: [2, -1] as [number, number] };
describe('customTrainingMap compiled collision export', () => {
it('共享payload逐字段roundtrip,不从种子生成布局,custom metadata严格一致', () => {
const layout = trainingTerrainFromCompiledScene([geometry()], pose, 6, coords);
expect(validateCustomTerrain(layout)).toEqual(shared);
const d = {
...deployment,
terrainPreset: 'custom_boxes',
terrain: layout,
terrainParams: { size: layout.size, friction: layout.friction },
};
expect(validatePolicyDeployment(d).terrain).toEqual(shared);
expect(() =>
validatePolicyDeployment({ ...d, terrainParams: { size: 8, friction: 0.8 } }),
).toThrow();
expect(() => validatePolicyDeployment({ ...d, terrain: undefined })).toThrow();
});
it('abs(R)*halfsize支持实例旋转/平移;排除机器人和装饰,所有碰撞障碍保留', () => {
const c = Math.SQRT1_2,
r = [c, -c, 0, c, c, 0, 0, 0, 1];
const layout = trainingTerrainFromCompiledScene(
[
geometry({
name: '__platform_map_one__ground',
position: [0, 0, -0.05],
size: [5, 5, 0.05],
}),
geometry({ position: [2, 2, 0.5], rotation: r }),
geometry({ position: [-2, 2, 0.5] }),
geometry({ static: false, map: false, position: [0, 0, 0.32], size: [20, 20, 20] }),
geometry({ collision: false, type: 7, size: [30, 30, 30] }),
],
pose,
4,
coords,
);
expect(layout.boxes).toHaveLength(3);
expect(layout.size).toBe(10);
expect(layout.boxes[1].pos).toEqual([2, 2, 0.5]);
expect(layout.boxes[1].size[0]).toBeCloseTo(0.7 * c, 12);
expect(layout.boxes[1].size[1]).toBeCloseTo(0.7 * c, 12);
expect(layout.approximation).toBe(true);
validateCustomTerrain(layout);
});
it('primitive bounds有解析证明;mesh/hfield、不明静态/坑底/混合摩擦拒绝', () => {
expect(collisionHalfBounds(2, [0.4, 0, 0], identity)).toEqual([0.4, 0.4, 0.4]);
expect(collisionHalfBounds(3, [0.2, 0.5, 0], identity)).toEqual([0.2, 0.2, 0.7]);
expect(collisionHalfBounds(4, [0.2, 0.3, 0.4], identity)).toEqual([0.2, 0.3, 0.4]);
expect(collisionHalfBounds(5, [0.2, 0.5, 0], identity)).toEqual([0.2, 0.2, 0.5]);
for (const patch of [
{ type: 1 },
{ type: 7 },
{ map: false },
{ position: [1, 2, -1] },
{ friction: [0.8, 0.01, 0.0001] },
{ position: [15, 2, 0.5] },
{ size: [0, 0.3, 0.5] },
{ size: [NaN, 0.3, 0.5] },
{ name: 'ground', type: 0, position: [0, 0, 1] },
{ name: 'ground', type: 0, position: [0, 0, 0], rotation: [1, 0, 0, 0, 0, -1, 0, 1, 0] },
])
expect(() => trainingTerrainFromCompiledScene([geometry(patch)], pose, 6, coords)).toThrow();
expect(() =>
trainingTerrainFromCompiledScene(
[geometry(), geometry({ friction: [1, 0.005, 0.0001] })],
pose,
6,
coords,
),
).toThrow(/混合摩擦/);
expect(() =>
trainingTerrainFromCompiledScene(
Array.from({ length: 257 }, () => geometry()),
pose,
6,
coords,
),
).toThrow(/256/);
// Only explicitly named horizontal surface at z=0 may become standard floor.
expect(
trainingTerrainFromCompiledScene(
[geometry({ name: 'floor', type: 0, position: [0, 0, 0], map: false })],
pose,
6,
coords,
).boxes,
).toHaveLength(1);
});
it('严格有限数值/正半尺寸/字段与数量校验,不扩尺寸、不偷偷清除安全区障碍', () => {
for (const bad of [NaN, Infinity, -Infinity, -1, 0, -0, true]) {
const t = structuredClone(shared);
t.boxes[1].size[0] = bad as number;
expect(() => validateCustomTerrain(t)).toThrow();
}
for (const patch of [
{ approximation: false },
{ actualObstacleCount: 5 },
{ path: '../../etc/passwd' },
{ spawnQuaternion: [0, 0, 0, 0] },
{ target: undefined },
{ size: 25 },
])
expect(() => validateCustomTerrain({ ...shared, ...patch })).toThrow();
expect(() => validateCustomTerrain({ ...shared, target: [1.9, 2] })).toThrow(/安全区/);
validateCustomTerrain({ ...shared, target: [1.91, 2] });
validateCustomTerrain({ ...shared, target: [1.8, 2.7] });
const t = trainingTerrainFromCompiledScene([geometry()], pose, 6, { target: [1, 2] });
expect(t.boxes).toHaveLength(2);
expect(() => validateCustomTerrain(t)).toThrow(/安全区/);
});
});
+191
View File
@@ -0,0 +1,191 @@
import dynamics from '../rl/fixtures/go2CompiledDynamics.json';
import { readFileSync } from 'node:fs';
import loadMujoco from '@mujoco/mujoco';
import { describe, it, expect } from 'vitest';
import { composeTrainingMap } from './trainingMap';
import fixture from '../rl/fixtures/obstacleDeployment.json';
import golden from '../rl/fixtures/obstacleRayGolden.json';
import multiGolden from '../rl/fixtures/multiRingGolden.json';
import multiFixture from '../rl/fixtures/multiRingDeployment.json';
import { validatePolicyDeployment } from '../rl/deployment';
import { SimulationSession } from '../simulation/SimulationSession';
import { Go2ObstacleAvoidanceBindings } from '../rl/runtime/Go2ObstacleAvoidanceBindings';
const deployment = validatePolicyDeployment(fixture);
const simple = new TextEncoder().encode(
'<mujoco><worldbody><geom type="plane" size="0 0 1"/><body name="base_link"><freejoint/><geom type="box" size=".1 .1 .1" mass="1"/>' +
deployment.jointNames
.map((name) => `<body><joint name="${name}"/><geom size=".01" mass=".01"/></body>`)
.join('') +
'</body></worldbody><keyframe><key name="old"/></keyframe></mujoco>',
);
describe('trainingMap', () => {
it('替换无限地面,保留精确halfSize/世界坐标/摩擦,添加PiP camera', () => {
const xml = new TextDecoder().decode(composeTrainingMap(simple, deployment));
const doc = new DOMParser().parseFromString(xml, 'application/xml');
expect(doc.querySelector('[type="plane"], keyframe')).toBeNull();
expect(doc.querySelectorAll('geom[name^="__training_terrain_"]')).toHaveLength(25);
expect(doc.querySelector('geom[name="__training_terrain_0"]')?.getAttribute('size')).toBe(
'6 6 0.1',
);
expect(doc.querySelector('camera[name="__platform_camera__"]')).not.toBeNull();
expect(doc.querySelector('body')?.getAttribute('pos')).toBe('-5 0 0.32');
});
it('真实WASM编译Go2+共享地图;出生姿态、81维观测、32ray golden、reset和资源生命周期', async () => {
const module = await loadMujoco({
wasmBinary: readFileSync('node_modules/@mujoco/mujoco/mujoco.wasm'),
});
const original = readFileSync(
'training_server/rl/src/assets/robots/unitree_go2/xmls/go2.xml',
'utf8',
);
const doc = new DOMParser().parseFromString(original, 'application/xml');
// Strip visual meshes only for the lightweight physics test; collision/inertia/joints remain original Go2.
doc.querySelectorAll('mesh, geom[mesh]').forEach((e) => e.remove());
const bytes = composeTrainingMap(
new TextEncoder().encode(new XMLSerializer().serializeToString(doc)),
deployment,
);
const path = '/training-map-test.xml';
module.FS.writeFile(path, bytes);
const session = new SimulationSession(module, path);
try {
session.configureDeployment(deployment);
expect(Array.from(session.data.qpos.slice(0, 7))).toEqual([-5, 0, 0.32, 1, 0, 0, 0]);
expect(session.model.ncam).toBeGreaterThan(0);
for (const [name, expected] of Object.entries(dynamics.joints)) {
const joint = session.model.jnt(name);
try {
const address = Number(joint.dofadr);
expect(Number(session.model.dof_armature[address])).toBe(expected.dof_armature);
expect(Number(session.model.dof_damping[address])).toBe(expected.dof_damping);
expect(Number(session.model.dof_frictionloss[address])).toBe(expected.dof_frictionloss);
} finally {
joint.delete();
}
}
for (const [name, expected] of Object.entries(dynamics.geoms)) {
const geom = session.model.geom(name);
try {
expect(Number(geom.contype)).toBe(expected.geom_contype);
expect(Number(geom.conaffinity)).toBe(expected.geom_conaffinity);
expect(Number(geom.condim)).toBe(expected.geom_condim);
expect(Number(geom.priority)).toBe(expected.geom_priority);
expect(Number(geom.group)).toBe(expected.geom_group);
expect(Array.from(geom.solimp)).toEqual(expected.geom_solimp);
expect(Array.from(geom.friction)).toEqual(
name.includes('_foot_')
? [deployment.terrain!.friction, 0.005, 0.0001]
: expected.geom_friction,
);
} finally {
geom.delete();
}
}
expect(Array.from(session.model.qpos0.slice(7))).toEqual(Array(12).fill(0));
expect(Array.from(session.data.qpos.slice(7))).toEqual(deployment.defaultJointPosition);
const bindings = new Go2ObstacleAvoidanceBindings(
session.model,
session.data,
(id, value) => session.setActuator(id, value),
deployment,
);
for (const item of golden.cases) {
session.data.qpos.set([...item.position, ...item.quaternion]);
module.mj_forward(session.model, session.data);
const obs = bindings.observe(0, new Float32Array(12));
expect(obs).toHaveLength(81);
expect(obs.every(Number.isFinite)).toBe(true);
item.depth.forEach((distance, i) => expect(obs[47 + i]).toBeCloseTo(distance, 5));
expect(Array.from(obs.slice(11, 23))).toEqual(Array(12).fill(0));
}
session.reset();
const original = [...deployment.terrain!.target];
const target: [number, number] = [-5, 0];
bindings.setNavigationTarget(target);
target[0] = 999;
expect(bindings.navigationStatus().target).toEqual([-5, 0]);
expect(bindings.observe(0, new Float32Array(12)).slice(6, 9)).toEqual(new Float32Array(3));
bindings.setNavigationTarget([-5, 3]);
const changed = bindings.observe(0, new Float32Array(12));
expect(changed[79]).toBeCloseTo(0.5);
expect(changed[80]).toBeCloseTo(3 / deployment.terrain!.size);
expect(changed[8]).toBe(1);
expect(bindings.navigationStatus().distance).toBeCloseTo(3);
bindings.setNavigationTarget([999, -999]);
expect(bindings.navigationStatus().target).toEqual([5.5, -5.5]);
for (const value of [NaN, Infinity, -Infinity])
expect(() => bindings.setNavigationTarget([value, 0])).toThrow(/有限/);
const status = bindings.navigationStatus();
status.defaultTarget[0] = 999;
bindings.resetNavigationTarget();
expect(bindings.navigationStatus().target).toEqual(original);
expect(deployment.terrain!.target).toEqual(original);
bindings.setNavigationTarget([0, 2]);
bindings.reset(0);
expect(bindings.navigationStatus().target).toEqual(original);
const box = deployment.terrain!.boxes[1];
bindings.setNavigationTarget([box.pos[0], box.pos[1]]);
expect(bindings.navigationStatus().targetHeight).toBeCloseTo(box.pos[2] + box.size[2]);
bindings.apply(new Float32Array(12));
module.mj_step(session.model, session.data);
session.reset();
bindings.reset(0);
expect(Array.from(session.data.qpos.slice(0, 3))).toEqual(deployment.terrain!.spawn);
expect(bindings.observe(21, new Float32Array(12))).toHaveLength(81);
expect(bindings.rays).toHaveLength(32);
session.data.qpos[2] = 0.1;
module.mj_forward(session.model, session.data);
expect(() => bindings.observe(22, new Float32Array(12))).toThrow(/导航安全停止/);
} finally {
session.dispose();
module.FS.unlink(path);
}
}, 60_000);
});
it('真实WASM完整48ray/97obs:CPU完整姿态/低障碍/坑边golden,绑定缓存按模型隔离', async () => {
const module = await loadMujoco({
wasmBinary: readFileSync('node_modules/@mujoco/mujoco/mujoco.wasm'),
});
const original = readFileSync(
'training_server/rl/src/assets/robots/unitree_go2/xmls/go2.xml',
'utf8',
);
const doc = new DOMParser().parseFromString(original, 'application/xml');
doc.querySelectorAll('mesh, geom[mesh]').forEach((e) => e.remove());
const robot = new TextEncoder().encode(new XMLSerializer().serializeToString(doc));
for (const layout of ['default', 'low', 'edge'] as const) {
const boxes = multiGolden.layouts[layout];
const d = validatePolicyDeployment({
...multiFixture,
terrain: { ...multiFixture.terrain, boxes, actualObstacleCount: boxes.length - 1 },
});
const path = '/multi-ring-wasm.xml';
module.FS.writeFile(path, composeTrainingMap(robot, d));
const session = new SimulationSession(module, path);
try {
session.configureDeployment(d);
const bindings = new Go2ObstacleAvoidanceBindings(session.model, session.data, () => {}, d);
for (const c of multiGolden.cases.filter((c) => c.layout === layout)) {
session.data.qpos.set([...c.position, ...c.quaternion]);
module.mj_forward(session.model, session.data);
const obs = bindings.observe(0, new Float32Array(12));
expect(obs).toHaveLength(97);
expect(bindings.rays).toHaveLength(48);
c.depth.forEach((v, i) => expect(obs[47 + i]).toBeCloseTo(v, 5));
expect(Array.from(obs.slice(11, 23))).toEqual(Array(12).fill(0));
}
bindings.setNavigationTarget([-5, 3]);
session.reset();
const obs = bindings.observe(0, new Float32Array(12));
expect(obs[95]).toBeCloseTo(0.5);
expect(obs[96]).toBeCloseTo(0.25);
bindings.clear();
expect(bindings.rays).toEqual([]);
} finally {
session.dispose();
module.FS.unlink(path);
}
}
}, 60_000);
+199
View File
@@ -0,0 +1,199 @@
import {
validatePolicyDeployment,
type PolicyDeployment,
type TrainingTerrain,
} from '../rl/deployment';
/** Input must be flattened MJCF (mj_saveLastXML), so included ground cannot survive composition. */
export function composeTrainingMap(source: Uint8Array, input: PolicyDeployment): Uint8Array {
const deployment = validatePolicyDeployment(input),
terrain = deployment.terrain;
if (!terrain) throw new Error('缺少训练地图');
const doc = new DOMParser().parseFromString(new TextDecoder().decode(source), 'application/xml');
if (doc.querySelector('parsererror, include')) throw new Error('训练地图需要已展开的有效MJCF');
const world = doc.querySelector('mujoco > worldbody');
if (!world) throw new Error('缺少worldbody');
const robots = Array.from(world.children).filter(
(e) => e.tagName === 'body' && e.querySelector('freejoint, joint[type="free"]'),
);
if (robots.length !== 1) throw new Error('训练评测需要唯一浮动基座机器人');
const robot = robots[0];
for (const child of Array.from(world.children))
if (child !== robot && child.tagName !== 'light') child.remove();
doc.querySelector('keyframe')?.remove();
for (const key of ['euler', 'axisangle', 'xyaxes', 'zaxis']) robot.removeAttribute(key);
robot.setAttribute('pos', terrain.spawn.join(' '));
robot.setAttribute('quat', terrain.spawnQuaternion.join(' '));
// Match go2_constants.py FULL_COLLISION + obstacle terrain startup foot friction.
// Raw go2.xml does not contain these mjlab Entity compilation overrides.
for (const geom of robot.querySelectorAll('geom')) {
const name = geom.getAttribute('name') ?? '';
if (!name.endsWith('_collision')) continue;
const foot = /^[FR][LR]_foot_collision$/.test(name);
for (const [key, value] of Object.entries({
contype: '1',
conaffinity: '0',
condim: foot ? '3' : '1',
priority: foot ? '1' : '0',
group: '3',
friction: `${foot ? terrain.friction : 1} 0.005 0.0001`,
solimp: `0.9 0.95 ${foot ? 0.023 : 0.001} 0.5 2`,
}))
geom.setAttribute(key, value);
}
let actuators = doc.querySelector('mujoco > actuator');
if (!actuators) {
actuators = doc.createElement('actuator');
doc.documentElement.append(actuators);
}
for (const name of deployment.jointNames) {
const joint = robot.querySelector(`joint[name="${name}"]`);
if (!joint) throw new Error(`机器人缺少训练关节:${name}`);
joint.setAttribute('armature', name.includes('_calf_') ? '0.02' : '0.01');
joint.setAttribute('damping', '0');
joint.setAttribute('frictionloss', '0');
if (!actuators.querySelector(`[joint="${name}"]`)) {
const motor = doc.createElement('motor');
motor.setAttribute('name', `${name}_motor`);
motor.setAttribute('joint', name);
motor.setAttribute('gear', '1');
actuators.append(motor);
}
}
for (let i = 0; i < terrain.boxes.length; i++) {
const box = terrain.boxes[i],
geom = doc.createElement('geom');
for (const [key, value] of Object.entries({
name: `__training_terrain_${i}`,
type: 'box',
pos: box.pos.join(' '),
size: box.size.join(' '),
quat: '1 0 0 0',
group: '2',
contype: '1',
conaffinity: '1',
friction: `${terrain.friction} 0.005 0.0001`,
priority: '1',
condim: '3',
rgba: i === 0 ? '0.35 0.4 0.45 1' : '0.65 0.35 0.2 1',
}))
geom.setAttribute(key, value);
world.append(geom);
}
// Camera's -Z optical axis looks along robot +X, +Y is image-up along robot +Z.
doc.querySelector('camera[name="__platform_camera__"]')?.remove();
const camera = doc.createElement('camera');
for (const [key, value] of Object.entries({
name: '__platform_camera__',
pos: '0.3 0 0.05',
xyaxes: '0 -1 0 0 0 1',
fovy: '60',
}))
camera.setAttribute(key, value);
robot.append(camera);
return new TextEncoder().encode(new XMLSerializer().serializeToString(doc));
}
export interface TrainingSceneCoordinates {
spawn?: [number, number];
target?: [number, number];
}
export type TrainingSceneCompiler = (coordinates?: TrainingSceneCoordinates) => TrainingTerrain;
export interface CompiledMapGeometry {
name: string;
type: number;
position: number[];
rotation: number[];
size: number[];
friction: number[];
collision: boolean;
static: boolean;
map: boolean;
}
/** Exact analytic world bounds for primitive MuJoCo collision shapes (not visual bounds). */
export function collisionHalfBounds(
type: number,
s: readonly number[],
r: readonly number[],
): number[] {
if (![2, 3, 4, 5, 6].includes(type))
throw new Error('不支持的静态碰撞形状:mesh/hfield无法精确提取;请改用box等基本几何');
return [0, 1, 2].map((i) => {
const row = r.slice(i * 3, i * 3 + 3);
if (type === 2) return s[0]; // sphere
if (type === 3) return s[0] + Math.abs(row[2]) * s[1]; // capsule, axial half-length
if (type === 4) return Math.hypot(...row.map((v, j) => v * s[j])); // ellipsoid
if (type === 5) return Math.hypot(row[0], row[1]) * s[0] + Math.abs(row[2]) * s[1];
return row.reduce((sum, v, j) => sum + Math.abs(v) * s[j], 0); // box: abs(R)*halfsize
});
}
/** Uses already compiled, applied physics. Deliberately does not run terrain RNG or parse XML.
* Returns a candidate so the UI can expose unsafe suggested coordinates for explicit correction.
* validateCustomTerrain is mandatory before declaring synchronization/upload success.
*/
export function trainingTerrainFromCompiledScene(
geometries: readonly CompiledMapGeometry[],
initialPose: readonly number[],
mapExtent: number,
coordinates: TrainingSceneCoordinates = {},
): TrainingTerrain {
const obstacles: TrainingTerrain['boxes'] = [];
let friction: number | undefined;
let extent = Math.max(4, mapExtent);
for (const g of geometries) {
if (!g.static || !g.collision) continue;
if (![...g.position, ...g.rotation, ...g.size, ...g.friction].every(Number.isFinite))
throw new Error(`静态碰撞几何 ${g.name} 含非有限数值`);
const supportName = /(?:^|[_-])(?:floor|ground|flat)(?:[_-]|$)/i.test(g.name);
if (!g.map && !(g.type === 0 && supportName))
throw new Error(`发现未声明为已应用地图的静态碰撞几何 ${g.name},不能静默漏障碍`);
if (
g.friction.length !== 3 ||
Math.abs(g.friction[1] - 0.005) > 1e-9 ||
Math.abs(g.friction[2] - 0.0001) > 1e-9 ||
(friction !== undefined && Math.abs(g.friction[0] - friction) > 1e-9)
)
throw new Error(
'地图混合摩擦无法用单一friction表示;请先显式统一各碰撞几何摩擦为 [f,0.005,0.0001]',
);
friction = g.friction[0];
const horizontal = Math.abs(Math.abs(g.rotation[8]) - 1) < 1e-9;
if (g.type === 0) {
if (!supportName || !horizontal || g.rotation[8] < 0 || Math.abs(g.position[2]) > 1e-6)
throw new Error('仅支持明确命名floor/ground的水平z=0支撑plane;倾斜/非零高度plane无法导出');
continue;
}
const half = collisionHalfBounds(g.type, g.size, g.rotation);
if (half.some((v) => !Number.isFinite(v) || v <= 0))
throw new Error('碰撞半尺寸必须严格大于零');
extent = Math.max(extent, ...[0, 1].map((i) => Math.abs(g.position[i]) + half[i]));
const top = g.position[2] + half[2],
bottom = g.position[2] - half[2];
if (g.type === 6 && supportName && horizontal && Math.abs(top) <= 1e-6 && bottom < 0) continue;
if (bottom < -1e-6)
throw new Error(`几何 ${g.name} 含地下/坑底结构;标准floor会填平,拒绝同步`);
obstacles.push({ pos: [...g.position], size: half, yaw: 0 });
if (obstacles.length > 256) throw new Error('自定义障碍物超过256,不能截断');
}
if (friction === undefined) throw new Error('已应用地图没有可导出的静态碰撞几何');
if (extent > 12 + 1e-6) throw new Error('世界地图范围超过24m上限;不平移或裁剪几何');
const size = Math.max(8, extent * 2);
const spawn = coordinates.spawn ?? [initialPose[0], initialPose[1]];
const q = initialPose.slice(3, 7);
const yaw = Math.atan2(2 * (q[0] * q[3] + q[1] * q[2]), 1 - 2 * (q[2] ** 2 + q[3] ** 2));
const target = coordinates.target ?? [spawn[0] + 3 * Math.cos(yaw), spawn[1] + 3 * Math.sin(yaw)];
return {
representation: 'boxes-v1',
approximation: true,
size,
friction,
boxes: [{ pos: [0, 0, -0.1], size: [size / 2, size / 2, 0.1], yaw: 0 }, ...obstacles],
spawn: [spawn[0], spawn[1], 0.32],
spawnQuaternion: q,
target: [...target],
actualObstacleCount: obstacles.length,
};
}
@@ -0,0 +1,64 @@
import { fireEvent, render, screen } from '@testing-library/react';
import { RLPolicyPanel } from './RLPolicyPanel';
import type { RLPolicyStatus } from './types';
it('仅避障策略显示实时导航状态,支持暂停时设定/复位', () => {
const status: RLPolicyStatus = {
taskId: 'Unitree-Go2-ObstacleAvoidance',
taskName: '避障',
path: 'policy.onnx',
loaded: true,
enabled: false,
controlHz: 50,
observationSize: 81,
actionSize: 12,
inputName: 'in',
outputName: 'out',
command: { linearX: 0, linearY: 0, angularZ: 0 },
inferenceCount: 0,
lastInferenceMs: 0,
navigation: { target: [2, -1], defaultTarget: [5, 0], distance: 3.25, targetHeight: 0 },
};
const props = {
paths: [],
loading: false,
status,
onSelectPath: vi.fn(),
onLoadPath: vi.fn(),
onImport: vi.fn(),
onToggle: vi.fn(),
onCommand: vi.fn(),
onRemove: vi.fn(),
onNavigationTargetMode: vi.fn(),
onResetNavigationTarget: vi.fn(),
};
const view = render(<RLPolicyPanel {...props} />);
expect(screen.getByText('(2.00, -1.00)')).toBeInTheDocument();
expect(screen.getByText('3.25 m')).toBeInTheDocument();
fireEvent.click(screen.getByRole('button', { name: /^设定目标$/ }));
expect(props.onNavigationTargetMode).toHaveBeenCalledWith(true);
fireEvent.click(screen.getByRole('button', { name: '复位目标点' }));
expect(props.onResetNavigationTarget).toHaveBeenCalledOnce();
expect(props.onToggle).not.toHaveBeenCalled();
expect(screen.getByText(/策略未启用:设定目标不会自动启动/)).toBeVisible();
view.rerender(<RLPolicyPanel {...props} status={{ ...status, enabled: true }} />);
expect(screen.getByText(/仿真已暂停/)).toBeVisible();
view.rerender(<RLPolicyPanel {...props} status={{ ...status, error: '跌倒' }} />);
expect(screen.getByText(/导航安全停止:请重置并重新启用/)).toBeVisible();
view.rerender(
<RLPolicyPanel
{...props}
navigationTargetMode
status={{ ...status, navigation: { ...status.navigation!, distance: 1.1 } }}
/>,
);
expect(screen.getByText('1.10 m')).toBeInTheDocument();
expect(screen.getByText(/Esc 取消/)).toBeInTheDocument();
view.rerender(
<RLPolicyPanel
{...props}
status={{ ...status, taskId: 'Unitree-Go2-Flat', observationSize: 47 }}
/>,
);
expect(screen.queryByText('复位目标点')).not.toBeInTheDocument();
});
+79 -30
View File
@@ -1,6 +1,7 @@
import { useRef, type ChangeEvent } from 'react';
import { BrainCircuit, FileUp, Power, RotateCw, Trash2 } from 'lucide-react';
import type { RLCommand, RLPolicyStatus } from './types';
import { useAppStore } from '../stores/useAppStore';
import { Badge, Button, PropertyRow, Select } from '../components/ui';
export interface RLPolicyPanelProps {
@@ -14,6 +15,9 @@ export interface RLPolicyPanelProps {
onToggle(enabled: boolean): void;
onCommand(command: RLCommand): void;
onRemove(): void;
navigationTargetMode?: boolean;
onNavigationTargetMode?(active: boolean): void;
onResetNavigationTarget?(): void;
}
export function RLPolicyPanel({
@@ -27,8 +31,12 @@ export function RLPolicyPanel({
onToggle,
onCommand,
onRemove,
navigationTargetMode = false,
onNavigationTargetMode,
onResetNavigationTarget,
}: RLPolicyPanelProps) {
const input = useRef<HTMLInputElement>(null);
const paused = useAppStore((state) => state.paused);
const importFile = (event: ChangeEvent<HTMLInputElement>) => {
const file = event.target.files?.[0];
if (file) onImport(file);
@@ -98,36 +106,77 @@ export function RLPolicyPanel({
/>
<PropertyRow label="推理次数" value={status.inferenceCount} />
<PropertyRow label="上次推理" value={`${status.lastInferenceMs.toFixed(2)} ms`} />
<div className="mt-3 border-t border-border pt-3">
<p className="mb-2 text-[10px] text-text-tertiary">速度指令(机身坐标系)</p>
<CommandInput
label="前向 m/s"
value={command.linearX}
min={-0.5}
max={1}
onChange={(linearX) => onCommand({ ...command, linearX })}
/>
<CommandInput
label="侧向 m/s"
value={command.linearY}
min={-0.5}
max={0.5}
onChange={(linearY) => onCommand({ ...command, linearY })}
/>
<CommandInput
label="偏航 rad/s"
value={command.angularZ}
min={-1}
max={1}
onChange={(angularZ) => onCommand({ ...command, angularZ })}
/>
<Button
className="mt-1 w-full"
onClick={() => onCommand({ linearX: 0, linearY: 0, angularZ: 0 })}
>
停止移动
</Button>
</div>
{status.taskId === 'Unitree-Go2-ObstacleAvoidance' && status.navigation && (
<div className="mt-3 border-t border-border pt-3">
<PropertyRow
label="当前目标 (X, Y)"
value={`(${status.navigation.target[0].toFixed(2)}, ${status.navigation.target[1].toFixed(2)})`}
/>
<PropertyRow label="剩余距离" value={`${status.navigation.distance.toFixed(2)} m`} />
<p role="status" className="text-xs text-text-secondary">
{status.error
? '导航安全停止:请重置并重新启用策略。'
: !status.enabled
? '策略未启用:设定目标不会自动启动,请手动启用。'
: paused
? '仿真已暂停:目标已设定,请恢复仿真后运动。'
: '导航控制运行中(不保证到达);可持续更换目标。'}
</p>
<div className="mt-2 grid grid-cols-2 gap-2">
<Button
disabled={loading}
aria-pressed={navigationTargetMode}
onClick={() => onNavigationTargetMode?.(!navigationTargetMode)}
>
{navigationTargetMode ? '取消设定目标' : '设定目标'}
</Button>
<Button disabled={loading} onClick={onResetNavigationTarget}>
复位目标点
</Button>
</div>
{navigationTargetMode && (
<p className="mt-2 text-xs text-accent">
点击主视口地形设定目标,Esc 取消;不会选择或移动地图对象。
</p>
)}
</div>
)}
{status.observationSize === 81 || status.observationSize === 97 ? (
<p className="mt-3 text-xs">
自动导航到配套目标,到达后停止指令。评测20秒/跌倒/越界时停止(非训练端自动重置);请重置后启用。水平射线有矮障碍/跌落盲区,Go2-W不是Go2同构模型。
</p>
) : (
<div className="mt-3 border-t border-border pt-3">
<p className="mb-2 text-[10px] text-text-tertiary">速度指令(机身坐标系)</p>
<CommandInput
label="前向 m/s"
value={command.linearX}
min={-0.5}
max={1}
onChange={(linearX) => onCommand({ ...command, linearX })}
/>
<CommandInput
label="侧向 m/s"
value={command.linearY}
min={-0.5}
max={0.5}
onChange={(linearY) => onCommand({ ...command, linearY })}
/>
<CommandInput
label="偏航 rad/s"
value={command.angularZ}
min={-1}
max={1}
onChange={(angularZ) => onCommand({ ...command, angularZ })}
/>
<Button
className="mt-1 w-full"
onClick={() => onCommand({ linearX: 0, linearY: 0, angularZ: 0 })}
>
停止移动
</Button>
</div>
)}
{status.error && (
<p
role="alert"
+127
View File
@@ -0,0 +1,127 @@
import { describe, it, expect } from 'vitest';
import fixture from './fixtures/obstacleDeployment.json';
import {
readPolicyDeployment,
validatePolicyDeployment,
resolvePolicyDeployment,
} from './deployment';
function varint(value: number): number[] {
const b: number[] = [];
do {
const part = value % 128;
value = Math.floor(value / 128);
b.push(part | (value ? 128 : 0));
} while (value);
return b;
}
function field(tag: number, content: number[]): number[] {
return [...varint(tag * 8 + 2), ...varint(content.length), ...content];
}
const string = (s: string) => Array.from(new TextEncoder().encode(s));
export function metadataModel(value: unknown): Uint8Array {
return new Uint8Array(
field(14, [
...field(1, string('platform_deployment')),
...field(2, string(JSON.stringify(value))),
]),
);
}
describe('policy deployment metadata', () => {
it('解析ONNX ModelProto的部署元数据且忽略graph,不依赖ORT', () => {
const bytes = metadataModel(fixture);
expect(readPolicyDeployment(bytes)).toEqual(validatePolicyDeployment(fixture));
expect(readPolicyDeployment(new Uint8Array([8, 9]))).toBeUndefined();
expect(() => readPolicyDeployment(bytes.slice(0, -1))).toThrow(/截断|长度/);
expect(() => readPolicyDeployment(new Uint8Array([...bytes, ...bytes]))).toThrow(/重复/);
});
it.each([
['version', 2],
['taskId', 'Unitree-Go2-Rough'],
['browserCompatible', false],
['observationSize', 47],
['actionSize', 16],
['controlHz', 200],
['seed', NaN],
['jointNames', []],
['effortLimits', []],
['observationTerms', []],
])('拒绝不兼容%s', (key, value) =>
expect(() => validatePolicyDeployment({ ...fixture, [key]: value })).toThrow(),
);
it('拒绝不支持sensor、超限box、不有限数字及出生点偏移', () => {
expect(() =>
validatePolicyDeployment({ ...fixture, sensorCfg: { ...fixture.sensorCfg, rayCount: 64 } }),
).toThrow();
expect(() =>
validatePolicyDeployment({
...fixture,
sensorCfg: { ...fixture.sensorCfg, type: 'camera_depth' },
}),
).toThrow();
expect(() =>
validatePolicyDeployment({
...fixture,
terrain: { ...fixture.terrain, boxes: Array(258).fill(fixture.terrain.boxes[0]) },
}),
).toThrow();
expect(() =>
validatePolicyDeployment({ ...fixture, terrain: { ...fixture.terrain, friction: Infinity } }),
).toThrow();
expect(() =>
validatePolicyDeployment({
...fixture,
terrain: { ...fixture.terrain, spawn: [0, 0, 0.32] },
}),
).toThrow();
});
});
it('新server默认Flat允许旧无metadata导出,仅返回需要真实graph47→12校验的预期契约', () => {
const { terrain, terrainPreset, terrainParams, sensorCfg, navigation, ...base } =
validatePolicyDeployment(fixture);
const flat = validatePolicyDeployment({
...base,
taskId: 'Unitree-Go2-Flat',
observationSize: 47,
observationTerms: base.observationTerms.slice(0, 7),
});
const legacy = new Uint8Array([8, 9]);
expect(resolvePolicyDeployment(legacy, flat)).toEqual(flat);
expect(() => resolvePolicyDeployment(legacy, validatePolicyDeployment(fixture))).toThrow(
/不一致/,
);
expect(() =>
resolvePolicyDeployment(legacy, { ...flat, terrain, terrainPreset, terrainParams }),
).toThrow(/不一致/);
expect(() => resolvePolicyDeployment(legacy, { ...flat, sensorCfg })).toThrow(/不一致/);
expect(() => resolvePolicyDeployment(legacy, { ...flat, navigation })).toThrow(/不一致/);
});
it('custom_boxes ONNX metadata保持完整布局且绝不降级旧Flat或preset', async () => {
const { default: terrain } =
await import('../../../training_server/tests/fixtures/custom-boxes.json');
const custom = validatePolicyDeployment({
...fixture,
terrainPreset: 'custom_boxes',
terrain,
terrainParams: { size: terrain.size, friction: terrain.friction },
});
const bytes = metadataModel(custom); // ModelProto metadata fixture, not a trained graph.
expect(readPolicyDeployment(bytes)).toEqual(custom);
expect(resolvePolicyDeployment(bytes, custom)).toEqual(custom);
expect(() => resolvePolicyDeployment(new Uint8Array([8, 9]), custom)).toThrow(/不一致/);
expect(() => resolvePolicyDeployment(metadataModel(fixture), custom)).toThrow(/不一致/);
});
it('调参导航速度有限且在专属范围,旧0.6契约仍有效', () => {
for (const speed of [0.3, 0.6, 0.9, 1.2]) {
const value = structuredClone(fixture);
value.navigation.speed = speed;
expect(readPolicyDeployment(metadataModel(value))?.navigation?.speed).toBe(speed);
}
for (const speed of [0, 1.21, NaN, Infinity]) {
const value = structuredClone(fixture);
value.navigation.speed = speed;
expect(() => validatePolicyDeployment(value)).toThrow();
}
});
+406
View File
@@ -0,0 +1,406 @@
import { GO2W_VELOCITY_TASK } from './tasks/go2wVelocity';
export const OBSTACLE_TASK_ID = 'Unitree-Go2-ObstacleAvoidance';
export interface TrainingTerrain {
representation: 'boxes-v1';
approximation: boolean;
size: number;
friction: number;
boxes: { pos: number[]; size: number[]; yaw: number }[];
spawn: number[];
spawnQuaternion: number[];
target: number[];
actualObstacleCount: number;
}
export interface ObstacleSensorConfig {
type: 'raycast';
fov: number;
maxDistance: number;
safetyDistance: number;
avoidanceWeight: number;
rayCount: 32 | 48;
sensorMode?: 'single_ring_raycast' | 'multi_ring_raycast';
pitchAngles?: number[];
yawCount?: number;
yawAngles?: number[];
angleUnit?: 'deg';
rayOrder?: 'layer-major';
offset: number[];
alignment: 'base';
terrainOnly: true;
includeGround: true;
}
export interface PolicyDeployment {
version: 1;
taskId: string;
browserCompatible: true;
observationSize: number;
actionSize: 12;
controlHz: 50;
gaitPeriod: number;
jointNames: string[];
defaultJointPosition: number[];
actionScale: number[];
stiffness: number[];
damping: number[];
effortLimits: number[];
observationTerms: string[];
seed: number;
terrainPreset?: string;
terrainParams?: Record<string, number>;
terrain?: TrainingTerrain;
sensorCfg?: ObstacleSensorConfig;
navigation?: {
speed: number;
arrivalRadius: number;
distanceScale: number;
headingScale: number;
yawGain: number;
maxYawRate: number;
episodeSeconds: number;
onArrival: 'stop';
onReset: 'respawn';
};
}
function requireValue(condition: unknown, label: string): asserts condition {
if (!condition) throw new Error(`不支持或无效的策略部署配置:${label}`);
}
function record(value: unknown): Record<string, unknown> {
requireValue(value !== null && typeof value === 'object' && !Array.isArray(value), '对象');
return value as Record<string, unknown>;
}
function number(value: unknown, min: number, max: number): asserts value is number {
requireValue(
typeof value === 'number' && Number.isFinite(value) && value >= min && value <= max,
'有限数值范围',
);
}
function vector(value: unknown, length: number, min = -100, max = 100): asserts value is number[] {
requireValue(Array.isArray(value) && value.length === length, `向量长度 ${length}`);
value.forEach((v) => number(v, min, max));
}
function canonical(value: unknown): string {
if (Array.isArray(value)) return `[${value.map(canonical).join(',')}]`;
if (value !== null && typeof value === 'object')
return `{${Object.entries(value)
.sort(([a], [b]) => a.localeCompare(b))
.map(([key, v]) => `${JSON.stringify(key)}:${canonical(v)}`)
.join(',')}}`;
return JSON.stringify(value);
}
export function policyDeploymentsMatch(a: PolicyDeployment, b: PolicyDeployment): boolean {
return canonical(validatePolicyDeployment(a)) === canonical(validatePolicyDeployment(b));
}
function equal(actual: unknown, expected: unknown, label: string) {
requireValue(canonical(actual) === canonical(expected), label);
}
/** Strict custom layout boundary shared by scene export, upload and ONNX metadata. */
export function validateCustomTerrain(value: unknown): TrainingTerrain {
const t = record(value);
const keys = (v: Record<string, unknown>, expected: string[]) =>
equal(Object.keys(v).sort(), expected.sort(), '自定义布局字段(不接受路径/MJCF)');
keys(t, [
'representation',
'approximation',
'size',
'friction',
'boxes',
'spawn',
'spawnQuaternion',
'target',
'actualObstacleCount',
]);
equal(t.representation, 'boxes-v1', '地图表示');
equal(t.approximation, true, 'custom_boxes 必须标记AABB/底板标准化近似');
number(t.size, 8, 24);
number(t.friction, 0.2, 2);
requireValue(
Array.isArray(t.boxes) && t.boxes.length >= 1 && t.boxes.length <= 257,
'最多256障碍物',
);
for (const item of t.boxes) {
const b = record(item);
keys(b, ['pos', 'size', 'yaw']);
vector(b.pos, 3, -12, 12);
vector(b.size, 3, 0, 12);
requireValue(
b.size.every((v) => v > 0),
'半尺寸必须严格大于零',
);
equal(b.yaw, 0, 'box yaw');
for (let i = 0; i < 2; i++)
requireValue(Math.abs(b.pos[i]) + b.size[i] <= t.size / 2 + 1e-6, '世界地图边界');
requireValue(b.pos[2] - b.size[2] >= -0.2 - 1e-6 && b.pos[2] + b.size[2] <= 12, 'box高度边界');
}
equal(
t.boxes[0],
{ pos: [0, 0, -0.1], size: [t.size / 2, t.size / 2, 0.1], yaw: 0 },
'标准floor',
);
equal(t.actualObstacleCount, t.boxes.length - 1, '实际障碍数');
vector(t.spawn, 3, -12, 12);
equal(t.spawn[2], 0.32, '出生高度');
vector(t.target, 2, -12, 12);
vector(t.spawnQuaternion, 4, -1, 1);
requireValue(
Math.abs(t.spawnQuaternion.reduce((sum, v) => sum + v * v, 0) - 1) <= 1e-6,
'归一化出生四元数',
);
for (const point of [t.spawn, t.target]) {
requireValue(
point.slice(0, 2).every((v) => Math.abs(v) <= (t.size as number) / 2 - 0.5),
'起终点0.5m安全区越界',
);
for (const item of t.boxes.slice(1)) {
const b = item as TrainingTerrain['boxes'][number];
const distanceSquared = [0, 1].reduce(
(sum, i) => sum + Math.max(Math.abs(point[i] - b.pos[i]) - b.size[i], 0) ** 2,
0,
);
requireValue(
distanceSquared > 0.25,
'障碍物侵占起终点0.5m圆形安全区;请修改坐标,不会清除障碍',
);
}
}
return structuredClone(t) as unknown as TrainingTerrain;
}
/** A single spelling for the supported angular patterns; legacy metadata defaults to 32. */
export function obstacleSensorPattern(mode: unknown, fov: unknown) {
requireValue(
mode === undefined || mode === 'single_ring_raycast' || mode === 'multi_ring_raycast',
'sensorMode',
);
number(fov, 30, 120);
const multi = mode === 'multi_ring_raycast';
const yawCount = multi ? 16 : 32;
return {
sensorMode: multi ? ('multi_ring_raycast' as const) : ('single_ring_raycast' as const),
rayCount: multi ? (48 as const) : (32 as const),
pitchAngles: multi ? [0, -20, -45] : [0],
yawCount,
yawAngles: Array.from({ length: yawCount }, (_, i) => -fov / 2 + (i * fov) / (yawCount - 1)),
angleUnit: 'deg' as const,
rayOrder: 'layer-major' as const,
};
}
/** v1 only: reject incompatible contracts before altering the simulation or allocating ORT. */
export function validatePolicyDeployment(value: unknown): PolicyDeployment {
const d = structuredClone(record(value));
requireValue(JSON.stringify(d).length <= 100_000, '配置大小');
requireValue(
d.version === 1 && d.browserCompatible === true,
'版本/浏览器兼容性(旧 Rough 不支持)',
);
requireValue(d.taskId === OBSTACLE_TASK_ID || d.taskId === 'Unitree-Go2-Flat', '任务');
if (d.taskId !== OBSTACLE_TASK_ID) equal(d.observationSize, 47, '观测维数');
equal(d.actionSize, 12, '动作维数');
equal(d.controlHz, 50, '控制频率');
equal(d.gaitPeriod, 0.6, '步态周期');
for (const key of [
'jointNames',
'defaultJointPosition',
'actionScale',
'stiffness',
'damping',
] as const)
equal(d[key], GO2W_VELOCITY_TASK[key], key);
equal(
d.effortLimits,
[23.5, 23.5, 45, 23.5, 23.5, 45, 23.5, 23.5, 45, 23.5, 23.5, 45],
'力矩限幅',
);
equal(
d.observationTerms,
[
'base_ang_vel',
'projected_gravity',
'command',
'phase',
'joint_pos',
'joint_vel',
'actions',
...(d.taskId === OBSTACLE_TASK_ID ? ['forward_depth', 'target_error'] : []),
],
'观测顺序',
);
number(d.seed, 0, 2147483647);
requireValue(Number.isInteger(d.seed), 'seed');
if (d.terrain !== undefined) {
if (d.terrainPreset === 'custom_boxes') {
const t = validateCustomTerrain(d.terrain);
equal(d.terrainParams, { size: t.size, friction: t.friction }, 'custom_boxes参数与布局一致');
} else {
const t = record(d.terrain);
equal(t.representation, 'boxes-v1', '地图表示');
requireValue(typeof t.approximation === 'boolean', '近似标记');
number(t.size, 8, 24);
number(t.friction, 0.2, 2);
requireValue(
Array.isArray(t.boxes) && t.boxes.length >= 1 && t.boxes.length <= 257,
'box数量',
);
for (const item of t.boxes) {
const box = record(item);
vector(box.pos, 3, -12, 12);
vector(box.size, 3, 0.0001, 12);
equal(box.yaw, 0, 'box yaw');
for (let i = 0; i < 2; i++)
requireValue(Math.abs(box.pos[i]) + box.size[i] <= t.size / 2 + 1e-6, 'box边界');
}
equal(t.boxes[0], { pos: [0, 0, -0.1], size: [t.size / 2, t.size / 2, 0.1], yaw: 0 }, '地板');
equal(t.spawn, [-t.size / 2 + 1, 0, 0.32], '出生点');
equal(t.spawnQuaternion, [1, 0, 0, 0], '出生姿态');
equal(t.target, [t.size / 2 - 1, 0], '目标');
equal(t.actualObstacleCount, t.boxes.length - 1, '实际障碍数');
requireValue(
['plane', 'discrete_obstacles', 'rough', 'wave', 'pyramid_stairs'].includes(
String(d.terrainPreset),
),
'地形预设',
);
const params = record(d.terrainParams);
for (const value of Object.values(params)) number(value, 0, 100);
equal(params.size, t.size, '地图尺寸');
equal(params.friction, t.friction, '摩擦');
}
} else {
requireValue(d.terrainPreset !== 'custom_boxes', '缺失custom_boxes布局');
}
if (d.taskId === OBSTACLE_TASK_ID) {
requireValue(d.terrain, '避障地图');
const s = record(d.sensorCfg);
equal(s.type, 'raycast', '传感器');
const pattern = obstacleSensorPattern(s.sensorMode, s.fov);
const allowed = new Set([
'type',
'fov',
'maxDistance',
'safetyDistance',
'avoidanceWeight',
'offset',
'alignment',
'terrainOnly',
'includeGround',
...Object.keys(pattern),
]);
requireValue(
Object.keys(s).every((key) => allowed.has(key)),
'传感器未知字段',
);
for (const [key, expected] of Object.entries(pattern)) {
if (s[key] === undefined) continue;
if (key === 'yawAngles') {
vector(s[key], (expected as number[]).length, -60, 60);
requireValue(
(s[key] as number[]).every((v, i) => Math.abs(v - (expected as number[])[i]) <= 1e-10),
'yawAngles与FOV矛盾',
);
} else equal(s[key], expected, `sensorCfg.${key}`);
}
equal(s.rayCount, pattern.rayCount, '射线数量');
equal(d.observationSize, 49 + pattern.rayCount, '观测维数');
d.sensorCfg = { ...s, ...pattern };
equal(s.offset, [0.3, 0, 0.05], '射线偏移');
equal(s.alignment, 'base', '射线姿态');
equal(s.terrainOnly, true, '仅地形');
equal(s.includeGround, true, '包含地面');
number(s.fov, 30, 120);
number(s.maxDistance, 1, 5);
number(s.safetyDistance, 0.1, 1);
number(s.avoidanceWeight, 0, 10);
requireValue(s.safetyDistance < s.maxDistance, '安全距离');
const n = record(d.navigation);
number(n.speed, 0.3, 1.2);
for (const [key, expected] of Object.entries({
arrivalRadius: 0.5,
distanceScale: record(d.terrain).size,
headingScale: Math.PI,
yawGain: 1,
maxYawRate: 1,
episodeSeconds: 20,
onArrival: 'stop',
onReset: 'respawn',
}))
equal(n[key], expected, `导航 ${key}`);
}
return structuredClone(d) as unknown as PolicyDeployment;
}
/** Read only ModelProto.metadata_props (field 14); skip graph/tensors without decoding them. */
export function readPolicyDeployment(bytes: Uint8Array): PolicyDeployment | undefined {
requireValue(bytes.length <= 64 * 1024 * 1024, 'ONNX超过64MiB');
const entries: Uint8Array[] = [];
function fields(data: Uint8Array, visit: (field: number, value: Uint8Array) => void) {
let p = 0;
const varint = () => {
let value = 0;
for (let i = 0; i < 10; i++) {
requireValue(p < data.length, 'ONNX截断');
const b = data[p++];
value += (b & 127) * 2 ** (i * 7);
requireValue(Number.isSafeInteger(value), 'ONNX整数');
if (!(b & 128)) return value;
}
throw new Error('无效ONNX varint');
};
while (p < data.length) {
const tag = varint(),
field = Math.floor(tag / 8),
wire = tag % 8;
requireValue(field > 0, 'ONNX字段');
if (wire === 0) {
varint();
continue;
}
const length = wire === 2 ? varint() : wire === 1 ? 8 : wire === 5 ? 4 : -1;
requireValue(length >= 0 && p + length <= data.length, 'ONNX字段长度');
if (wire === 2) visit(field, data.subarray(p, p + length));
p += length;
}
}
fields(bytes, (field, value) => {
if (field === 14) {
requireValue(value.length < 100_000 && entries.length < 100, 'ONNX metadata大小');
entries.push(value);
}
});
let deployment: PolicyDeployment | undefined;
const decoder = new TextDecoder('utf-8', { fatal: true });
for (const entry of entries) {
let key = '',
value = '';
fields(entry, (field, bytes) => {
if (field === 1) key = decoder.decode(bytes);
if (field === 2) value = decoder.decode(bytes);
});
if (key === 'platform_deployment') {
requireValue(!deployment, '重复deployment');
deployment = validatePolicyDeployment(JSON.parse(value));
}
}
return deployment;
}
/** Only the original non-custom Flat job may import older external-trainer ONNX without metadata.
* Runtime.load still validates the real graph, not this JSON declaration. */
export function resolvePolicyDeployment(
bytes: Uint8Array,
expected?: PolicyDeployment,
): PolicyDeployment | undefined {
const embedded = readPolicyDeployment(bytes);
if (!expected) return embedded;
const validated = validatePolicyDeployment(expected);
if (embedded && policyDeploymentsMatch(validated, embedded)) return embedded;
const legacyFlat =
validated.taskId === 'Unitree-Go2-Flat' &&
validated.terrain === undefined &&
validated.terrainPreset === undefined &&
validated.terrainParams === undefined &&
validated.sensorCfg === undefined &&
validated.navigation === undefined;
if (!embedded && legacyFlat) return validated;
throw new Error('下载的策略与训练作业部署配置不一致');
}
@@ -0,0 +1,279 @@
{
"source": "Entity(get_go2_robot_cfg()).compile(); go2_constants.py FULL_COLLISION and armature. Foot friction is overridden by custom terrain startup.",
"joints": {
"floating_base_joint": {
"dof_armature": 0.0,
"dof_damping": 0.0,
"dof_frictionloss": 0.0
},
"FL_hip_joint": {
"dof_armature": 0.01,
"dof_damping": 0.0,
"dof_frictionloss": 0.0
},
"FL_thigh_joint": {
"dof_armature": 0.01,
"dof_damping": 0.0,
"dof_frictionloss": 0.0
},
"FL_calf_joint": {
"dof_armature": 0.02,
"dof_damping": 0.0,
"dof_frictionloss": 0.0
},
"FR_hip_joint": {
"dof_armature": 0.01,
"dof_damping": 0.0,
"dof_frictionloss": 0.0
},
"FR_thigh_joint": {
"dof_armature": 0.01,
"dof_damping": 0.0,
"dof_frictionloss": 0.0
},
"FR_calf_joint": {
"dof_armature": 0.02,
"dof_damping": 0.0,
"dof_frictionloss": 0.0
},
"RL_hip_joint": {
"dof_armature": 0.01,
"dof_damping": 0.0,
"dof_frictionloss": 0.0
},
"RL_thigh_joint": {
"dof_armature": 0.01,
"dof_damping": 0.0,
"dof_frictionloss": 0.0
},
"RL_calf_joint": {
"dof_armature": 0.02,
"dof_damping": 0.0,
"dof_frictionloss": 0.0
},
"RR_hip_joint": {
"dof_armature": 0.01,
"dof_damping": 0.0,
"dof_frictionloss": 0.0
},
"RR_thigh_joint": {
"dof_armature": 0.01,
"dof_damping": 0.0,
"dof_frictionloss": 0.0
},
"RR_calf_joint": {
"dof_armature": 0.02,
"dof_damping": 0.0,
"dof_frictionloss": 0.0
}
},
"geoms": {
"base1_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"base2_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"base3_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"FL_hip_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"FL_thigh_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"FL_calf1_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"FL_calf2_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"FL_foot_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 3,
"geom_priority": 1,
"geom_solimp": [0.9, 0.95, 0.023, 0.5, 2.0],
"geom_friction": [0.6, 0.005, 0.0001],
"geom_group": 3
},
"FR_hip_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"FR_thigh_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"FR_calf1_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"FR_calf2_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"FR_foot_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 3,
"geom_priority": 1,
"geom_solimp": [0.9, 0.95, 0.023, 0.5, 2.0],
"geom_friction": [0.6, 0.005, 0.0001],
"geom_group": 3
},
"RL_hip_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"RL_thigh_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"RL_calf1_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"RL_calf2_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"RL_foot_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 3,
"geom_priority": 1,
"geom_solimp": [0.9, 0.95, 0.023, 0.5, 2.0],
"geom_friction": [0.6, 0.005, 0.0001],
"geom_group": 3
},
"RR_hip_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"RR_thigh_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"RR_calf1_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"RR_calf2_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 1,
"geom_priority": 0,
"geom_solimp": [0.9, 0.95, 0.001, 0.5, 2.0],
"geom_friction": [1.0, 0.005, 0.0001],
"geom_group": 3
},
"RR_foot_collision": {
"geom_contype": 1,
"geom_conaffinity": 0,
"geom_condim": 3,
"geom_priority": 1,
"geom_solimp": [0.9, 0.95, 0.023, 0.5, 2.0],
"geom_friction": [0.6, 0.005, 0.0001],
"geom_group": 3
}
}
}
@@ -0,0 +1,221 @@
{
"version": 1,
"taskId": "Unitree-Go2-ObstacleAvoidance",
"browserCompatible": true,
"observationSize": 97,
"actionSize": 12,
"controlHz": 50,
"gaitPeriod": 0.6,
"jointNames": [
"FL_hip_joint",
"FL_thigh_joint",
"FL_calf_joint",
"FR_hip_joint",
"FR_thigh_joint",
"FR_calf_joint",
"RL_hip_joint",
"RL_thigh_joint",
"RL_calf_joint",
"RR_hip_joint",
"RR_thigh_joint",
"RR_calf_joint"
],
"defaultJointPosition": [-0.1, 0.9, -1.8, 0.1, 0.9, -1.8, -0.1, 0.9, -1.8, 0.1, 0.9, -1.8],
"actionScale": [0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25],
"stiffness": [20, 20, 40, 20, 20, 40, 20, 20, 40, 20, 20, 40],
"damping": [1, 1, 2, 1, 1, 2, 1, 1, 2, 1, 1, 2],
"effortLimits": [23.5, 23.5, 45, 23.5, 23.5, 45, 23.5, 23.5, 45, 23.5, 23.5, 45],
"observationTerms": [
"base_ang_vel",
"projected_gravity",
"command",
"phase",
"joint_pos",
"joint_vel",
"actions",
"forward_depth",
"target_error"
],
"seed": 42,
"terrainPreset": "discrete_obstacles",
"terrainParams": {
"size": 12,
"obstacle_count": 24,
"obstacle_height_min": 0.2,
"obstacle_height_max": 0.6,
"spacing": 1.2,
"friction": 0.8,
"roughness": 0.06,
"step_height": 0.08,
"wave_amplitude": 0.08
},
"terrain": {
"representation": "boxes-v1",
"approximation": false,
"size": 12,
"friction": 0.8,
"boxes": [
{
"pos": [0, 0, -0.1],
"size": [6.0, 6.0, 0.1],
"yaw": 0
},
{
"pos": [0.666666666666667, 1.875, 0.261425654654876],
"size": [0.2, 0.2, 0.261425654654876],
"yaw": 0
},
{
"pos": [-0.6666666666666665, 4.375, 0.24594635733876358],
"size": [0.2, 0.2, 0.24594635733876358],
"yaw": 0
},
{
"pos": [0.666666666666667, -4.375, 0.20724561829094013],
"size": [0.2, 0.2, 0.20724561829094013],
"yaw": 0
},
{
"pos": [-2.0, -1.875, 0.29462315279587414],
"size": [0.2, 0.2, 0.29462315279587414],
"yaw": 0
},
{
"pos": [-0.6666666666666665, 1.875, 0.17570687544167068],
"size": [0.2, 0.2, 0.17570687544167068],
"yaw": 0
},
{
"pos": [-2.0, -3.125, 0.2104081262546454],
"size": [0.2, 0.2, 0.2104081262546454],
"yaw": 0
},
{
"pos": [0.666666666666667, 3.125, 0.265880932850599],
"size": [0.2, 0.2, 0.265880932850599],
"yaw": 0
},
{
"pos": [2.0, 4.375, 0.2237039504728492],
"size": [0.2, 0.2, 0.2237039504728492],
"yaw": 0
},
{
"pos": [-2.0, 0.625, 0.27234138006215547],
"size": [0.2, 0.2, 0.27234138006215547],
"yaw": 0
},
{
"pos": [-3.3333333333333335, -0.625, 0.21547042905135239],
"size": [0.2, 0.2, 0.21547042905135239],
"yaw": 0
},
{
"pos": [2.0, -0.625, 0.24091436724298468],
"size": [0.2, 0.2, 0.24091436724298468],
"yaw": 0
},
{
"pos": [-3.3333333333333335, 0.625, 0.10916487673113245],
"size": [0.2, 0.2, 0.10916487673113245],
"yaw": 0
},
{
"pos": [3.333333333333333, -0.625, 0.14557965513030938],
"size": [0.2, 0.2, 0.14557965513030938],
"yaw": 0
},
{
"pos": [0.666666666666667, -1.875, 0.15787759272042143],
"size": [0.2, 0.2, 0.15787759272042143],
"yaw": 0
},
{
"pos": [-2.0, -0.625, 0.11595839538472551],
"size": [0.2, 0.2, 0.11595839538472551],
"yaw": 0
},
{
"pos": [2.0, -3.125, 0.14655817727220605],
"size": [0.2, 0.2, 0.14655817727220605],
"yaw": 0
},
{
"pos": [0.666666666666667, 0.625, 0.12020028588194583],
"size": [0.2, 0.2, 0.12020028588194583],
"yaw": 0
},
{
"pos": [3.333333333333333, -3.125, 0.15559472062201843],
"size": [0.2, 0.2, 0.15559472062201843],
"yaw": 0
},
{
"pos": [-0.6666666666666665, -1.875, 0.22713688885288003],
"size": [0.2, 0.2, 0.22713688885288003],
"yaw": 0
},
{
"pos": [0.666666666666667, 4.375, 0.17296643579401685],
"size": [0.2, 0.2, 0.17296643579401685],
"yaw": 0
},
{
"pos": [3.333333333333333, 3.125, 0.17403619342337653],
"size": [0.2, 0.2, 0.17403619342337653],
"yaw": 0
},
{
"pos": [2.0, -4.375, 0.14190140615429753],
"size": [0.2, 0.2, 0.14190140615429753],
"yaw": 0
},
{
"pos": [2.0, 0.625, 0.15339556440982266],
"size": [0.2, 0.2, 0.15339556440982266],
"yaw": 0
},
{
"pos": [0.666666666666667, -3.125, 0.2873309175424988],
"size": [0.2, 0.2, 0.2873309175424988],
"yaw": 0
}
],
"spawn": [-5.0, 0, 0.32],
"spawnQuaternion": [1, 0, 0, 0],
"target": [5.0, 0],
"actualObstacleCount": 24
},
"sensorCfg": {
"fov": 90,
"maxDistance": 4,
"safetyDistance": 0.5,
"avoidanceWeight": 2,
"type": "raycast",
"sensorMode": "multi_ring_raycast",
"rayCount": 48,
"pitchAngles": [0, -20, -45],
"yawCount": 16,
"yawAngles": [
-45.0, -39.0, -33.0, -27.0, -21.0, -15.0, -9.0, -3.0, 3.0, 9.0, 15.0, 21.0, 27.0, 33.0, 39.0,
45.0
],
"angleUnit": "deg",
"rayOrder": "layer-major",
"offset": [0.3, 0, 0.05],
"alignment": "base",
"terrainOnly": true,
"includeGround": true
},
"navigation": {
"speed": 0.6,
"arrivalRadius": 0.5,
"distanceScale": 12,
"headingScale": 3.141592653589793,
"yawGain": 1.0,
"maxYawRate": 1.0,
"episodeSeconds": 20,
"onArrival": "stop",
"onReset": "respawn"
}
}
@@ -0,0 +1,336 @@
{
"source": "CPU MuJoCo 3.5.0 mj_ray",
"layouts": {
"default": [
{
"pos": [0, 0, -0.1],
"size": [6.0, 6.0, 0.1],
"yaw": 0
},
{
"pos": [0.666666666666667, 1.875, 0.261425654654876],
"size": [0.2, 0.2, 0.261425654654876],
"yaw": 0
},
{
"pos": [-0.6666666666666665, 4.375, 0.24594635733876358],
"size": [0.2, 0.2, 0.24594635733876358],
"yaw": 0
},
{
"pos": [0.666666666666667, -4.375, 0.20724561829094013],
"size": [0.2, 0.2, 0.20724561829094013],
"yaw": 0
},
{
"pos": [-2.0, -1.875, 0.29462315279587414],
"size": [0.2, 0.2, 0.29462315279587414],
"yaw": 0
},
{
"pos": [-0.6666666666666665, 1.875, 0.17570687544167068],
"size": [0.2, 0.2, 0.17570687544167068],
"yaw": 0
},
{
"pos": [-2.0, -3.125, 0.2104081262546454],
"size": [0.2, 0.2, 0.2104081262546454],
"yaw": 0
},
{
"pos": [0.666666666666667, 3.125, 0.265880932850599],
"size": [0.2, 0.2, 0.265880932850599],
"yaw": 0
},
{
"pos": [2.0, 4.375, 0.2237039504728492],
"size": [0.2, 0.2, 0.2237039504728492],
"yaw": 0
},
{
"pos": [-2.0, 0.625, 0.27234138006215547],
"size": [0.2, 0.2, 0.27234138006215547],
"yaw": 0
},
{
"pos": [-3.3333333333333335, -0.625, 0.21547042905135239],
"size": [0.2, 0.2, 0.21547042905135239],
"yaw": 0
},
{
"pos": [2.0, -0.625, 0.24091436724298468],
"size": [0.2, 0.2, 0.24091436724298468],
"yaw": 0
},
{
"pos": [-3.3333333333333335, 0.625, 0.10916487673113245],
"size": [0.2, 0.2, 0.10916487673113245],
"yaw": 0
},
{
"pos": [3.333333333333333, -0.625, 0.14557965513030938],
"size": [0.2, 0.2, 0.14557965513030938],
"yaw": 0
},
{
"pos": [0.666666666666667, -1.875, 0.15787759272042143],
"size": [0.2, 0.2, 0.15787759272042143],
"yaw": 0
},
{
"pos": [-2.0, -0.625, 0.11595839538472551],
"size": [0.2, 0.2, 0.11595839538472551],
"yaw": 0
},
{
"pos": [2.0, -3.125, 0.14655817727220605],
"size": [0.2, 0.2, 0.14655817727220605],
"yaw": 0
},
{
"pos": [0.666666666666667, 0.625, 0.12020028588194583],
"size": [0.2, 0.2, 0.12020028588194583],
"yaw": 0
},
{
"pos": [3.333333333333333, -3.125, 0.15559472062201843],
"size": [0.2, 0.2, 0.15559472062201843],
"yaw": 0
},
{
"pos": [-0.6666666666666665, -1.875, 0.22713688885288003],
"size": [0.2, 0.2, 0.22713688885288003],
"yaw": 0
},
{
"pos": [0.666666666666667, 4.375, 0.17296643579401685],
"size": [0.2, 0.2, 0.17296643579401685],
"yaw": 0
},
{
"pos": [3.333333333333333, 3.125, 0.17403619342337653],
"size": [0.2, 0.2, 0.17403619342337653],
"yaw": 0
},
{
"pos": [2.0, -4.375, 0.14190140615429753],
"size": [0.2, 0.2, 0.14190140615429753],
"yaw": 0
},
{
"pos": [2.0, 0.625, 0.15339556440982266],
"size": [0.2, 0.2, 0.15339556440982266],
"yaw": 0
},
{
"pos": [0.666666666666667, -3.125, 0.2873309175424988],
"size": [0.2, 0.2, 0.2873309175424988],
"yaw": 0
}
],
"low": [
{
"pos": [0, 0, -0.1],
"size": [6.0, 6.0, 0.1],
"yaw": 0
},
{
"pos": [0.9, 0, 0.025],
"size": [0.2, 0.5, 0.025],
"yaw": 0
}
],
"edge": [
{
"pos": [0, 0, -0.1],
"size": [6.0, 6.0, 0.1],
"yaw": 0
}
]
},
"cases": [
{
"layout": "default",
"position": [-5, 0, 0.32],
"quaternion": [1.0, 0.0, 0.0, 0.0],
"distances": [
-1.0, 3.2168989147329183, 1.3910905083086054, 1.309380610573421, 1.2496691592432005, -1.0,
-1.0, -1.0, -1.0, 2.716792619137356, 2.5881904510252074, 5.534249133791317, -1.0,
7.750361403433658, -1.0, 5.904341622907672, 1.0818076280603424, 1.0818076280603424,
1.0818076280603424, 1.0818076280603424, 1.0818076280603424, 1.0818076280603424,
1.0818076280603424, 1.0818076280603424, 1.0818076280603424, 1.0818076280603424,
1.0818076280603424, 1.0818076280603424, 1.0818076280603424, 1.0818076280603424,
1.0818076280603424, 1.0818076280603424, 0.5232590180780452, 0.5232590180780452,
0.5232590180780452, 0.5232590180780452, 0.5232590180780452, 0.5232590180780452,
0.5232590180780452, 0.5232590180780452, 0.5232590180780452, 0.5232590180780452,
0.5232590180780452, 0.5232590180780452, 0.5232590180780452, 0.5232590180780452,
0.5232590180780452, 0.5232590180780452
],
"hitIds": [
-1, 4, 10, 10, 10, -1, -1, -1, -1, 9, 9, 1, -1, 8, -1, 2, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0
],
"depth": [
1, 0.8042247286832296, 0.34777262707715134, 0.32734515264335523, 0.31241728981080014, 1, 1,
1, 1, 0.679198154784339, 0.6470476127563018, 1, 1, 1, 1, 1, 0.2704519070150856,
0.2704519070150856, 0.2704519070150856, 0.2704519070150856, 0.2704519070150856,
0.2704519070150856, 0.2704519070150856, 0.2704519070150856, 0.2704519070150856,
0.2704519070150856, 0.2704519070150856, 0.2704519070150856, 0.2704519070150856,
0.2704519070150856, 0.2704519070150856, 0.2704519070150856, 0.1308147545195113,
0.1308147545195113, 0.1308147545195113, 0.1308147545195113, 0.1308147545195113,
0.1308147545195113, 0.1308147545195113, 0.1308147545195113, 0.1308147545195113,
0.1308147545195113, 0.1308147545195113, 0.1308147545195113, 0.1308147545195113,
0.1308147545195113, 0.1308147545195113, 0.1308147545195113
]
},
{
"layout": "default",
"position": [-3, 1, 0.42],
"quaternion": [
0.9542159228792829, 0.07102055199624979, -0.17699854315088892, 0.23043343819907972
],
"distances": [
-1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0,
-1.0, 5.4074561027806505, 4.629641840762559, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0,
-1.0, -1.0, -1.0, -1.0, -1.0, -1.0, 1.1610973299449852, 1.2130908885760254,
1.2644878853056216, 1.313884872400026, 1.3596758729124838, 1.400132832042811,
1.4335269836666904, 1.4582821425665438, 1.4731397038952894, 1.4773064832243956,
1.4705548683639014, 1.4532525276398263, 1.426314634300804, 1.3910898658423474,
1.349205630860882, 1.3024034756284684
],
"hitIds": [
-1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, 0, 23, -1, -1, -1, -1, -1,
-1, -1, -1, -1, -1, -1, -1, -1, -1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0
],
"depth": [
1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1,
1, 0.2902743324862463, 0.30327272214400636, 0.3161219713264054, 0.3284712181000065,
0.33991896822812095, 0.35003320801070276, 0.3583817459166726, 0.36457053564163594,
0.36828492597382234, 0.3693266208060989, 0.36763871709097534, 0.36331313190995657,
0.356578658575201, 0.34777246646058685, 0.3373014077152205, 0.3256008689071171
]
},
{
"layout": "default",
"position": [-2, -2, 0.5],
"quaternion": [
0.920495563973526, -0.20509659844134395, 0.07256269404109637, -0.32458890530384493
],
"distances": [
-1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, 4.141121232325999, 2.46859904634721,
2.3854692510759445, 2.821623442591415, 2.353034357115619, 2.037057192180342,
1.813411398667459, -1.0, -1.0, -1.0, -1.0, 3.2639344664632044, 2.6365745703702133,
2.2014351861213233, 1.8851396055702525, 1.647197681451036, 1.463506304912354,
1.3188672797050087, 1.2032560687123115, 1.1098187807096116, 1.0337317580203136,
0.9715203831532319, 0.9206357398775848, 0.8370986440649361, 1.207568318572572,
1.1431256747550447, 1.081361321440165, 1.0230591974383043, 0.9686938988212619,
0.9185003934090735, 0.8725376029719634, 0.8307419413609812, 0.7929698985660093,
0.7590304425905215, 0.7287087687688034, 0.7017831189893099, 0.6780362858286884,
0.657263179268657, 0.6392755644084297
],
"hitIds": [
-1, -1, -1, -1, -1, -1, -1, -1, -1, 22, 24, 24, 0, 0, 0, 0, -1, -1, -1, -1, 0, 0, 0, 0, 0,
0, 0, 0, 0, 0, 0, 0, 6, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0
],
"depth": [
1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 0.6171497615868025, 0.5963673127689861, 0.7054058606478537,
0.5882585892789047, 0.5092642980450856, 0.45335284966686473, 1, 1, 1, 1, 0.8159836166158011,
0.6591436425925533, 0.5503587965303308, 0.4712849013925631, 0.411799420362759,
0.3658765762280885, 0.3297168199262522, 0.3008140171780779, 0.2774546951774029,
0.2584329395050784, 0.24288009578830796, 0.2301589349693962, 0.20927466101623401,
0.301892079643143, 0.28578141868876117, 0.27034033036004124, 0.2557647993595761,
0.24217347470531547, 0.22962509835226838, 0.21813440074299084, 0.2076854853402453,
0.19824247464150233, 0.1897576106476304, 0.18217719219220085, 0.17544577974732747,
0.1695090714571721, 0.16431579481716424, 0.15981889110210742
]
},
{
"layout": "low",
"position": [0, 0, 0.32],
"quaternion": [1.0, 0.0, 0.0, 0.0],
"distances": [
-1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0,
-1.0, 1.0818076280603424, 1.0818076280603424, 0.9356174080521878, 0.9356174080521878,
1.0818076280603424, 1.0818076280603424, 1.0818076280603424, 1.0818076280603424,
1.0818076280603424, 1.0818076280603424, 1.0818076280603424, 1.0818076280603424,
0.9356174080521878, 0.9356174080521878, 1.0818076280603424, 1.0818076280603424,
0.5232590180780452, 0.5232590180780452, 0.5232590180780452, 0.5232590180780452,
0.5232590180780452, 0.5232590180780452, 0.5232590180780452, 0.5232590180780452,
0.5232590180780452, 0.5232590180780452, 0.5232590180780452, 0.5232590180780452,
0.5232590180780452, 0.5232590180780452, 0.5232590180780452, 0.5232590180780452
],
"hitIds": [
-1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, 0, 0, 1, 1, 0, 0, 0, 0, 0,
0, 0, 0, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0
],
"depth": [
1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 0.2704519070150856, 0.2704519070150856,
0.23390435201304696, 0.23390435201304696, 0.2704519070150856, 0.2704519070150856,
0.2704519070150856, 0.2704519070150856, 0.2704519070150856, 0.2704519070150856,
0.2704519070150856, 0.2704519070150856, 0.23390435201304696, 0.23390435201304696,
0.2704519070150856, 0.2704519070150856, 0.1308147545195113, 0.1308147545195113,
0.1308147545195113, 0.1308147545195113, 0.1308147545195113, 0.1308147545195113,
0.1308147545195113, 0.1308147545195113, 0.1308147545195113, 0.1308147545195113,
0.1308147545195113, 0.1308147545195113, 0.1308147545195113, 0.1308147545195113,
0.1308147545195113, 0.1308147545195113
]
},
{
"layout": "low",
"position": [0, 0, 0.32],
"quaternion": [
0.9927681027239722, 0.09709860284347692, 0.054650688238454904, -0.04468397715907572
],
"distances": [
1.5463415120082626, 1.593685768444766, 1.6628142035077924, 1.7583500416114888,
1.8874730476566093, 2.06143867460523, 2.298467554987033, 2.6296404516571896,
3.112149423056716, 3.863359519926606, 5.167240524912454, -1.0, -1.0, -1.0, -1.0, -1.0,
0.62209731944295, 0.629163205894759, 0.5724502740761562, 0.5541151601654627,
0.567642029183309, 0.5840261942036087, 0.6035165215332615, 0.626413352751139,
0.6530744139420617, 0.6839211881333935, 0.7194452948914949, 0.7602140509721573,
0.8068737942201454, 0.8601486324595742, 1.0831765236086466, 1.1642551456388435,
0.3961557091780092, 0.39829918966870376, 0.4012471161917819, 0.40500178050690006,
0.4095651162997735, 0.4149379498546247, 0.4211190461681518, 0.4281039291887044,
0.4358834560016674, 0.44444212949443834, 0.4537561437057185, 0.4637911723119031,
0.4744999352392512, 0.48581961273489255, 0.4976692212346268, 0.5099471205068027
],
"hitIds": [
0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, -1, -1, -1, -1, -1, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1,
1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0
],
"depth": [
0.38658537800206566, 0.3984214421111915, 0.4157035508769481, 0.4395875104028722,
0.4718682619141523, 0.5153596686513074, 0.5746168887467582, 0.6574101129142974,
0.778037355764179, 0.9658398799816516, 1, 1, 1, 1, 1, 1, 0.1555243298607375,
0.15729080147368976, 0.14311256851903906, 0.13852879004136567, 0.14191050729582724,
0.14600654855090217, 0.15087913038331538, 0.15660333818778474, 0.16326860348551542,
0.17098029703334838, 0.17986132372287372, 0.19005351274303933, 0.20171844855503634,
0.21503715811489355, 0.27079413090216164, 0.29106378640971087, 0.0990389272945023,
0.09957479741717594, 0.10031177904794547, 0.10125044512672501, 0.10239127907494337,
0.10373448746365617, 0.10527976154203796, 0.1070259822971761, 0.10897086400041685,
0.11111053237360959, 0.11343903592642962, 0.11594779307797577, 0.1186249838098128,
0.12145490318372314, 0.1244173053086567, 0.12748678012670067
]
},
{
"layout": "edge",
"position": [5.45, 0, 0.32],
"quaternion": [1.0, 0.0, 0.0, 0.0],
"distances": [
-1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0,
-1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0,
-1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0, -1.0,
-1.0, -1.0, -1.0
],
"hitIds": [
-1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1,
-1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1,
-1, -1
],
"depth": [
1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1,
1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1
]
}
]
}
@@ -0,0 +1,212 @@
{
"version": 1,
"taskId": "Unitree-Go2-ObstacleAvoidance",
"browserCompatible": true,
"observationSize": 81,
"actionSize": 12,
"controlHz": 50,
"gaitPeriod": 0.6,
"jointNames": [
"FL_hip_joint",
"FL_thigh_joint",
"FL_calf_joint",
"FR_hip_joint",
"FR_thigh_joint",
"FR_calf_joint",
"RL_hip_joint",
"RL_thigh_joint",
"RL_calf_joint",
"RR_hip_joint",
"RR_thigh_joint",
"RR_calf_joint"
],
"defaultJointPosition": [-0.1, 0.9, -1.8, 0.1, 0.9, -1.8, -0.1, 0.9, -1.8, 0.1, 0.9, -1.8],
"actionScale": [0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25],
"stiffness": [20, 20, 40, 20, 20, 40, 20, 20, 40, 20, 20, 40],
"damping": [1, 1, 2, 1, 1, 2, 1, 1, 2, 1, 1, 2],
"effortLimits": [23.5, 23.5, 45, 23.5, 23.5, 45, 23.5, 23.5, 45, 23.5, 23.5, 45],
"observationTerms": [
"base_ang_vel",
"projected_gravity",
"command",
"phase",
"joint_pos",
"joint_vel",
"actions",
"forward_depth",
"target_error"
],
"seed": 42,
"terrainPreset": "discrete_obstacles",
"terrainParams": {
"size": 12,
"obstacle_count": 24,
"obstacle_height_min": 0.2,
"obstacle_height_max": 0.6,
"spacing": 1.2,
"friction": 0.8,
"roughness": 0.06,
"step_height": 0.08,
"wave_amplitude": 0.08
},
"terrain": {
"representation": "boxes-v1",
"approximation": false,
"size": 12,
"friction": 0.8,
"boxes": [
{
"pos": [0, 0, -0.1],
"size": [6.0, 6.0, 0.1],
"yaw": 0
},
{
"pos": [0.666666666666667, 1.875, 0.261425654654876],
"size": [0.2, 0.2, 0.261425654654876],
"yaw": 0
},
{
"pos": [-0.6666666666666665, 4.375, 0.24594635733876358],
"size": [0.2, 0.2, 0.24594635733876358],
"yaw": 0
},
{
"pos": [0.666666666666667, -4.375, 0.20724561829094013],
"size": [0.2, 0.2, 0.20724561829094013],
"yaw": 0
},
{
"pos": [-2.0, -1.875, 0.29462315279587414],
"size": [0.2, 0.2, 0.29462315279587414],
"yaw": 0
},
{
"pos": [-0.6666666666666665, 1.875, 0.17570687544167068],
"size": [0.2, 0.2, 0.17570687544167068],
"yaw": 0
},
{
"pos": [-2.0, -3.125, 0.2104081262546454],
"size": [0.2, 0.2, 0.2104081262546454],
"yaw": 0
},
{
"pos": [0.666666666666667, 3.125, 0.265880932850599],
"size": [0.2, 0.2, 0.265880932850599],
"yaw": 0
},
{
"pos": [2.0, 4.375, 0.2237039504728492],
"size": [0.2, 0.2, 0.2237039504728492],
"yaw": 0
},
{
"pos": [-2.0, 0.625, 0.27234138006215547],
"size": [0.2, 0.2, 0.27234138006215547],
"yaw": 0
},
{
"pos": [-3.3333333333333335, -0.625, 0.21547042905135239],
"size": [0.2, 0.2, 0.21547042905135239],
"yaw": 0
},
{
"pos": [2.0, -0.625, 0.24091436724298468],
"size": [0.2, 0.2, 0.24091436724298468],
"yaw": 0
},
{
"pos": [-3.3333333333333335, 0.625, 0.10916487673113245],
"size": [0.2, 0.2, 0.10916487673113245],
"yaw": 0
},
{
"pos": [3.333333333333333, -0.625, 0.14557965513030938],
"size": [0.2, 0.2, 0.14557965513030938],
"yaw": 0
},
{
"pos": [0.666666666666667, -1.875, 0.15787759272042143],
"size": [0.2, 0.2, 0.15787759272042143],
"yaw": 0
},
{
"pos": [-2.0, -0.625, 0.11595839538472551],
"size": [0.2, 0.2, 0.11595839538472551],
"yaw": 0
},
{
"pos": [2.0, -3.125, 0.14655817727220605],
"size": [0.2, 0.2, 0.14655817727220605],
"yaw": 0
},
{
"pos": [0.666666666666667, 0.625, 0.12020028588194583],
"size": [0.2, 0.2, 0.12020028588194583],
"yaw": 0
},
{
"pos": [3.333333333333333, -3.125, 0.15559472062201843],
"size": [0.2, 0.2, 0.15559472062201843],
"yaw": 0
},
{
"pos": [-0.6666666666666665, -1.875, 0.22713688885288003],
"size": [0.2, 0.2, 0.22713688885288003],
"yaw": 0
},
{
"pos": [0.666666666666667, 4.375, 0.17296643579401685],
"size": [0.2, 0.2, 0.17296643579401685],
"yaw": 0
},
{
"pos": [3.333333333333333, 3.125, 0.17403619342337653],
"size": [0.2, 0.2, 0.17403619342337653],
"yaw": 0
},
{
"pos": [2.0, -4.375, 0.14190140615429753],
"size": [0.2, 0.2, 0.14190140615429753],
"yaw": 0
},
{
"pos": [2.0, 0.625, 0.15339556440982266],
"size": [0.2, 0.2, 0.15339556440982266],
"yaw": 0
},
{
"pos": [0.666666666666667, -3.125, 0.2873309175424988],
"size": [0.2, 0.2, 0.2873309175424988],
"yaw": 0
}
],
"spawn": [-5.0, 0, 0.32],
"spawnQuaternion": [1, 0, 0, 0],
"target": [5.0, 0],
"actualObstacleCount": 24
},
"sensorCfg": {
"fov": 90,
"maxDistance": 4,
"safetyDistance": 0.5,
"avoidanceWeight": 2,
"type": "raycast",
"rayCount": 32,
"offset": [0.3, 0, 0.05],
"alignment": "base",
"terrainOnly": true,
"includeGround": true
},
"navigation": {
"speed": 0.6,
"arrivalRadius": 0.5,
"distanceScale": 12,
"headingScale": 3.141592653589793,
"yawGain": 1.0,
"maxYawRate": 1.0,
"episodeSeconds": 20,
"onArrival": "stop",
"onReset": "respawn"
}
}
@@ -0,0 +1,90 @@
{
"source": "CPU MuJoCo 3.5.0 mj_ray; task_config seed42 discrete_obstacles; distances normalized by4",
"cases": [
{
"position": [-5, 0, 0.32],
"quaternion": [1, 0, 0, 0],
"depth": [
1, 1, 0.8064353265688073, 0.7754070525042633, 0.3493131891686196, 0.33844992227063814,
0.32906116691883763, 0.3209809620250947, 0.31407497396875733, 0.3285019206818616,
0.3862286490933719, 1, 1, 1, 1, 1, 1, 1, 1, 0.634959321852795, 0.6416072675585701,
0.650082300931107, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1
]
},
{
"position": [-2, 1, 0.32],
"quaternion": [0.9689124217106447, 0, 0.24740395925452294, 0],
"depth": [
0.16227742543146714, 0.154643351828059, 0.1480582588705743, 0.1423615934562481,
0.13742675650547637, 0.1331529312359637, 0.12945920813695258, 0.12628029481538672,
0.12356334175298786, 0.12126556878831947, 0.11935247679185605, 0.11779649499930897,
0.11657595910035341, 0.11567434601043744, 0.1150797130480744, 0.1147843050751654,
0.1147843050751654, 0.1150797130480744, 0.11567434601043744, 0.11657595910035341,
0.11779649499930897, 0.11935247679185605, 0.12126556878831947, 0.12356334175298787,
0.12628029481538672, 0.12945920813695258, 0.1331529312359637, 0.13742675650547637,
0.14236159345624813, 0.14805825887057428, 0.154643351828059, 0.16227742543146714
]
},
{
"position": [0, -1, 0.32],
"quaternion": [0.955336489125606, 0.29552020666133955, 0, 0],
"depth": [
0.2262087980617112, 0.23859992736340305, 0.2531146297321152, 0.27024831583146797,
0.2906703480610164, 0.31530672670613974, 0.34547502041327477, 0.3831145628697801,
0.431200828086753, 0.49454232783018703, 0.5814468749554897, 0.707609556942636,
0.9066658372041383, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1
]
},
{
"position": [2, 2, 0.32],
"quaternion": [0.8253356149096783, 0, 0, 0.5646424733950354],
"depth": [
1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 0.47425413995403354,
0.4738672755338355, 0.4746985930390542, 0.47675881629202255, 0.48007471719272266, 1, 1, 1,
1, 1
]
}
],
"boxEdges": [
{
"origin": [0, 0, 0],
"direction": [1, 0, 0],
"distance": 1.0
},
{
"origin": [1, 0, 0],
"direction": [1, 0, 0],
"distance": 0.0
},
{
"origin": [1, 0, 0],
"direction": [-1, 0, 0],
"distance": -0.0
},
{
"origin": [1, 0, 0],
"direction": [0, 1, 0],
"distance": 1.0
},
{
"origin": [2, 0, 0],
"direction": [0, 1, 0],
"distance": -1.0
},
{
"origin": [2, 0, 0],
"direction": [-1, 0, 0],
"distance": 1.0
},
{
"origin": [1, 1, 0],
"direction": [0, 0, 1],
"distance": 1.0
},
{
"origin": [2, 2, 0],
"direction": [1, 0, 0],
"distance": -1.0
}
]
}
@@ -0,0 +1,55 @@
import { expect, it, vi } from 'vitest';
import type { MjData, MjModel } from '@mujoco/mujoco';
import fixture from '../fixtures/obstacleDeployment.json';
import { validatePolicyDeployment } from '../deployment';
vi.mock('./Go2wPolicyBindings', () => ({
Go2wPolicyBindings: class {
baseBodyId = 0;
constructor(
public model: MjModel,
public data: MjData,
) {}
clear() {}
observe(
_time: number,
_last: Float32Array,
command: { linearX: number; linearY: number; angularZ: number },
) {
const obs = new Float32Array(47);
obs.set([command.linearX, command.linearY, command.angularZ], 6);
return obs;
}
},
}));
import { Go2ObstacleAvoidanceBindings } from './Go2ObstacleAvoidanceBindings';
it('浏览器连续导航超过20秒仍观察;换点改变command/目标误差,跌倒越界仍明确停止', () => {
const original = validatePolicyDeployment(fixture);
const deployment = { ...original, terrain: { ...original.terrain!, boxes: [] } };
const data = {
time: 0,
xpos: new Float64Array([0, 0, 0.32]),
xquat: new Float64Array([1, 0, 0, 0]),
} as unknown as MjData;
const bindings = new Go2ObstacleAvoidanceBindings(
{ ngeom: 0 } as MjModel,
data,
vi.fn(),
deployment,
);
bindings.setNavigationTarget([3, 0]);
const before = bindings.observe(21, new Float32Array(12));
bindings.setNavigationTarget([0, 3]);
const after = bindings.observe(22, new Float32Array(12));
expect(before[6]).toBeGreaterThan(0);
expect(after[6]).toBeCloseTo(0);
expect(after[8]).toBe(1);
expect(after[79]).not.toBe(before[79]);
data.xpos[2] = 0.1;
expect(() => bindings.observe(23, new Float32Array(12))).toThrow(/跌倒\/越界/);
bindings.setNavigationTarget([2, 0]);
expect(() => bindings.observe(24, new Float32Array(12))).toThrow(/安全停止/);
data.xpos[2] = 0.32;
data.xpos[0] = 100;
expect(() => bindings.observe(25, new Float32Array(12))).toThrow(/安全停止/);
});
@@ -0,0 +1,165 @@
import type { MjModel, MjData } from '@mujoco/mujoco';
import type { NavigationStatus } from '../types';
import type { PolicyDeployment } from '../deployment';
import {
buildGo2ObstacleAvoidanceObservation,
obstacleNavigation,
rayBoxDistance,
sampleForwardRays,
forwardRayDirections,
type PerceptionRay,
} from '../tasks/go2ObstacleAvoidance';
import { Go2wPolicyBindings } from './Go2wPolicyBindings';
/** Only named boxes from the compiled training scene are sensed: never robot visual/collision geoms. */
export class Go2ObstacleAvoidanceBindings extends Go2wPolicyBindings {
private readonly terrainIds: number[] = [];
private readonly boxes: {
center: number[];
half: number[];
rotation: number[];
axisAligned: boolean;
}[] = [];
private readonly directions: number[][];
private startTime: number;
private readonly defaultTarget: [number, number];
private currentTarget: [number, number];
rays: PerceptionRay[] = [];
constructor(
model: MjModel,
data: MjData,
setActuator: (id: number, value: number) => void,
private readonly deployment: PolicyDeployment,
) {
super(model, data, setActuator, deployment.effortLimits);
if (!deployment.terrain || !deployment.sensorCfg)
throw new Error('避障策略缺少地图/传感器配置');
this.directions = forwardRayDirections(deployment.sensorCfg);
this.defaultTarget = [deployment.terrain.target[0], deployment.terrain.target[1]];
this.currentTarget = [...this.defaultTarget];
this.startTime = Number(data.time);
for (let id = 0; id < model.ngeom; id++) {
const geom = model.geom(id);
try {
if (!geom.name.startsWith('__training_terrain_')) continue;
if (
Number(model.geom_type[id]) !== 6 ||
Number(model.geom_bodyid[id]) !== 0 ||
Number(model.geom_group[id]) !== 2
)
throw new Error('训练地图必须是静态group2 box');
this.terrainIds.push(id);
this.boxes.push({
center: Array.from(data.geom_xpos.subarray(id * 3, id * 3 + 3), Number),
half: Array.from(model.geom_size.subarray(id * 3, id * 3 + 3), Number),
rotation: Array.from(data.geom_xmat.subarray(id * 9, id * 9 + 9), Number),
axisAligned: Array.from(data.geom_xmat.subarray(id * 9, id * 9 + 9), Number).every(
(v, i) => v === (i % 4 === 0 ? 1 : 0),
),
});
} finally {
geom.delete();
}
}
if (this.terrainIds.length !== deployment.terrain.boxes.length)
throw new Error('请先加载策略配套训练地图');
}
setNavigationTarget(target: [number, number]): void {
if (target.length !== 2 || !target.every(Number.isFinite))
throw new Error('导航目标必须是两个有限坐标');
const limit = this.deployment.terrain!.size / 2 - 0.5;
this.currentTarget = target.map((v) => Math.max(-limit, Math.min(limit, v))) as [
number,
number,
];
}
resetNavigationTarget(): void {
this.currentTarget = [...this.defaultTarget];
}
navigationStatus(): NavigationStatus {
const [x, y] = this.currentTarget;
let height = 0;
for (const id of this.terrainIds) {
const distance = rayBoxDistance(
[x, y, 10000],
[0, 0, -1],
this.data.geom_xpos.subarray(id * 3, id * 3 + 3),
this.model.geom_size.subarray(id * 3, id * 3 + 3),
this.data.geom_xmat.subarray(id * 9, id * 9 + 9),
);
if (distance >= 0) height = Math.max(height, 10000 - distance);
}
return {
target: [...this.currentTarget],
defaultTarget: [...this.defaultTarget],
targetHeight: height,
distance: Math.hypot(
x - Number(this.data.xpos[this.baseBodyId * 3]),
y - Number(this.data.xpos[this.baseBodyId * 3 + 1]),
),
};
}
reset(time: number) {
this.resetNavigationTarget();
this.startTime = time;
this.rays = [];
}
override clear() {
super.clear();
this.rays = [];
}
override observe(time: number, lastAction: Float32Array): Float32Array {
const terrain = this.deployment.terrain!,
sensor = this.deployment.sensorCfg!;
const position = Array.from(
this.data.xpos.subarray(this.baseBodyId * 3, this.baseBodyId * 3 + 3),
Number,
);
const quaternion = Array.from(
this.data.xquat.subarray(this.baseBodyId * 4, this.baseBodyId * 4 + 4),
Number,
);
// Interactive navigation has no episode deadline. Training/evaluation retain 20s.
// Safety failures still require an explicit reset/re-enable, never a new target.
if (
position[2] < 0.12 ||
Math.abs(position[0]) > terrain.size / 2 - 0.3 ||
Math.abs(position[1]) > terrain.size / 2 - 0.3
)
throw new Error('导航安全停止(跌倒/越界),请重置仿真后重新启用');
const navigation = obstacleNavigation(
position,
quaternion,
this.currentTarget,
terrain.size,
this.deployment.navigation?.speed,
);
const sample = sampleForwardRays(
position,
quaternion,
sensor,
(origin, direction) => {
let nearest = -1;
for (const box of this.boxes) {
const distance = rayBoxDistance(
origin,
direction,
box.center,
box.half,
box.rotation,
box.axisAligned,
);
if (distance >= 0 && (nearest < 0 || distance < nearest)) nearest = distance;
}
return nearest;
},
this.directions,
);
this.rays = sample.rays;
return buildGo2ObstacleAvoidanceObservation(
super.observe(time - this.startTime, lastAction, navigation.command),
sample.depth,
navigation.targetError,
);
}
}
@@ -27,15 +27,16 @@ function rotateInverse(
/** 将 mjlab Go2 velocity 的 47 维 actor 观测和 12 维关节位置动作映射到 MuJoCo。 */
export class Go2wPolicyBindings implements PolicyRuntimeBindings {
private readonly joints: BoundJoint[];
private readonly baseBodyId: number;
protected readonly baseBodyId: number;
private readonly baseFreeJointId: number;
private readonly gyroSensorId?: number;
private readonly wheelActuatorIds: number[];
constructor(
private readonly model: MjModel,
private readonly data: MjData,
protected readonly model: MjModel,
protected readonly data: MjData,
private readonly setActuator: (id: number, value: number) => void,
private readonly effortLimits?: readonly number[],
) {
const jointIds = new Map<string, number>(),
actuatorIds = new Map<string, number>(),
@@ -199,7 +200,11 @@ export class Go2wPolicyBindings implements PolicyRuntimeBindings {
GO2W_VELOCITY_TASK.damping[index] * Number(this.data.qvel[item.qvelAddress]);
this.setActuator(
item.actuatorId,
item.positionActuator ? target : torque / item.controlScale,
item.positionActuator
? target
: (this.effortLimits
? Math.max(-this.effortLimits[index], Math.min(this.effortLimits[index], torque))
: torque) / item.controlScale,
);
}
for (const id of this.wheelActuatorIds) this.setActuator(id, 0);
@@ -0,0 +1,217 @@
import { beforeEach, describe, expect, it, vi } from 'vitest';
import fixture from '../fixtures/obstacleDeployment.json';
import { validatePolicyDeployment } from '../deployment';
const mock = vi.hoisted(() => ({
create: vi.fn(),
tensors: [] as { dispose: ReturnType<typeof vi.fn> }[],
}));
vi.mock('onnxruntime-web/wasm', () => ({
env: { wasm: {} },
InferenceSession: { create: mock.create },
Tensor: class {
dispose = vi.fn();
constructor(
public type: string,
public data: Float32Array,
public dims: number[],
) {
mock.tensors.push(this);
}
},
}));
import { OnnxPolicyRuntime } from './OnnxPolicyRuntime';
const deployment = validatePolicyDeployment(fixture);
const flush = async () => {
for (let i = 0; i < 10; i++) await Promise.resolve();
};
function setup(size = 81) {
const session = {
inputNames: ['obs'],
outputNames: ['actions'],
inputMetadata: [{ isTensor: true, type: 'float32', shape: [1, size] }],
outputMetadata: [{ isTensor: true, type: 'float32', shape: [1, 12] }],
run: vi.fn(),
release: vi.fn().mockResolvedValue(undefined),
};
const bindings = {
observe: vi.fn(() => new Float32Array(size)),
apply: vi.fn(),
clear: vi.fn(),
reset: vi.fn(),
};
mock.create.mockResolvedValue(session);
return { session, bindings };
}
beforeEach(() => {
vi.clearAllMocks();
mock.tensors.length = 0;
});
describe('OnnxPolicyRuntime held-action', () => {
it('动态目标转发不重建ORT、不启用策略或打断in-flight,status实时读取', async () => {
const { session, bindings } = setup();
let finish!: (outputs: Record<string, unknown>) => void;
session.run.mockImplementation(
() =>
new Promise((resolve) => {
finish = resolve;
}),
);
const navigation = {
target: [1, 2] as [number, number],
defaultTarget: [5, 0] as [number, number],
distance: 3,
targetHeight: 0,
};
const targetBindings = {
...bindings,
navigationStatus: () => navigation,
setNavigationTarget: vi.fn(),
resetNavigationTarget: vi.fn(),
};
const runtime = await OnnxPolicyRuntime.load(
new Uint8Array(),
'goal.onnx',
targetBindings,
deployment,
);
runtime.setNavigationTarget([2, 3]);
expect(targetBindings.setNavigationTarget).toHaveBeenCalledWith([2, 3]);
expect(runtime.status().enabled).toBe(false);
expect(runtime.status().navigation).toBe(navigation);
runtime.setEnabled(true, 0);
runtime.step(0);
runtime.setNavigationTarget([3, 4]);
runtime.step(0.02);
expect(session.run).toHaveBeenCalledOnce();
expect(mock.create).toHaveBeenCalledOnce();
runtime.resetNavigationTarget();
expect(targetBindings.resetNavigationTarget).toHaveBeenCalledOnce();
finish({ actions: { type: 'float32', data: new Float32Array(12), dispose: vi.fn() } });
await flush();
runtime.dispose();
expect(runtime.status().navigation).toBeUndefined();
});
it('81维契约单in-flight,持有最新动作,不阻塞物理步;reset拒绝旧promise', async () => {
const { session, bindings } = setup();
let finish!: (outputs: Record<string, unknown>) => void;
session.run.mockImplementation(
() =>
new Promise((resolve) => {
finish = resolve;
}),
);
const runtime = await OnnxPolicyRuntime.load(
new Uint8Array([1]),
'policy.onnx',
bindings,
deployment,
);
runtime.setEnabled(true, 0);
runtime.step(0);
runtime.step(0.02);
runtime.step(0.04);
expect(session.run).toHaveBeenCalledOnce();
expect(bindings.apply).toHaveBeenCalledTimes(3);
const output = { type: 'float32', data: new Float32Array(12).fill(0.3), dispose: vi.fn() };
finish({ actions: output });
await flush();
runtime.step(0.06);
expect(bindings.apply.mock.lastCall?.[0][0]).toBeCloseTo(0.3);
expect(session.run).toHaveBeenCalledTimes(2);
expect(output.dispose).toHaveBeenCalledOnce();
runtime.reset(0);
finish({ actions: { ...output, data: new Float32Array(12).fill(0.9) } });
await flush();
runtime.step(0);
expect(bindings.apply.mock.lastCall?.[0][0]).toBe(0);
expect(bindings.reset).toHaveBeenCalledWith(0);
runtime.dispose();
expect(session.release).not.toHaveBeenCalled();
finish({ actions: output });
await flush();
expect(session.release).toHaveBeenCalledOnce();
expect(mock.tensors.every((t) => t.dispose.mock.calls.length === 1)).toBe(true);
});
it('81策略不能默认为47、Rough维度不能加载,失败释放session', async () => {
const { session, bindings } = setup();
await expect(OnnxPolicyRuntime.load(new Uint8Array(), 'wrong.onnx', bindings)).rejects.toThrow(
/维度/,
);
expect(session.release).toHaveBeenCalledOnce();
session.inputMetadata[0].shape = [1, 234];
await expect(
OnnxPolicyRuntime.load(new Uint8Array(), 'rough.onnx', bindings, deployment),
).rejects.toThrow(/维度/);
});
it('原47维Flat保持向后兼容,非法观测/输出失败关闭并清动作', async () => {
const { session, bindings } = setup(47);
session.run.mockResolvedValue({
actions: { type: 'float32', data: new Float32Array(12).fill(NaN), dispose: vi.fn() },
});
const runtime = await OnnxPolicyRuntime.load(new Uint8Array(), 'flat.onnx', bindings);
expect(runtime.status().observationSize).toBe(47);
runtime.setEnabled(true, 0);
runtime.step(0);
await flush();
expect(runtime.status().enabled).toBe(false);
expect(runtime.status().error).toMatch(/非有限/);
runtime.reset(0);
bindings.observe.mockReturnValue(new Float32Array(81));
runtime.setEnabled(true, 0);
runtime.step(0);
expect(runtime.status().error).toMatch(/维度/);
runtime.dispose();
await flush();
});
});
it('默认Flat作业的兼容契约仍检查真实graph,拒绝81维及动态特征', async () => {
const flat = validatePolicyDeployment({
...fixture,
taskId: 'Unitree-Go2-Flat',
observationSize: 47,
observationTerms: fixture.observationTerms.slice(0, 7),
terrain: undefined,
terrainPreset: undefined,
terrainParams: undefined,
sensorCfg: undefined,
navigation: undefined,
});
const { session, bindings } = setup(47);
const runtime = await OnnxPolicyRuntime.load(new Uint8Array(), 'legacy.onnx', bindings, flat);
expect(runtime.status().observationSize).toBe(47);
runtime.dispose();
await flush();
session.inputMetadata[0].shape = [1, 81];
await expect(
OnnxPolicyRuntime.load(new Uint8Array(), 'legacy-wrong.onnx', bindings, flat),
).rejects.toThrow(/维度/);
session.inputMetadata[0].shape = [1, -1];
await expect(
OnnxPolicyRuntime.load(new Uint8Array(), 'legacy-dynamic.onnx', bindings, flat),
).rejects.toThrow(/固定/);
});
it('97 metadata不能冒充81 graph;真实97 shape沿用single-flight生命周期', async () => {
const multi = validatePolicyDeployment({
...fixture,
observationSize: 97,
sensorCfg: { ...fixture.sensorCfg, sensorMode: 'multi_ring_raycast', rayCount: 48 },
});
const { session, bindings } = setup(81);
await expect(
OnnxPolicyRuntime.load(new Uint8Array([1]), 'multi.onnx', bindings, multi),
).rejects.toThrow(/维度/);
expect(session.release).toHaveBeenCalledOnce();
const next = setup(97);
const runtime = await OnnxPolicyRuntime.load(
new Uint8Array([1]),
'multi.onnx',
next.bindings,
multi,
);
expect(runtime.status().observationSize).toBe(97);
runtime.dispose();
await flush();
expect(next.session.release).toHaveBeenCalledOnce();
});
@@ -1,6 +1,8 @@
import * as ort from 'onnxruntime-web/wasm';
import { GO2W_VELOCITY_TASK, clampGo2wCommand } from '../tasks/go2wVelocity';
import type { RLCommand, RLPolicyStatus } from '../types';
import type { PolicyDeployment } from '../deployment';
import { GO2_OBSTACLE_AVOIDANCE_TASK } from '../tasks/go2ObstacleAvoidance';
import type { NavigationStatus, RLCommand, RLPolicyStatus } from '../types';
ort.env.wasm.numThreads = 1;
ort.env.wasm.proxy = false;
@@ -9,6 +11,10 @@ export interface PolicyRuntimeBindings {
observe(time: number, lastAction: Float32Array, command: RLCommand): Float32Array;
apply(action: Float32Array): void;
clear(): void;
reset?(time: number): void;
navigationStatus?(): NavigationStatus;
setNavigationTarget?(target: [number, number]): void;
resetNavigationTarget?(): void;
}
function message(error: unknown): string {
@@ -38,13 +44,25 @@ export class OnnxPolicyRuntime {
private readonly path: string,
private readonly inputName: string,
private readonly outputName: string,
private readonly task: {
id: string;
name: string;
observationSize: number;
actionSize: number;
controlHz: number;
},
) {}
static async load(
model: Uint8Array,
path: string,
bindings: PolicyRuntimeBindings,
deployment?: PolicyDeployment,
): Promise<OnnxPolicyRuntime> {
const task =
deployment?.taskId === GO2_OBSTACLE_AVOIDANCE_TASK.id
? { ...GO2_OBSTACLE_AVOIDANCE_TASK, observationSize: deployment.observationSize }
: GO2W_VELOCITY_TASK;
const session = await ort.InferenceSession.create(model.slice(), {
executionProviders: ['wasm'],
graphOptimizationLevel: 'all',
@@ -71,28 +89,19 @@ export class OnnxPolicyRuntime {
throw new Error(`策略输入 batch 必须为 1 或动态维度,实际为 ${inputBatch}`);
if (typeof outputBatch === 'number' && outputBatch !== -1 && outputBatch !== 1)
throw new Error(`策略输出 batch 必须为 1 或动态维度,实际为 ${outputBatch}`);
if (
typeof fixedInput === 'number' &&
fixedInput > 0 &&
fixedInput !== GO2W_VELOCITY_TASK.observationSize
)
throw new Error(
`策略观测维度不匹配:模型 ${fixedInput},任务 ${GO2W_VELOCITY_TASK.observationSize}`,
);
if (
typeof fixedOutput === 'number' &&
fixedOutput > 0 &&
fixedOutput !== GO2W_VELOCITY_TASK.actionSize
)
throw new Error(
`策略动作维度不匹配:模型 ${fixedOutput},任务 ${GO2W_VELOCITY_TASK.actionSize}`,
);
if (typeof fixedInput === 'number' && fixedInput > 0 && fixedInput !== task.observationSize)
throw new Error(`策略观测维度不匹配:模型 ${fixedInput},任务 ${task.observationSize}`);
if (typeof fixedOutput === 'number' && fixedOutput > 0 && fixedOutput !== task.actionSize)
throw new Error(`策略动作维度不匹配:模型 ${fixedOutput},任务 ${task.actionSize}`);
if (deployment && (fixedInput !== task.observationSize || fixedOutput !== task.actionSize))
throw new Error('部署策略必须声明固定的观测/动作特征维度');
return new OnnxPolicyRuntime(
session,
bindings,
path,
session.inputNames[0],
session.outputNames[0],
task,
);
} catch (error) {
await session.release();
@@ -102,22 +111,32 @@ export class OnnxPolicyRuntime {
status(): RLPolicyStatus {
return {
taskId: GO2W_VELOCITY_TASK.id,
taskName: GO2W_VELOCITY_TASK.name,
taskId: this.task.id,
taskName: this.task.name,
path: this.path,
loaded: !this.disposed,
enabled: this.enabled,
controlHz: GO2W_VELOCITY_TASK.controlHz,
observationSize: GO2W_VELOCITY_TASK.observationSize,
actionSize: GO2W_VELOCITY_TASK.actionSize,
controlHz: this.task.controlHz,
observationSize: this.task.observationSize,
actionSize: this.task.actionSize,
inputName: this.inputName,
outputName: this.outputName,
command: { ...this.commandValue },
inferenceCount: this.inferenceCount,
lastInferenceMs: this.lastInferenceMs,
navigation: this.navigationStatus(),
error: this.error,
};
}
navigationStatus(): NavigationStatus | undefined {
return this.disposed ? undefined : this.bindings.navigationStatus?.();
}
setNavigationTarget(target: [number, number]): void {
if (!this.disposed) this.bindings.setNavigationTarget?.(target);
}
resetNavigationTarget(): void {
if (!this.disposed) this.bindings.resetNavigationTarget?.();
}
setCommand(command: RLCommand): void {
this.commandValue = clampGo2wCommand(command);
}
@@ -138,6 +157,7 @@ export class OnnxPolicyRuntime {
this.nextInferenceTime = time;
this.error = undefined;
this.bindings.clear();
this.bindings.reset?.(time);
}
step(time: number): void {
@@ -147,12 +167,14 @@ export class OnnxPolicyRuntime {
let observation: Float32Array;
try {
observation = this.bindings.observe(time, this.action, this.commandValue);
if (observation.length !== this.task.observationSize || !observation.every(Number.isFinite))
throw new Error('策略观测维度错误或包含非有限数');
} catch (error) {
this.fail(error);
return;
}
this.inFlight = true;
this.nextInferenceTime = time + 1 / GO2W_VELOCITY_TASK.controlHz;
this.nextInferenceTime = time + 1 / this.task.controlHz;
const started = performance.now(),
epoch = this.epoch;
const input = new ort.Tensor('float32', observation, [1, observation.length]);
@@ -163,9 +185,9 @@ export class OnnxPolicyRuntime {
const output = outputs[this.outputName];
if (!output || output.type !== 'float32')
throw new Error(`找不到 float32 输出:${this.outputName}`);
if (output.data.length !== GO2W_VELOCITY_TASK.actionSize)
if (output.data.length !== this.task.actionSize)
throw new Error(
`策略动作维度错误:期望 ${GO2W_VELOCITY_TASK.actionSize},实际 ${output.data.length}`,
`策略动作维度错误:期望 ${this.task.actionSize},实际 ${output.data.length}`,
);
const next = Float32Array.from(output.data as Float32Array, Number);
for (const value of next)
@@ -0,0 +1,114 @@
import { describe, it, expect } from 'vitest';
import deploymentJSON from '../fixtures/obstacleDeployment.json';
import golden from '../fixtures/obstacleRayGolden.json';
import { validatePolicyDeployment } from '../deployment';
import {
buildGo2ObstacleAvoidanceObservation,
obstacleNavigation,
rayBoxDistance,
sampleForwardRays,
} from './go2ObstacleAvoidance';
const deployment = validatePolicyDeployment(deploymentJSON);
const identity = [1, 0, 0, 0, 1, 0, 0, 0, 1];
describe('go2ObstacleAvoidance', () => {
it('严格47+32+2顺序及有限值、归一化范围校验', () => {
const base = Array.from({ length: 47 }, (_, i) => i),
depth = Array(32).fill(0.75);
const result = buildGo2ObstacleAvoidanceObservation(base, depth, [0.25, 0.5]);
expect(result).toBeInstanceOf(Float32Array);
expect(result).toHaveLength(81);
expect(Array.from(result.slice(0, 47))).toEqual(base);
expect(Array.from(result.slice(47, 79))).toEqual(depth);
expect(Array.from(result.slice(79))).toEqual([0.25, 0.5]);
expect(() => buildGo2ObstacleAvoidanceObservation([], depth, [0, 0])).toThrow(/维度/);
expect(() => buildGo2ObstacleAvoidanceObservation(base, depth, [NaN, 0])).toThrow(/非有限/);
expect(() =>
buildGo2ObstacleAvoidanceObservation(base, Array(32).fill(Infinity), [0, 0]),
).toThrow(/非有限/);
expect(() => buildGo2ObstacleAvoidanceObservation(base, Array(32).fill(1.1), [0, 0])).toThrow(
/0~1/,
);
});
it('目标前向、身后、yaw旋转、到达停止及离开恢复与后端公式一致', () => {
expect(obstacleNavigation([-5, 0, 0.32], [1, 0, 0, 0], [5, 0], 12)).toEqual({
command: { linearX: 0.6, linearY: 0, angularZ: 0 },
targetError: [0, 10 / 12],
});
expect(obstacleNavigation([5, 0, 0.32], [1, 0, 0, 0], [-5, 0], 12).command).toEqual({
linearX: 0,
linearY: 0,
angularZ: 1,
});
expect(obstacleNavigation([4.8, 0, 0.32], [0, 0, 0, 1], [5, 0], 12).command.linearX).toBe(0);
expect(obstacleNavigation([4.4, 0, 0.32], [1, 0, 0, 0], [5, 0], 12).command.linearX).toBe(0.6);
expect(
obstacleNavigation([0, 0, 0], [Math.SQRT1_2, 0, 0, Math.SQRT1_2], [0, 2], 12).targetError[0],
).toBeCloseTo(0);
});
it('动态换目标维持81维末尾归一化误差,到达后恢复转向', () => {
const position = [0, 0, 0.32],
quaternion = [1, 0, 0, 0];
expect(obstacleNavigation(position, quaternion, [0, 0], 12).command.angularZ).toBe(0);
for (const [target, heading] of [
[[0, 3], 0.5],
[[0, -3], -0.5],
] as const) {
const navigation = obstacleNavigation(position, quaternion, target, 12);
const observation = buildGo2ObstacleAvoidanceObservation(
Array(47).fill(0),
Array(32).fill(1),
navigation.targetError,
);
expect(observation).toHaveLength(81);
expect(Array.from(observation.slice(79))).toEqual([heading, 0.25]);
expect(navigation.command.angularZ).toBe(Math.sign(heading));
}
expect(obstacleNavigation(position, quaternion, [24, 0], 12).targetError[1]).toBe(1);
});
it('全部32射线与CPU mj_ray golden对照:完整yaw/pitch/roll且包含地面命中', () => {
for (const fixture of golden.cases) {
const result = sampleForwardRays(
fixture.position,
fixture.quaternion,
deployment.sensorCfg!,
(origin, direction) => {
const hits = deployment
.terrain!.boxes.map((box) =>
rayBoxDistance(origin, direction, box.pos, box.size, identity),
)
.filter((d) => d >= 0);
return hits.length ? Math.min(...hits) : -1;
},
);
result.depth.forEach((d, i) => expect(d).toBeCloseTo(fixture.depth[i], 6));
expect(result.rays).toHaveLength(32);
}
});
it('box内部/表面/平行/边缘/miss与CPU mj_ray一致', () => {
golden.boxEdges.forEach((item) =>
expect(
rayBoxDistance(item.origin, item.direction, [0, 0, 0], [1, 1, 1], identity),
).toBeCloseTo(item.distance, 8),
);
});
it('按右到左含端点采样,maxRange等号视为命中,越界/miss为1', () => {
let count = 0;
const sample = sampleForwardRays(
[0, 0, 0],
[1, 0, 0, 0],
deployment.sensorCfg!,
() => [4, 4.01, -1, 0][count++ % 4],
);
expect(sample.depth.slice(0, 4)).toEqual([1, 1, 1, 0]);
expect(sample.rays.slice(0, 4).map((r) => r.hit)).toEqual([true, false, false, true]);
expect(sample.rays[0].end[1]).toBeLessThan(0);
expect(sample.rays[30].end[1]).toBeGreaterThan(0);
});
});
it('使用部署中的target_velocity命令幅值,不改变81维目标误差或到达停机', () => {
const result = obstacleNavigation([0, 0, 0.32], [1, 0, 0, 0], [5, 0], 12, 0.9);
expect(result.command.linearX).toBe(0.9);
expect(result.targetError).toEqual([0, 5 / 12]);
expect(obstacleNavigation([4.8, 0, 0.32], [1, 0, 0, 0], [5, 0], 12, 1.2).command.linearX).toBe(0);
});
@@ -0,0 +1,160 @@
import { GO2W_VELOCITY_TASK } from './go2wVelocity';
import { OBSTACLE_TASK_ID, obstacleSensorPattern, type ObstacleSensorConfig } from '../deployment';
import type { RLCommand } from '../types';
export const GO2_OBSTACLE_AVOIDANCE_TASK = {
...GO2W_VELOCITY_TASK,
id: OBSTACLE_TASK_ID,
name: 'Go2 前视射线避障导航',
observationSize: 81,
};
export interface PerceptionRay {
origin: number[];
end: number[];
hit: boolean;
}
export function rotateVector(q: readonly number[], v: readonly number[]): number[] {
const [w, x, y, z] = q,
[vx, vy, vz] = v;
const tx = 2 * (y * vz - z * vy),
ty = 2 * (z * vx - x * vz),
tz = 2 * (x * vy - y * vx);
return [
vx + w * tx + y * tz - z * ty,
vy + w * ty + z * tx - x * tz,
vz + w * tz + x * ty - y * tx,
];
}
export function obstacleNavigation(
position: readonly number[],
quaternion: readonly number[],
target: readonly number[],
size: number,
speed = 0.6,
): { command: RLCommand; targetError: number[] } {
const [w, x, y, z] = quaternion;
const yaw = Math.atan2(2 * (w * z + x * y), 1 - 2 * (y * y + z * z));
const dx = target[0] - position[0],
dy = target[1] - position[1],
distance = Math.hypot(dx, dy);
const delta = Math.atan2(dy, dx) - yaw;
const heading = distance < 0.5 ? 0 : Math.atan2(Math.sin(delta), Math.cos(delta));
return {
command:
distance < 0.5
? { linearX: 0, linearY: 0, angularZ: 0 }
: {
linearX: speed * Math.max(0, Math.cos(heading)),
linearY: 0,
angularZ: Math.max(-1, Math.min(1, heading)),
},
targetError: [heading / Math.PI, Math.min(1, distance / size)],
};
}
/** Slab intersection with a physical box in its local frame. Inside hits the exit, as mj_ray does. */
export function rayBoxDistance(
origin: readonly number[],
direction: readonly number[],
center: ArrayLike<number>,
halfSize: ArrayLike<number>,
rotation: ArrayLike<number>,
axisAligned = false,
): number {
if (axisAligned) {
let near = -Infinity,
far = Infinity;
for (let i = 0; i < 3; i++) {
const o = origin[i] - center[i],
d = direction[i],
h = halfSize[i];
if (Math.abs(d) < 1e-12) {
if (Math.abs(o) > h) return -1;
continue;
}
let a = (-h - o) / d,
b = (h - o) / d;
if (a > b) {
const t = a;
a = b;
b = t;
}
if (a > near) near = a;
if (b < far) far = b;
if (near > far || far < 0) return -1;
}
return near >= 0 ? near : far;
}
const x = origin[0] - center[0],
y = origin[1] - center[1],
z = origin[2] - center[2];
let near = -Infinity,
far = Infinity;
for (let i = 0; i < 3; i++) {
const o = rotation[i] * x + rotation[3 + i] * y + rotation[6 + i] * z;
const d =
rotation[i] * direction[0] + rotation[3 + i] * direction[1] + rotation[6 + i] * direction[2];
if (Math.abs(d) < 1e-12) {
if (Math.abs(o) > halfSize[i]) return -1;
continue;
}
const a = (-halfSize[i] - o) / d,
b = (halfSize[i] - o) / d;
near = Math.max(near, Math.min(a, b));
far = Math.min(far, Math.max(a, b));
if (near > far || far < 0) return -1;
}
return near >= 0 ? near : far;
}
/** Per-binding directions are precomputed, avoiding trig and temporary arrays per ray/box. */
export function forwardRayDirections(sensor: ObstacleSensorConfig): number[][] {
const pattern = obstacleSensorPattern(sensor.sensorMode, sensor.fov);
return pattern.pitchAngles.flatMap((pitch) =>
pattern.yawAngles.map((yaw) => {
const p = (pitch * Math.PI) / 180,
y = (yaw * Math.PI) / 180;
return [Math.cos(p) * Math.cos(y), Math.cos(p) * Math.sin(y), Math.sin(p)];
}),
);
}
export function sampleForwardRays(
position: readonly number[],
quaternion: readonly number[],
sensor: ObstacleSensorConfig,
cast: (origin: number[], direction: number[]) => number,
localDirections = forwardRayDirections(sensor),
): { depth: number[]; rays: PerceptionRay[] } {
const offset = rotateVector(quaternion, sensor.offset),
origin = position.map((v, i) => v + offset[i]);
const rays: PerceptionRay[] = [],
depth: number[] = [];
for (const local of localDirections) {
const direction = rotateVector(quaternion, local);
const distance = cast(origin, direction);
if (!Number.isFinite(distance)) throw new Error('物理射线返回非有限距离');
const hit = distance >= 0 && distance <= sensor.maxDistance;
const length = hit ? distance : sensor.maxDistance;
depth.push(Math.max(0, Math.min(1, length / sensor.maxDistance)));
rays.push({ origin, end: direction.map((v, axis) => origin[axis] + v * length), hit });
}
return { depth, rays };
}
export function buildGo2ObstacleAvoidanceObservation(
base: ArrayLike<number>,
depth: ArrayLike<number>,
targetError: ArrayLike<number>,
): Float32Array {
if (
base.length !== 47 ||
(depth.length !== 32 && depth.length !== 48) ||
targetError.length !== 2
)
throw new Error('避障观测维度必须为47+(32或48)+2');
const result = Float32Array.from([
...Array.from(base),
...Array.from(depth),
...Array.from(targetError),
]);
if (!result.every(Number.isFinite)) throw new Error('避障观测包含非有限数');
if (Array.from(depth).some((v) => v < 0 || v > 1)) throw new Error('归一化深度必须在0~1之间');
return result;
}
@@ -0,0 +1,90 @@
import { expect, it } from 'vitest';
import legacy from '../fixtures/obstacleDeployment.json';
import fixture from '../fixtures/multiRingDeployment.json';
import golden from '../fixtures/multiRingGolden.json';
import { policyDeploymentsMatch, validatePolicyDeployment } from '../deployment';
import {
buildGo2ObstacleAvoidanceObservation,
forwardRayDirections,
rayBoxDistance,
sampleForwardRays,
} from './go2ObstacleAvoidance';
const deployment = validatePolicyDeployment(fixture);
const rotation = [1, 0, 0, 0, 1, 0, 0, 0, 1];
it('legacy缺字段与规范化metadata语义相等,multi拒绝矛盾及未知组合', () => {
expect(
policyDeploymentsMatch(
legacy as unknown as typeof deployment,
validatePolicyDeployment(legacy),
),
).toBe(true);
expect(deployment.observationSize).toBe(97);
for (const patch of [
{ rayCount: 32 },
{ sensorMode: 'unknown' },
{ pitchAngles: [0, -45, -20] },
{ yawCount: 32 },
{ yawAngles: Array(16).fill(0) },
{ angleUnit: 'rad' },
{ rayOrder: 'yaw-major' },
{ garbage: 1 },
])
expect(() =>
validatePolicyDeployment({ ...fixture, sensorCfg: { ...fixture.sensorCfg, ...patch } }),
).toThrow();
expect(() => validatePolicyDeployment({ ...fixture, observationSize: 81 })).toThrow(/观测维数/);
expect(() =>
validatePolicyDeployment({
...fixture,
sensorCfg: { ...fixture.sensorCfg, sensorMode: undefined },
}),
).toThrow();
});
it('所有48ray CPU mj_ray golden:完整姿态、5cm障碍、坑边/地板终点、层序', () => {
for (const c of golden.cases) {
const boxes = golden.layouts[c.layout as keyof typeof golden.layouts];
const sample = sampleForwardRays(c.position, c.quaternion, deployment.sensorCfg!, (o, d) => {
let nearest = -1;
for (const b of boxes) {
const t = rayBoxDistance(o, d, b.pos, b.size, rotation);
if (t >= 0 && (nearest < 0 || t < nearest)) nearest = t;
}
return nearest;
});
expect(sample.rays).toHaveLength(48);
c.depth.forEach((d, i) => expect(sample.depth[i]).toBeCloseTo(d, 6));
const obs = buildGo2ObstacleAvoidanceObservation(
Array(47).fill(0.2),
sample.depth,
[-0.5, 0.3],
);
expect(obs).toHaveLength(97);
expect(obs[95]).toBe(-0.5);
expect(obs[96]).toBeCloseTo(0.3);
}
const low = golden.cases.find((c) => c.layout === 'low')!;
expect(low.hitIds.slice(0, 16)).not.toContain(1);
expect(low.hitIds.slice(16)).toContain(1);
const rays = forwardRayDirections(deployment.sensorCfg!);
for (const [layer, pitch] of [0, -20, -45].entries()) {
expect(rays[layer * 16][2]).toBeCloseTo(Math.sin((pitch * Math.PI) / 180));
expect(rays[layer * 16][1]).toBeLessThan(0);
expect(rays[layer * 16 + 15][1]).toBeGreaterThan(0);
}
});
it('48ray完整range边界、miss、有限值,floor实际距离不隐藏', () => {
let i = 0;
const result = sampleForwardRays(
[0, 0, 0.32],
[1, 0, 0, 0],
deployment.sensorCfg!,
() => [-1, 0, 4, 4 + 1e-9][i++ % 4],
);
expect(result.depth).toEqual(Array.from({ length: 48 }, (_, j) => [1, 0, 1, 1][j % 4]));
expect(result.rays.map((r) => r.hit)).toEqual(
Array.from({ length: 48 }, (_, j) => [false, true, true, false][j % 4]),
);
expect(() =>
sampleForwardRays([0, 0, 0.32], [1, 0, 0, 0], deployment.sensorCfg!, () => NaN),
).toThrow();
});
@@ -0,0 +1,60 @@
/** Reproducible full-frame benchmark, not a single-ray microbenchmark. See TRAINING.md. */
import { forwardRayDirections, rayBoxDistance, sampleForwardRays } from './go2ObstacleAvoidance';
import { validatePolicyDeployment } from '../deployment';
import fixture from '../fixtures/multiRingDeployment.json';
export function benchmarkForwardRays(warmup = 3000, samples = 10000) {
const deployment = validatePolicyDeployment(fixture),
sensor = deployment.sensorCfg!;
const directions = forwardRayDirections(sensor);
const rotation = [1, 0, 0, 0, 1, 0, 0, 0, 1];
const defaultBoxes = deployment.terrain!.boxes;
// 16x16 grid + floor, same maximal box count as rough/wave deployment.
const worstBoxes = [
defaultBoxes[0],
...Array.from({ length: 256 }, (_, i) => ({
pos: [-3.75 + Math.floor(i / 16) * 0.5, -5.625 + (i % 16) * 0.75, 0.1],
size: [0.25, 0.375, 0.1],
yaw: 0,
})),
];
return [defaultBoxes, worstBoxes].map((boxes) => {
const cast = (o: number[], d: number[]) => {
let nearest = -1;
for (const box of boxes) {
const t = rayBoxDistance(o, d, box.pos, box.size, rotation, true);
if (t >= 0 && (nearest < 0 || t < nearest)) nearest = t;
}
return nearest;
};
let checksum = 0;
const frame = (i: number) => {
// Change pose each frame: no memoized results. Full attitude, origin, all intersections + output.
const yaw = (i % 100) * 0.004,
q = [
Math.cos(yaw / 2) * Math.cos(0.1),
Math.sin(0.1) * Math.cos(yaw / 2),
Math.sin(0.1) * Math.sin(yaw / 2),
Math.sin(yaw / 2) * Math.cos(0.1),
];
const r = sampleForwardRays([-5 + (i % 10) * 0.05, 0, 0.32], q, sensor, cast, directions);
checksum += r.depth[i % 48];
};
for (let i = 0; i < warmup; i++) frame(i);
const times = [];
for (let i = 0; i < samples; i++) {
const start = performance.now();
frame(i);
times.push(performance.now() - start);
}
times.sort((a, b) => a - b);
return {
boxes: boxes.length,
rays: 48,
warmup,
samples,
p50Ms: times[Math.floor(samples * 0.5)],
p95Ms: times[Math.floor(samples * 0.95)],
checksum,
};
});
}
+9 -1
View File
@@ -4,8 +4,15 @@ export interface RLCommand {
angularZ: number;
}
export interface NavigationStatus {
target: [number, number];
defaultTarget: [number, number];
distance: number;
targetHeight: number;
}
export interface RLPolicyStatus {
taskId: 'unitree-go2w-velocity';
taskId: string;
taskName: string;
path: string;
loaded: boolean;
@@ -18,6 +25,7 @@ export interface RLPolicyStatus {
command: RLCommand;
inferenceCount: number;
lastInferenceMs: number;
navigation?: NavigationStatus;
error?: string;
}
@@ -0,0 +1,259 @@
import { readFileSync } from 'node:fs';
import { beforeEach, describe, expect, it, vi } from 'vitest';
import type { ProjectManifest } from '../project/types';
import fixture from '../rl/fixtures/obstacleDeployment.json';
import { validatePolicyDeployment } from '../rl/deployment';
const ort = vi.hoisted(() => ({ create: vi.fn() }));
vi.mock('onnxruntime-web/wasm', () => ({
env: { wasm: {} },
InferenceSession: { create: ort.create },
}));
vi.mock('@mujoco/mujoco', async (original) => {
const actual = await original<typeof import('@mujoco/mujoco')>();
return {
...actual,
default: () =>
actual.default({ wasmBinary: readFileSync('node_modules/@mujoco/mujoco/mujoco.wasm') }),
};
});
import { MainThreadPhysicsAdapter } from './PhysicsAdapter';
const deployment = validatePolicyDeployment(fixture);
function graph(size = 47) {
return {
inputNames: ['obs'],
outputNames: ['actions'],
inputMetadata: [{ isTensor: true, type: 'float32', shape: [1, size] }],
outputMetadata: [{ isTensor: true, type: 'float32', shape: [1, 12] }],
release: vi.fn().mockResolvedValue(undefined),
};
}
function project(): ProjectManifest {
const doc = new DOMParser().parseFromString(
readFileSync('training_server/rl/src/assets/robots/unitree_go2/xmls/go2.xml', 'utf8'),
'application/xml',
);
doc.querySelectorAll('mesh, geom[mesh]').forEach((e) => e.remove());
const actuators = doc.createElement('actuator');
for (const name of deployment.jointNames) {
const motor = doc.createElement('motor');
motor.setAttribute('joint', name);
motor.setAttribute('name', `${name}_motor`);
actuators.append(motor);
}
doc.documentElement.append(actuators);
const data = new TextEncoder().encode(new XMLSerializer().serializeToString(doc));
return {
id: 'transaction',
name: 'go2',
files: [{ path: 'go2.xml', data, size: data.length, source: 'file', mimeType: 'text/xml' }],
entries: [{ path: 'go2.xml', format: 'mjcf', label: 'go2' }],
maps: [],
selectedEntry: 'go2.xml',
totalBytes: data.length,
};
}
async function existing() {
const adapter = new MainThreadPhysicsAdapter(),
source = project(),
oldGraph = graph();
await adapter.load(source, 'go2.xml');
ort.create.mockResolvedValueOnce(oldGraph);
await adapter.loadRLPolicy(new Uint8Array([1]), 'old-flat.onnx');
adapter.setRLPolicyEnabled(true);
const old = adapter.session!;
old.data.time = 7;
old.data.qvel[0] = 0.123;
old.data.qpos[0] = 1.25;
old.data.ctrl[0] = 0.3;
adapter.setPaused(false);
return { adapter, source, old, oldGraph, before: adapter.snapshot()! };
}
beforeEach(() => {
vi.clearAllMocks();
ort.create.mockReset();
});
describe('PhysicsAdapter training transaction', () => {
it.each(['wrong graph', 'ORT initialization', 'binding'] as const)(
'%s失败保留原session/workspace/策略及物理状态',
async (failure) => {
const { adapter, source, old, oldGraph, before } = await existing();
const workspace = adapter.workspace,
dispose = vi.spyOn(old, 'dispose');
const candidate = graph(47);
if (failure === 'ORT initialization')
ort.create.mockRejectedValueOnce(new Error('ORT init failed'));
else ort.create.mockResolvedValueOnce(candidate);
if (failure === 'binding') {
const text = new TextDecoder()
.decode(source.files[0].data)
.replace('name="FL_hip_joint_motor"', 'name="incompatible_name"');
source.files[0].data = new TextEncoder().encode(text);
}
try {
await expect(
adapter.load(source, 'go2.xml', {
trainingDeployment: deployment,
trainingPolicy: { data: new Uint8Array([2]), path: 'candidate.onnx' },
}),
).rejects.toThrow(
failure === 'wrong graph' ? /维度/ : failure === 'binding' ? /驱动器/ : /ORT init/,
);
expect(adapter.session).toBe(old);
expect(adapter.workspace).toBe(workspace);
expect(adapter.snapshot()).toEqual(before);
expect(dispose).not.toHaveBeenCalled();
expect(oldGraph.release).not.toHaveBeenCalled();
if (failure === 'wrong graph') expect(candidate.release).toHaveBeenCalledOnce();
} finally {
adapter.dispose();
}
},
);
it('完成ORT初始化前不发布候选,成功后仍保留旧资源以供viewer失败回滚', async () => {
const { adapter, source, old, oldGraph, before } = await existing();
let finish!: (value: ReturnType<typeof graph>) => void;
let started!: () => void;
const initialized = new Promise<void>((resolve) => {
started = resolve;
});
ort.create.mockImplementationOnce(() => {
started();
return new Promise((resolve) => {
finish = resolve;
});
});
const candidate = graph(81);
try {
const loading = adapter.load(source, 'go2.xml', {
trainingDeployment: deployment,
trainingPolicy: { data: new Uint8Array([2]), path: 'new.onnx' },
});
await initialized;
expect(adapter.session).toBe(old);
expect(oldGraph.release).not.toHaveBeenCalled();
finish(candidate);
await loading;
expect(adapter.session).not.toBe(old);
expect(adapter.snapshot()?.rlPolicy).toMatchObject({
path: 'new.onnx',
observationSize: 81,
enabled: true,
});
expect(oldGraph.release).not.toHaveBeenCalled();
adapter.rollbackRetired();
expect(adapter.session).toBe(old);
expect(adapter.snapshot()).toEqual(before);
for (let i = 0; i < 10; i++) await Promise.resolve();
expect(candidate.release).toHaveBeenCalledOnce();
expect(oldGraph.release).not.toHaveBeenCalled();
} finally {
adapter.dispose();
}
});
});
it('候选ORT等待期间另一场景完成加载,迟到候选不得覆盖最新场景', async () => {
const { adapter, source } = await existing();
let finish!: (value: ReturnType<typeof graph>) => void;
let started!: () => void;
const initialized = new Promise<void>((resolve) => {
started = resolve;
});
ort.create.mockImplementationOnce(() => {
started();
return new Promise((resolve) => {
finish = resolve;
});
});
const candidate = graph(81);
try {
const stale = adapter.load(source, 'go2.xml', {
trainingDeployment: deployment,
trainingPolicy: { data: new Uint8Array([2]), path: 'stale.onnx' },
});
await initialized;
await adapter.load(source, 'go2.xml');
const latest = adapter.session,
before = adapter.snapshot();
const rejected = expect(stale).rejects.toThrow(/取消/);
finish(candidate);
await rejected;
expect(adapter.session).toBe(latest);
expect(adapter.snapshot()).toEqual(before);
for (let i = 0; i < 10; i++) await Promise.resolve();
expect(candidate.release).toHaveBeenCalledOnce();
} finally {
adapter.dispose();
}
});
describe('PhysicsAdapter custom_boxes applied scene compiler', () => {
it('真实WASM读取工程路径/嵌套body变换/旋转多实例;与共享CPU布局halfsize和原点一致', async () => {
const { default: shared } =
await import('../../../training_server/tests/fixtures/custom-boxes.json');
const source = project();
const descriptor = {
schemaVersion: 1,
id: 'warehouse',
name: '仓库',
coordinateSystem: { units: 'm', up: 'Z', forward: '+X' },
physics: { source: 'physics/world.xml' },
spawnPoints: [],
};
for (const [path, text] of [
['maps/warehouse/map.json', JSON.stringify(descriptor)],
[
'maps/warehouse/physics/world.xml',
'<mujoco><worldbody><body name="nested" pos="0 1 .5"><geom name="wall" type="box" size=".4 .3 .5" friction=".8 .005 .0001"/><geom name="decoration" type="sphere" size="10" contype="0" conaffinity="0"/></body></worldbody></mujoco>',
],
]) {
const data = new TextEncoder().encode(text);
source.files.push({ path, data, size: data.length, source: 'file', mimeType: 'text/xml' });
}
const assets = [
{
id: 'one',
name: 'one',
selection: {
kind: 'project' as const,
descriptorPath: 'maps/warehouse/map.json',
positionX: 1,
positionY: 1,
},
},
];
const adapter = new MainThreadPhysicsAdapter();
try {
await adapter.load(source, 'go2.xml', { mapAssets: assets });
const initial = adapter.exportTrainingTerrain(assets, { spawn: [-2, -1], target: [2, -1] });
// Geometry from actual WASM matches the JSON consumed by CPU TerrainGenerator tests.
expect(initial.boxes[1]).toEqual(shared.boxes[1]);
expect(initial.actualObstacleCount).toBe(1);
adapter.session!.data.qpos[0] = 8;
expect(adapter.exportTrainingTerrain(assets).spawn[0]).not.toBe(8);
expect(() => adapter.exportTrainingTerrain([])).toThrow(/过时/);
const multi = [
...assets,
{
id: 'two',
name: 'two',
selection: { ...assets[0].selection, positionX: -2, positionY: 1, yawDeg: 45 },
},
];
await adapter.load(source, 'go2.xml', { mapAssets: multi });
const layout = adapter.exportTrainingTerrain(multi, { spawn: [-2, -1], target: [2, -1] });
expect(layout.boxes).toHaveLength(3);
expect(layout.boxes[2].pos[0]).toBeCloseTo(-2 - Math.SQRT1_2, 12);
expect(layout.boxes[2].pos[1]).toBeCloseTo(1 + Math.SQRT1_2, 12);
expect(layout.boxes[2].size[0]).toBeCloseTo(0.7 * Math.SQRT1_2, 12);
expect(layout.boxes[2].size[1]).toBeCloseTo(0.7 * Math.SQRT1_2, 12);
adapter.rollbackRetired();
expect(adapter.exportTrainingTerrain(assets).actualObstacleCount).toBe(1);
// A deployment replacement cannot pretend to be the original applied editor map.
await adapter.load(source, 'go2.xml', { mapAssets: assets, trainingDeployment: deployment });
expect(() => adapter.exportTrainingTerrain(assets)).toThrow(/已应用/);
} finally {
adapter.dispose();
}
});
});
+146 -4
View File
@@ -1,3 +1,10 @@
import {
composeTrainingMap,
trainingTerrainFromCompiledScene,
type CompiledMapGeometry,
type TrainingSceneCoordinates,
} from '../map/trainingMap';
import type { PolicyDeployment, TrainingTerrain } from '../rl/deployment';
import type { MainModule } from '@mujoco/mujoco';
import type { ProjectFile, ProjectManifest } from '../project/types';
import { prepareProjectForMujoco } from '../project/importer';
@@ -36,6 +43,9 @@ export interface PhysicsLoadProgress {
}
export interface PhysicsLoadOptions {
trainingDeployment?: PolicyDeployment;
/** Candidate policy is initialized/validated before replacing the active session. */
trainingPolicy?: { data: Uint8Array; path: string };
urdfMode?: UrdfLoadMode;
baseMode?: UrdfBaseMode;
enhancements?: UrdfEnhancementOptions;
@@ -68,9 +78,15 @@ export interface PhysicsAdapter {
setControllerEnabled(enabled: boolean): void;
sendControllerCommand(command: ControllerCommand): void;
removeController(): void;
loadRLPolicy(model: Uint8Array, path: string): Promise<RLPolicyStatus>;
loadRLPolicy(
model: Uint8Array,
path: string,
deployment?: PolicyDeployment,
): Promise<RLPolicyStatus>;
setRLPolicyEnabled(enabled: boolean): void;
setRLCommand(command: RLCommand): void;
setNavigationTarget(target: [number, number]): void;
resetNavigationTarget(): void;
removeRLPolicy(): void;
configureDataRecorder(config: Partial<DataRecorderConfig>): DataRecorderStatus | undefined;
startDataRecording(): DataRecorderStatus | undefined;
@@ -82,6 +98,10 @@ export interface PhysicsAdapter {
releaseRetired(): void;
rollbackRetired(): void;
exportMjcf(): Uint8Array;
exportTrainingTerrain(
assets: readonly PlacedMapAsset[],
coordinates?: TrainingSceneCoordinates,
): TrainingTerrain;
dispose(): void;
}
@@ -108,6 +128,15 @@ export class MainThreadPhysicsAdapter implements PhysicsAdapter {
session: SimulationSession | null = null;
workspace: MemfsWorkspace | null = null;
private supportFiles: ProjectFile[] = [];
private trainingScenes = new WeakMap<
SimulationSession,
{
assets: string;
geometries: CompiledMapGeometry[];
initialPose: number[];
extent: number;
}
>();
private loadGeneration = 0;
private disposed = false;
private retiredSession: SimulationSession | null = null;
@@ -196,7 +225,7 @@ export class MainThreadPhysicsAdapter implements PhysicsAdapter {
intermediate.dispose();
}
}
if (hasMaps) {
if (hasMaps && !options.trainingDeployment?.terrain) {
report(0.72, '组合机器人与物理地图');
if (entry?.format === 'urdf' && urdfMode === 'native')
throw new Error('原生 URDF 模式暂不支持地图,请切换为转换模式');
@@ -232,13 +261,102 @@ export class MainThreadPhysicsAdapter implements PhysicsAdapter {
modelRelativePath = mapPath;
modelPath = workspace.path(modelRelativePath);
}
if (options.trainingDeployment?.terrain) {
if (entry?.format === 'urdf' && urdfMode === 'native')
throw new Error('训练地图需要URDF转换模式');
const intermediate = new SimulationSession(module, modelPath);
try {
const flatPath = `${modelRelativePath}.training.xml`;
if (!module.mj_saveLastXML(workspace.path(flatPath), intermediate.model))
throw new Error('无法展开训练场景');
workspace.writeGenerated(
flatPath,
composeTrainingMap(
new TextEncoder().encode(workspace.readText(flatPath)),
options.trainingDeployment,
),
);
modelPath = workspace.path(flatPath);
} finally {
intermediate.dispose();
}
warnings.push('已用策略配套训练布局替换场景地形;编辑器地图未更改。Go2-W动力学不等同Go2。');
}
report(0.84, '编译模型与物理数据');
console.info('[MuJoCo] 编译模型', modelPath);
nextSession = new SimulationSession(module, modelPath, warnings);
if (options.trainingDeployment) nextSession.configureDeployment(options.trainingDeployment);
if (entry?.format === 'urdf' && urdfMode === 'native') {
const offset = nextSession.alignLowestPointToGround();
warnings.push(`原生 URDF 已整体平移 ${offset.toFixed(4)} m,使最低点位于 z=0`);
}
if (placedMapAssets?.length && !options.trainingDeployment) {
const sanitize = (id: string) => id.replace(/[^a-zA-Z0-9_-]/g, '_');
const prefixes = placedMapAssets.map((asset) =>
asset.selection.kind === 'builtin'
? `__platform_map_${sanitize(asset.id)}__`
: `__platform_map_${sanitize(`${resolveProjectMap(prepared.manifest, asset.selection.descriptorPath).definition.id}_${asset.id}`)}_`,
);
const model = nextSession.model,
data = nextSession.data;
const geometries: CompiledMapGeometry[] = [];
for (let id = 0; id < model.ngeom; id++) {
const geom = model.geom(id);
try {
const name = geom.name;
geometries.push({
name,
type: Number(model.geom_type[id]),
position: Array.from(data.geom_xpos.slice(id * 3, id * 3 + 3)),
rotation: Array.from(data.geom_xmat.slice(id * 9, id * 9 + 9)),
size: Array.from(model.geom_size.slice(id * 3, id * 3 + 3)),
friction: Array.from(model.geom_friction.slice(id * 3, id * 3 + 3)),
collision: Boolean(model.geom_contype[id] || model.geom_conaffinity[id]),
static: Number(model.body_weldid[Number(model.geom_bodyid[id])]) === 0,
map:
prefixes.some((prefix) => name.startsWith(prefix)) ||
name === '__platform_map_ground__',
});
} finally {
geom.delete();
}
}
const roots = Array.from({ length: model.njnt }, (_, id) => id).filter(
(id) => Number(model.jnt_type[id]) === 0,
);
if (roots.length === 1) {
const address = Number(model.jnt_qposadr[roots[0]]);
this.trainingScenes.set(nextSession, {
assets: JSON.stringify(placedMapAssets),
geometries,
initialPose: Array.from(data.qpos.slice(address, address + 7)),
extent: Math.max(
4,
...placedMapAssets.map((asset) =>
asset.selection.kind === 'builtin'
? Math.max(
Math.abs(asset.selection.config.positionX),
Math.abs(asset.selection.config.positionY),
) +
asset.selection.config.size / 2
: 0,
),
),
});
}
}
if (options.trainingPolicy) {
if (!options.trainingDeployment?.terrain) throw new Error('事务策略加载需要配套训练地图');
report(0.9, '校验候选场景的 ONNX 策略');
await nextSession.loadRLPolicy(
options.trainingPolicy.data,
options.trainingPolicy.path,
options.trainingDeployment,
);
// Configure/reset already supplied the correct spawn. Enable while still paused;
// the candidate cannot advance until the viewer transaction commits.
nextSession.setRLPolicyEnabled(true);
}
if (warnings.length) console.info('[MuJoCo] 兼容与地图处理', warnings);
console.info('[MuJoCo] 模型编译完成');
report(0.94, '生成初始仿真状态');
@@ -263,6 +381,20 @@ export class MainThreadPhysicsAdapter implements PhysicsAdapter {
);
}
}
exportTrainingTerrain(
assets: readonly PlacedMapAsset[],
coordinates?: TrainingSceneCoordinates,
): TrainingTerrain {
const scene = this.session && this.trainingScenes.get(this.session);
if (!scene || scene.assets !== JSON.stringify(assets))
throw new Error('没有匹配的已应用碰撞场景/唯一浮动机器人,或场景已过时;请先应用地图');
return trainingTerrainFromCompiledScene(
scene.geometries,
scene.initialPose,
scene.extent,
coordinates,
);
}
advance(now: number): FrameResult {
return this.session?.advance(now) ?? { steps: 0, stepMs: 0, overBudget: false };
}
@@ -315,13 +447,23 @@ export class MainThreadPhysicsAdapter implements PhysicsAdapter {
removeController(): void {
this.session?.removeController();
}
async loadRLPolicy(model: Uint8Array, path: string): Promise<RLPolicyStatus> {
async loadRLPolicy(
model: Uint8Array,
path: string,
deployment?: PolicyDeployment,
): Promise<RLPolicyStatus> {
if (!this.session) throw new Error('请先加载模型');
return this.session.loadRLPolicy(model, path);
return this.session.loadRLPolicy(model, path, deployment);
}
setRLPolicyEnabled(enabled: boolean): void {
this.session?.setRLPolicyEnabled(enabled);
}
setNavigationTarget(target: [number, number]): void {
this.session?.setNavigationTarget(target);
}
resetNavigationTarget(): void {
this.session?.resetNavigationTarget();
}
setRLCommand(command: RLCommand): void {
this.session?.setRLCommand(command);
}
@@ -0,0 +1,71 @@
import { describe, it, expect, vi } from 'vitest';
const pending = vi.hoisted(() => ({
resolvers: [] as (() => void)[],
runtimes: [] as { dispose: ReturnType<typeof vi.fn>; status: ReturnType<typeof vi.fn> }[],
}));
vi.mock('../rl/runtime/Go2wPolicyBindings', () => ({
Go2wPolicyBindings: class {
constructor(
_model: unknown,
_data: unknown,
private write: (id: number, value: number) => void,
) {}
clear() {
this.write(0, 0);
}
},
}));
vi.mock('../rl/runtime/OnnxPolicyRuntime', () => ({
OnnxPolicyRuntime: {
load: (_model: unknown, _path: string, bindings: { clear(): void }) =>
new Promise((resolve) => {
const runtime = {
dispose: vi.fn(() => bindings.clear()),
status: vi.fn(() => ({ loaded: true })),
};
pending.runtimes.push(runtime);
pending.resolvers.push(() => resolve(runtime));
}),
},
}));
import { SimulationSession } from './SimulationSession';
async function flush() {
for (let i = 0; i < 20; i++) await Promise.resolve();
}
function session(): SimulationSession {
return Object.assign(Object.create(SimulationSession.prototype) as object, {
model: {},
data: { ctrl: new Float64Array(1) },
rlPolicyLoadGeneration: 0,
disposed: false,
setActuator: vi.fn(),
}) as unknown as SimulationSession;
}
describe('SimulationSession policy loading race', () => {
it('模型释放后迟到ORT session必须release,不能写已释放MjData', async () => {
pending.resolvers.length = pending.runtimes.length = 0;
const s = session();
const load = s.loadRLPolicy(new Uint8Array(), 'late.onnx');
await flush();
Object.assign(s, { disposed: true });
pending.resolvers[0]();
await expect(load).rejects.toThrow(/取消/);
expect(pending.runtimes[0].dispose).toHaveBeenCalledOnce();
expect(s.setActuator).not.toHaveBeenCalled();
});
it('同一模型两次加载只接受最新一份,旧clear不能清新策略动作', async () => {
pending.resolvers.length = pending.runtimes.length = 0;
const s = session();
const first = s.loadRLPolicy(new Uint8Array(), 'first.onnx');
await flush();
const second = s.loadRLPolicy(new Uint8Array(), 'second.onnx');
await flush();
pending.resolvers[1]();
await expect(second).resolves.toMatchObject({ loaded: true });
pending.resolvers[0]();
await expect(first).rejects.toThrow(/取消/);
expect(pending.runtimes[0].dispose).toHaveBeenCalledOnce();
expect(pending.runtimes[1].dispose).not.toHaveBeenCalled();
expect(s.setActuator).not.toHaveBeenCalled();
});
});
@@ -1,3 +1,10 @@
import { Go2ObstacleAvoidanceBindings } from '../rl/runtime/Go2ObstacleAvoidanceBindings';
import {
OBSTACLE_TASK_ID,
policyDeploymentsMatch,
validatePolicyDeployment,
type PolicyDeployment,
} from '../rl/deployment';
import type { MainModule, MjData, MjModel, MjvPerturb, MjvScene } from '@mujoco/mujoco';
import { meshIdFromSceneDataId } from './geometry';
import { PythonControllerRuntime } from '../controller/PythonControllerRuntime';
@@ -106,6 +113,9 @@ export class SimulationSession {
private pythonController?: PythonControllerRuntime;
private controllerLoadGeneration = 0;
private rlPolicy?: OnnxPolicyRuntime;
private obstacleBindings?: Go2ObstacleAvoidanceBindings;
private deployment?: PolicyDeployment;
private deploymentInitialQpos?: Float64Array;
private rlPolicyLoadGeneration = 0;
private dataRecorder!: DataRecorder;
@@ -189,6 +199,7 @@ export class SimulationSession {
reset(): void {
this.setPaused(true);
this.module.mj_resetData(this.model, this.data);
if (this.deploymentInitialQpos) this.data.qpos.set(this.deploymentInitialQpos);
this.module.mj_forward(this.model, this.data);
this.clearExternalForce();
this.data.ctrl.fill(0);
@@ -253,20 +264,132 @@ export class SimulationSession {
if (!enabled) this.data.ctrl.fill(0);
}
async loadRLPolicy(model: Uint8Array, path: string): Promise<RLPolicyStatus> {
navigationStatus() {
return this.rlPolicy?.navigationStatus();
}
setNavigationTarget(target: [number, number]): void {
this.rlPolicy?.setNavigationTarget(target);
}
resetNavigationTarget(): void {
this.rlPolicy?.resetNavigationTarget();
}
perceptionRays() {
return this.obstacleBindings?.rays ?? [];
}
configureDeployment(input: PolicyDeployment): void {
const d = validatePolicyDeployment(input);
if (!d.terrain) return;
const seen = new Set<number>();
for (let id = 0; id < this.model.ngeom; id++) {
const geom = this.model.geom(id);
try {
if (Number(this.model.geom_bodyid[id]) !== 0) {
if (geom.name.startsWith('__training_terrain_'))
throw new Error('机器人几何不能使用训练地图保留名称');
continue;
}
const match = /^__training_terrain_(\d+)$/.exec(geom.name);
const index = match ? Number(match[1]) : -1;
const box = d.terrain.boxes[index];
if (
!box ||
seen.has(index) ||
Number(this.model.geom_type[id]) !== 6 ||
Number(this.model.geom_group[id]) !== 2
)
throw new Error('配套地图只支持声明的静态box,不允许额外plane或其他形状');
seen.add(index);
for (let axis = 0; axis < 3; axis++) {
if (
Math.abs(Number(this.data.geom_xpos[id * 3 + axis]) - box.pos[axis]) > 1e-6 ||
Math.abs(Number(this.model.geom_size[id * 3 + axis]) - box.size[axis]) > 1e-6
)
throw new Error('编译后的训练地图坐标/半尺寸不匹配');
}
} finally {
geom.delete();
}
}
if (seen.size !== d.terrain.boxes.length) throw new Error('配套训练地图缺少box');
// Construction validates robot topology before changing the initial pose.
new Go2wPolicyBindings(this.model, this.data, (id, value) => this.setActuator(id, value));
const initialQpos = Float64Array.from(this.model.qpos0);
let freeCount = 0;
for (let id = 0; id < this.model.njnt; id++) {
const joint = this.model.jnt(id);
try {
const address = Number(joint.qposadr);
if (Number(this.model.jnt_type[id]) === 0) {
freeCount++;
initialQpos.set([...d.terrain.spawn, ...d.terrain.spawnQuaternion], address);
} else {
const index = d.jointNames.indexOf(joint.name);
if (index >= 0) initialQpos[address] = d.defaultJointPosition[index];
}
} finally {
joint.delete();
}
}
if (freeCount !== 1) throw new Error('训练评测需要唯一浮动基座');
for (let id = 0; id < this.model.nactuator; id++) {
const actuator = this.model.actuator(id);
try {
const jointId = Number(actuator.trnid[0]);
if (jointId < 0 || jointId >= this.model.njnt) continue;
const joint = this.model.jnt(jointId);
try {
const index = d.jointNames.indexOf(joint.name);
if (index >= 0) {
actuator.forcelimited = 1;
actuator.forcerange[0] = -d.effortLimits[index];
actuator.forcerange[1] = d.effortLimits[index];
}
} finally {
joint.delete();
}
} finally {
actuator.delete();
}
}
this.deployment = d;
// qpos0 contains hinge reference angles, not the desired joint pose. Do not mutate it.
this.deploymentInitialQpos = initialQpos;
this.reset();
}
async loadRLPolicy(
model: Uint8Array,
path: string,
deployment?: PolicyDeployment,
): Promise<RLPolicyStatus> {
const generation = ++this.rlPolicyLoadGeneration;
const bindings = new Go2wPolicyBindings(this.model, this.data, (id, value) =>
this.setActuator(id, value),
);
if (deployment) deployment = validatePolicyDeployment(deployment);
if (
deployment?.terrain &&
(!this.deployment || !policyDeploymentsMatch(deployment, this.deployment))
)
throw new Error('策略配套地图尚未加载或配置不匹配');
let cancelled = false;
const writeActuator = (id: number, value: number) => {
if (!this.disposed && !cancelled) this.setActuator(id, value);
};
const bindings =
deployment?.taskId === OBSTACLE_TASK_ID
? new Go2ObstacleAvoidanceBindings(this.model, this.data, writeActuator, deployment)
: new Go2wPolicyBindings(this.model, this.data, writeActuator, deployment?.effortLimits);
const { OnnxPolicyRuntime: Runtime } = await import('../rl/runtime/OnnxPolicyRuntime');
const runtime = await Runtime.load(model, path, bindings);
const runtime = await Runtime.load(model, path, bindings, deployment);
if (this.disposed || generation !== this.rlPolicyLoadGeneration) {
cancelled = true;
runtime.dispose();
throw new Error('模型已切换,ONNX 策略加载已取消');
}
this.data.ctrl.fill(0);
this.rlPolicy?.dispose();
this.rlPolicy = runtime;
this.obstacleBindings = bindings instanceof Go2ObstacleAvoidanceBindings ? bindings : undefined;
return runtime.status();
}
@@ -306,6 +429,7 @@ export class SimulationSession {
this.rlPolicyLoadGeneration += 1;
this.rlPolicy?.dispose();
this.rlPolicy = undefined;
this.obstacleBindings = undefined;
this.data.ctrl.fill(0);
}
@@ -148,3 +148,66 @@ describe('LocalTrainingClient', () => {
expect(() => new LocalTrainingClient('http://localhost:8765', '')).toThrow('访问令牌');
});
});
it('LocalTrainingClient完整透传自定义任务与嵌套地形/传感器参数', async () => {
const fetchMock = vi.fn().mockResolvedValue(Response.json({ id: 'job' }));
vi.stubGlobal('fetch', fetchMock);
const request = {
taskId: 'Unitree-Go2-ObstacleAvoidance',
numEnvs: 16,
maxIterations: 1,
seed: 42,
runName: 'avoid',
device: 'gpu' as const,
gpuIds: [0],
wandbMode: 'disabled' as const,
terrainPreset: 'discrete_obstacles',
terrainParams: { friction: 0.9, obstacle_count: 12 },
sensorType: 'raycast' as const,
sensorCfg: { fov: 60, maxDistance: 3 },
};
await new LocalTrainingClient('http://127.0.0.1:8765', 'test').start(request);
expect(JSON.parse(String(fetchMock.mock.calls[0][1].body))).toEqual(request);
});
it('上传使用File原始body、显式模板、Bearer及AbortSignal,不发送JSON/base64/服务器路径', async () => {
const fetchMock = vi
.fn()
.mockResolvedValue(new Response(JSON.stringify({ id: 'source' }), { status: 201 }));
vi.stubGlobal('fetch', fetchMock);
const file = new File(['weights'], '单个 actor.ONNX');
const controller = new AbortController();
await new LocalTrainingClient('http://localhost:8765', 'secret-token').uploadPretrained(
file,
'go2-legacy47-v1',
controller.signal,
);
const [url, init] = fetchMock.mock.calls[0] as [string, RequestInit];
const query = new URL(url).searchParams;
expect(query.get('format')).toBe('onnx');
expect(query.get('template')).toBe('go2-legacy47-v1');
expect(query.get('name')).toBe(file.name);
expect(url).not.toContain('secret-token');
expect(init.body).toBe(file);
expect(init.signal).toBe(controller.signal);
expect(new Headers(init.headers).get('Content-Type')).toBe('application/octet-stream');
expect(new Headers(init.headers).get('Authorization')).toBe('Bearer secret-token');
expect(new Headers(init.headers).has('Content-Length')).toBe(false); // Browser owns this forbidden header.
});
it('上传前拒绝ZIP、空文件和超过各格式上限,仍由服务验证实际模型', () => {
const fetchMock = vi.fn();
vi.stubGlobal('fetch', fetchMock);
const client = new LocalTrainingClient('http://localhost:8765', 'token');
for (const [name, size] of [
['model.zip', 1],
['model.pt', 0],
['model.pt', 256 * 1024 ** 2 + 1],
['model.onnx', 64 * 1024 ** 2 + 1],
] as const) {
const file = new File(['x'], name);
Object.defineProperty(file, 'size', { value: size });
expect(() => client.uploadPretrained(file, 'go2-legacy47-v1')).toThrow();
}
expect(fetchMock).not.toHaveBeenCalled();
});
@@ -1,5 +1,6 @@
import type {
ParameterConstraint,
PretrainedSource,
RewardPreset,
TuningCapability,
TuningCreateRequest,
@@ -57,6 +58,26 @@ export class LocalTrainingClient {
health(): Promise<TrainingServerInfo> {
return this.json('/api/training/health');
}
uploadPretrained(
file: File,
template: 'go2-legacy47-v1',
signal?: AbortSignal,
): Promise<PretrainedSource> {
const format = file.name.split('.').pop()?.toLowerCase();
if (format !== 'pt' && format !== 'onnx')
throw new Error('请选择单个.pt或.onnx文件,不支持ZIP/目录');
const limit = (format === 'pt' ? 256 : 64) * 1024 ** 2;
if (!file.size || file.size > limit)
throw new Error(`文件不能为空,${format}上限为${limit / 1024 ** 2}MiB`);
if (template !== 'go2-legacy47-v1') throw new Error('请先确认Go2 legacy47模板');
const query = new URLSearchParams({ format, template, name: file.name });
return this.json(`/api/training/pretrained-sources/upload?${query}`, {
method: 'POST',
headers: { 'Content-Type': 'application/octet-stream' },
body: file,
signal,
});
}
start(request: TrainingRequest): Promise<TrainingJob> {
return this.json('/api/training/jobs', {
method: 'POST',
@@ -1,3 +1,8 @@
import customLayout from '../../../training_server/tests/fixtures/custom-boxes.json';
import { validateCustomTerrain, validatePolicyDeployment } from '../rl/deployment';
import fixture from '../rl/fixtures/obstacleDeployment.json';
import { DEFAULT_PHYSICAL_MAP_CONFIG } from '../map/types';
import { TRAINING_JOB_KEY } from './storage';
import { fireEvent, render, screen, waitFor } from '@testing-library/react';
import { beforeEach, describe, expect, it, vi } from 'vitest';
import { LocalTrainingPanel } from './LocalTrainingPanel';
@@ -161,3 +166,308 @@ describe('LocalTrainingPanel', () => {
);
});
});
const customTasks = ['Unitree-Go2-Flat', 'Unitree-Go2-ObstacleAvoidance', 'Unitree-Go2-Rough'];
const customMetadata = customTasks.map((id) => ({
id,
name: id,
browserCompatible: id !== 'Unitree-Go2-Rough',
terrainPresets: [
'plane',
'discrete_obstacles',
'rough',
'pyramid_stairs',
'wave',
'custom_boxes',
],
terrainParameters: {
size: { min: 8, max: 24, default: 12 },
friction: { min: 0.2, max: 2, default: 0.8 },
obstacle_count: { min: 1, max: 100, default: 24, integer: true },
},
sensorTypes: id.includes('Obstacle') ? ['raycast'] : [],
sensorModes: ['single_ring_raycast', 'multi_ring_raycast'],
sensorParameters: id.includes('Obstacle')
? { fov: { min: 30, max: 120, default: 90 }, maxDistance: { min: 1, max: 5, default: 4 } }
: {},
mapSyncScope: '单块预设参数',
}));
function customServer(job?: unknown) {
const fetchMock = vi.fn(async (url: string, init?: RequestInit) => {
if (url.endsWith('/health'))
return Response.json({
ready: true,
tasks: customTasks,
taskMetadata: customMetadata,
trainerRoot: '/local',
});
if (url.includes('/presets')) return Response.json({ presets: [] });
if (url.endsWith('/policy.onnx')) return new Response(new Uint8Array([8, 9]));
if (init?.method === 'POST' || job)
return Response.json(
job ?? {
id: 'c'.repeat(32),
taskId: 'Unitree-Go2-ObstacleAvoidance',
state: 'queued',
progress: 0,
iteration: 0,
maxIterations: 2,
message: '',
logs: [],
artifactReady: false,
},
);
return Response.json({ error: 'not found' }, { status: 404 });
});
vi.stubGlobal('fetch', fetchMock);
return fetchMock;
}
async function connectCustom() {
fireEvent.change(screen.getByLabelText('训练服务访问令牌'), { target: { value: 'token' } });
fireEvent.click(screen.getByRole('button', { name: '连接' }));
await screen.findByText('/local');
}
describe('LocalTrainingPanel 自定义任务', () => {
it('动态任务/地图/传感器配置传入payload,切换任务清掉sensor与reward状态', async () => {
const fetchMock = customServer();
render(<LocalTrainingPanel onPolicyReady={vi.fn()} />);
await connectCustom();
fireEvent.change(screen.getByLabelText('训练任务'), {
target: { value: 'Unitree-Go2-ObstacleAvoidance' },
});
expect(screen.getByLabelText('训练地形')).toHaveValue('discrete_obstacles');
expect(screen.getByLabelText('奖励配置 preset')).toBeDisabled();
fireEvent.change(screen.getByLabelText('感知角 FOV'), { target: { value: '60' } });
fireEvent.change(screen.getByLabelText('训练地形'), { target: { value: 'rough' } });
expect(screen.getByText(/训练专用 box/)).toBeInTheDocument();
fireEvent.change(screen.getByLabelText('训练任务'), { target: { value: 'Unitree-Go2-Flat' } });
expect(screen.queryByLabelText('感知角 FOV')).not.toBeInTheDocument();
expect(screen.getByLabelText('训练地形')).toHaveValue('');
fireEvent.change(screen.getByLabelText('训练任务'), {
target: { value: 'Unitree-Go2-ObstacleAvoidance' },
});
expect(screen.getByLabelText('感知角 FOV')).toHaveValue(90);
fireEvent.change(screen.getByLabelText('感知角 FOV'), { target: { value: '60' } });
fireEvent.click(screen.getByRole('button', { name: '发起本地训练' }));
await waitFor(() =>
expect(fetchMock.mock.calls.some((call) => call[1]?.method === 'POST')).toBe(true),
);
const payload = JSON.parse(
String(fetchMock.mock.calls.find((call) => call[1]?.method === 'POST')?.[1]?.body),
);
expect(payload).toMatchObject({
taskId: 'Unitree-Go2-ObstacleAvoidance',
terrainPreset: 'discrete_obstacles',
sensorType: 'raycast',
sensorCfg: { fov: 60 },
});
expect(payload.rewardPresetId).toBeUndefined();
});
it('无地图/多实例明确报错而不提交;超限参数阻止请求', async () => {
const fetchMock = customServer();
render(<LocalTrainingPanel onPolicyReady={vi.fn()} />);
await connectCustom();
fireEvent.click(screen.getByRole('button', { name: '同步当前场景地图' }));
expect(screen.getByRole('alert')).toHaveTextContent('没有已应用');
fireEvent.change(screen.getByLabelText('训练地形'), { target: { value: 'plane' } });
fireEvent.change(screen.getByLabelText('地图尺寸 m'), { target: { value: '100' } });
fireEvent.click(screen.getByRole('button', { name: '发起本地训练' }));
await waitFor(() => expect(screen.getByRole('alert')).toHaveTextContent('超出允许范围'));
expect(fetchMock.mock.calls.some((call) => call[1]?.method === 'POST')).toBe(false);
});
it('旧Rough作业禁用导入,不误当Flat', async () => {
customServer({
id: 'r'.repeat(32),
taskId: 'Unitree-Go2-Rough',
state: 'succeeded',
artifactReady: true,
progress: 1,
iteration: 1,
maxIterations: 1,
logs: [],
message: '',
});
localStorage.setItem(TRAINING_JOB_KEY, 'r'.repeat(32));
render(<LocalTrainingPanel onPolicyReady={vi.fn()} />);
await connectCustom();
expect(await screen.findByRole('button', { name: '导入策略' })).toBeDisabled();
expect(screen.getByText(/234维 Rough/)).toBeInTheDocument();
});
});
describe('LocalTrainingPanel 配套策略交接', () => {
it('同步权威布局并上传完整boxes,不再重跑预设种子', async () => {
const fetchMock = customServer();
render(
<LocalTrainingPanel
compileScene={() => validateCustomTerrain(customLayout)}
onPolicyReady={vi.fn()}
sceneMaps={[
{
id: 'one',
name: 'one',
selection: {
kind: 'builtin',
config: {
...DEFAULT_PHYSICAL_MAP_CONFIG,
preset: 'discrete_obstacles',
size: 16,
friction: 1.2,
seed: 123,
obstacleCount: 17,
},
},
},
]}
/>,
);
await connectCustom();
fireEvent.click(screen.getByRole('button', { name: '同步当前场景地图' }));
expect(screen.getByLabelText('训练地形')).toHaveValue('custom_boxes');
expect(screen.getByText('已将视口中 1 个自定义障碍物编译为训练地图布局')).toBeInTheDocument();
expect(screen.getByLabelText('出生 X')).toHaveValue(-2);
expect(screen.getByLabelText('目标 Y')).toHaveValue(-1);
expect(screen.getByText(/旋转障碍会膨胀/)).toBeInTheDocument();
fireEvent.click(screen.getByRole('button', { name: '发起本地训练' }));
await waitFor(() =>
expect(fetchMock.mock.calls.some((call) => call[1]?.method === 'POST')).toBe(true),
);
const request = JSON.parse(
String(fetchMock.mock.calls.find((call) => call[1]?.method === 'POST')?.[1]?.body),
);
expect(request.customTerrainBoxes).toEqual(customLayout);
expect(request.terrainPreset).toBe('custom_boxes');
expect(request.terrainParams).toEqual({ size: 12, friction: 0.8 });
});
it('等待地图/策略异步回调完成,失败显示错误,不提前解除busy', async () => {
localStorage.setItem(TRAINING_JOB_KEY, 'd'.repeat(32));
customServer({
id: 'd'.repeat(32),
taskId: fixture.taskId,
state: 'succeeded',
artifactReady: true,
progress: 1,
iteration: 1,
maxIterations: 1,
logs: ['Mean value loss: 0.0141'],
message: '',
deployment: fixture,
});
let reject!: (error: Error) => void;
const ready = vi.fn(
() =>
new Promise<void>((_resolve, failure) => {
reject = failure;
}),
);
render(<LocalTrainingPanel onPolicyReady={ready} />);
await connectCustom();
const button = await screen.findByRole('button', { name: '导入策略' });
fireEvent.click(button);
await waitFor(() => expect(ready).toHaveBeenCalledOnce());
expect(ready.mock.calls[0]).toEqual([expect.any(File), validatePolicyDeployment(fixture)]);
expect(button).toBeDisabled();
expect(screen.getByText('价值损失')).toBeInTheDocument();
reject(new Error('配套地图失败'));
await waitFor(() => expect(screen.getByRole('alert')).toHaveTextContent('配套地图失败'));
expect(button).toBeEnabled();
});
});
it('custom_boxes拒绝草稿、过时编译场景和未重新同步的坐标;安全区报错不删障碍', async () => {
const fetchMock = customServer();
const assets = [
{
id: 'one',
name: 'one',
selection: {
kind: 'builtin' as const,
config: { ...DEFAULT_PHYSICAL_MAP_CONFIG, preset: 'flat' as const },
},
},
];
let current = validateCustomTerrain(customLayout);
const compileScene = vi.fn(
(coordinates?: import('../map/trainingMap').TrainingSceneCoordinates) => ({
...structuredClone(current),
...(coordinates?.spawn ? { spawn: [...coordinates.spawn, 0.32] } : {}),
...(coordinates?.target ? { target: [...coordinates.target] } : {}),
}),
);
const { rerender } = render(
<LocalTrainingPanel
onPolicyReady={vi.fn()}
sceneMaps={assets}
sceneDirty
compileScene={compileScene}
/>,
);
await connectCustom();
fireEvent.click(screen.getByRole('button', { name: '同步当前场景地图' }));
expect(screen.getByRole('alert')).toHaveTextContent('草稿');
expect(compileScene).not.toHaveBeenCalled();
rerender(
<LocalTrainingPanel onPolicyReady={vi.fn()} sceneMaps={assets} compileScene={compileScene} />,
);
fireEvent.click(screen.getByRole('button', { name: '同步当前场景地图' }));
fireEvent.change(screen.getByLabelText('目标 X'), { target: { value: 1 } });
fireEvent.change(screen.getByLabelText('目标 Y'), { target: { value: 2 } });
fireEvent.click(screen.getByRole('button', { name: '发起本地训练' }));
await waitFor(() => expect(screen.getByRole('alert')).toHaveTextContent('重新同步'));
fireEvent.click(screen.getByRole('button', { name: '同步当前场景地图' }));
expect(screen.getByRole('alert')).toHaveTextContent('安全区');
expect(current.boxes).toHaveLength(2);
fireEvent.change(screen.getByLabelText('目标 X'), { target: { value: 2 } });
fireEvent.change(screen.getByLabelText('目标 Y'), { target: { value: -1 } });
fireEvent.click(screen.getByRole('button', { name: '同步当前场景地图' }));
current = { ...current, boxes: [current.boxes[0], { ...current.boxes[1], pos: [1, 3, 0.5] }] };
fireEvent.click(screen.getByRole('button', { name: '发起本地训练' }));
await waitFor(() => expect(screen.getByRole('alert')).toHaveTextContent('过时'));
expect(fetchMock.mock.calls.some((call) => call[1]?.method === 'POST')).toBe(false);
});
it('显式选择multi48传入训练请求,默认与切换任务仍single32', async () => {
const fetchMock = customServer();
render(<LocalTrainingPanel onPolicyReady={vi.fn()} />);
await connectCustom();
fireEvent.change(screen.getByLabelText('训练任务'), {
target: { value: 'Unitree-Go2-ObstacleAvoidance' },
});
expect(screen.getByLabelText('传感器模式')).toHaveValue('single_ring_raycast');
fireEvent.change(screen.getByLabelText('传感器模式'), {
target: { value: 'multi_ring_raycast' },
});
fireEvent.click(screen.getByRole('button', { name: '发起本地训练' }));
await waitFor(() => expect(fetchMock.mock.calls.some((c) => c[1]?.method === 'POST')).toBe(true));
const payload = JSON.parse(
String(fetchMock.mock.calls.find((c) => c[1]?.method === 'POST')?.[1]?.body),
);
expect(payload.sensorCfg.sensorMode).toBe('multi_ring_raycast');
});
it('Flat奖励菜单只展示服务确认归属Flat的preset,Obstacle和无身份项不展示', async () => {
const entries = [
{ id: 'a'.repeat(32), name: '合法历史Flat', taskId: 'Unitree-Go2-Flat' },
{ id: 'b'.repeat(32), name: 'Obstacle不能混入', taskId: 'Unitree-Go2-ObstacleAvoidance' },
{ id: 'c'.repeat(32), name: '缺少权威身份' },
];
vi.stubGlobal(
'fetch',
vi.fn(async (url: string) => {
if (url.endsWith('/health'))
return Response.json({
ready: true,
tasks: customTasks,
taskMetadata: customMetadata,
trainerRoot: '/local',
});
if (url.endsWith('/presets')) return Response.json({ presets: entries });
return Response.json({ error: 'not found' }, { status: 404 });
}),
);
render(<LocalTrainingPanel onPolicyReady={vi.fn()} />);
await connectCustom();
expect(await screen.findByRole('option', { name: '合法历史Flat' })).toBeInTheDocument();
expect(screen.queryByRole('option', { name: 'Obstacle不能混入' })).not.toBeInTheDocument();
expect(screen.queryByRole('option', { name: '缺少权威身份' })).not.toBeInTheDocument();
});
+451 -23
View File
@@ -1,4 +1,17 @@
import { useEffect, useState, type ReactNode } from 'react';
import { PretrainedIdentity, PretrainedSourceSelect } from './PretrainedSourceSelect';
import { pretrainedSelectionError } from './pretrainedSelection';
import { TrainingMetricsPanel } from './TrainingMetricsPanel';
import { trainingLosses } from './trainingLosses';
import type { PolicyDeployment, TrainingTerrain } from '../rl/deployment';
import {
OBSTACLE_TASK_ID,
validatePolicyDeployment,
validateCustomTerrain,
readPolicyDeployment,
} from '../rl/deployment';
import type { TrainingSceneCompiler } from '../map/trainingMap';
import type { PlacedMapAsset } from '../map/types';
import { useEffect, useRef, useState, type ReactNode } from 'react';
import { Download, ExternalLink, Link, Play, Server, Square } from 'lucide-react';
import { Badge, Button, ProgressBar, PropertyRow, Select } from '../components/ui';
import { LocalTrainingClient } from './LocalTrainingClient';
@@ -33,7 +46,17 @@ function stateLabel(state: TrainingJob['state']): string {
}[state];
}
export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File): void }) {
export function LocalTrainingPanel({
onPolicyReady,
compileScene,
sceneMaps = [],
sceneDirty = false,
}: {
onPolicyReady(file: File, deployment?: PolicyDeployment): void | Promise<void>;
compileScene?: TrainingSceneCompiler;
sceneMaps?: readonly PlacedMapAsset[];
sceneDirty?: boolean;
}) {
const [endpoint, setEndpoint] = useState(() =>
localStored(TRAINING_ENDPOINT_KEY, DEFAULT_TRAINING_ENDPOINT),
);
@@ -42,6 +65,9 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
const [job, setJob] = useState<TrainingJob>();
const [presets, setPresets] = useState<RewardPreset[]>([]);
const [rewardPresetId, setRewardPresetId] = useState('');
const [pretrainedSourceId, setPretrainedSourceId] = useState('');
const [uploading, setUploading] = useState(false);
const [connectionRevision, setConnectionRevision] = useState(0);
const [busy, setBusy] = useState(false),
[error, setError] = useState<string>();
const [taskId, setTaskId] = useState('Unitree-Go2-Flat'),
@@ -53,15 +79,85 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
[gpuIds, setGpuIds] = useState('0'),
[wandbMode, setWandbMode] = useState<WandbMode>('offline');
const [terrainPreset, setTerrainPreset] = useState('');
const [customTerrainBoxes, setCustomTerrainBoxes] = useState<TrainingTerrain>();
const [syncedScene, setSyncedScene] = useState<string>();
const [terrainParams, setTerrainParams] = useState<Record<string, number>>({});
const [sensorMode, setSensorMode] = useState<'single_ring_raycast' | 'multi_ring_raycast'>(
'single_ring_raycast',
);
const [sensorCfg, setSensorCfg] = useState<Record<string, number>>({});
const metadata = server?.taskMetadata?.find((item) => item.id === taskId);
const sourceSelectionError = pretrainedSelectionError(
server?.pretrainedSources,
taskId,
pretrainedSourceId,
);
const selectTask = (id: string) => {
setTaskId(id);
setCustomTerrainBoxes(undefined);
setSyncedScene(undefined);
setRewardPresetId('');
setTerrainParams({});
setSensorCfg({});
setSensorMode('single_ring_raycast');
setTerrainPreset(id === OBSTACLE_TASK_ID ? 'discrete_obstacles' : '');
};
const syncMap = () => {
try {
if (sceneDirty) throw new Error('请先应用地图草稿,再同步训练地图');
if (!compileScene || !sceneMaps.length) throw new Error('没有已应用的碰撞地图');
if (!metadata?.terrainPresets.includes('custom_boxes'))
throw new Error('当前服务/任务不支持custom_boxes,请升级训练服务');
setSyncedScene(undefined);
const layout = compileScene(
customTerrainBoxes
? {
spawn: [customTerrainBoxes.spawn[0], customTerrainBoxes.spawn[1]],
target: [customTerrainBoxes.target[0], customTerrainBoxes.target[1]],
}
: undefined,
);
setCustomTerrainBoxes(layout);
setTerrainPreset('custom_boxes');
setTerrainParams({ size: layout.size, friction: layout.friction });
validateCustomTerrain(layout);
setSyncedScene(JSON.stringify(sceneMaps));
setError(undefined);
} catch (value) {
setError(errorText(value));
}
};
const connectionEpoch = useRef(0);
const connected = () => {
connectionEpoch.current += 1;
setConnectionRevision(connectionEpoch.current);
setServer(undefined);
setJob(undefined);
setPresets([]);
setRewardPresetId('');
setError(undefined);
};
useEffect(
() => () => {
connectionEpoch.current += 1;
},
[],
);
const connect = async () => {
setBusy(true);
setError(undefined);
const epoch = ++connectionEpoch.current;
setConnectionRevision(epoch);
try {
const client = new LocalTrainingClient(endpoint, token),
info = await client.health();
if (epoch !== connectionEpoch.current) return;
setServer(info);
try {
setPresets(await client.presets());
const nextPresets = await client.presets();
if (epoch !== connectionEpoch.current) return;
setPresets(nextPresets);
} catch {
setPresets([]);
}
@@ -70,11 +166,12 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
} catch {
/* 当前会话仍可连接 */
}
if (info.tasks.length && !info.tasks.includes(taskId)) setTaskId(info.tasks[0]);
if (info.tasks.length && !info.tasks.includes(taskId)) selectTask(info.tasks[0]);
const remembered = info.activeJobId ?? localStored(TRAINING_JOB_KEY);
if (remembered) {
try {
const recovered = await client.job(remembered);
if (epoch !== connectionEpoch.current) return;
setJob(recovered);
try {
localStorage.setItem(TRAINING_JOB_KEY, recovered.id);
@@ -129,10 +226,57 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
typeof event.data !== 'object'
)
return;
const data = event.data as { type?: string; sessionId?: string; policy?: unknown };
const data = event.data as {
type?: string;
sessionId?: string;
policy?: unknown;
taskId?: string;
};
const source = event.source as Window;
if (data.type === 'mujoco-tuning-ready') {
source.postMessage({ type: 'mujoco-tuning-credentials', endpoint, token }, event.origin);
try {
if (taskId === OBSTACLE_TASK_ID && terrainPreset === 'custom_boxes') {
if (
sceneDirty ||
syncedScene !== JSON.stringify(sceneMaps) ||
!compileScene ||
!customTerrainBoxes
)
throw new Error('自定义地图已过时,请重新同步后打开调参');
const current = compileScene({
spawn: [customTerrainBoxes.spawn[0], customTerrainBoxes.spawn[1]],
target: [customTerrainBoxes.target[0], customTerrainBoxes.target[1]],
});
if (JSON.stringify(current) !== JSON.stringify(customTerrainBoxes))
throw new Error('碰撞场景已过时,请重新同步');
}
source.postMessage(
{
type: 'mujoco-tuning-credentials',
endpoint,
token,
trainingContext: {
taskId: taskId === OBSTACLE_TASK_ID ? taskId : 'Unitree-Go2-Flat',
seed,
...(pretrainedSourceId ? { pretrainedSourceId } : {}),
...(taskId === OBSTACLE_TASK_ID
? {
taskConfig: {
terrainPreset,
terrainParams,
sensorType: 'raycast',
sensorCfg: { ...sensorCfg, sensorMode },
...(terrainPreset === 'custom_boxes' ? { customTerrainBoxes } : {}),
},
}
: {}),
},
},
event.origin,
);
} catch (value) {
setError(errorText(value));
}
}
if (data.type === 'mujoco-tuning-import-policy' && data.sessionId) {
const reply = (ok: boolean, message?: string) => {
@@ -159,7 +303,12 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
if (!(policy instanceof File) || !/\.onnx$/i.test(policy.name))
throw new Error('调参工作台返回的 ONNX 策略无效');
if (policy.size > 64 * 1024 * 1024) throw new Error('ONNX 策略不能超过 64 MiB');
onPolicyReady(policy);
if (data.taskId === OBSTACLE_TASK_ID) {
const deployment = readPolicyDeployment(new Uint8Array(await policy.arrayBuffer()));
if (deployment?.taskId !== OBSTACLE_TASK_ID)
throw new Error('避障最佳策略缺少匹配部署契约');
await onPolicyReady(policy, deployment);
} else await onPolicyReady(policy);
reply(true);
} catch (value) {
const message = errorText(value);
@@ -171,7 +320,23 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
};
window.addEventListener('message', receive);
return () => window.removeEventListener('message', receive);
}, [endpoint, onPolicyReady, token]);
}, [
endpoint,
pretrainedSourceId,
onPolicyReady,
token,
taskId,
seed,
terrainPreset,
terrainParams,
sensorCfg,
sensorMode,
customTerrainBoxes,
syncedScene,
sceneMaps,
sceneDirty,
compileScene,
]);
const openTuningDashboard = () => {
rememberTrainingConnection(endpoint, token);
@@ -179,9 +344,20 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
};
const start = async () => {
if (uploading) return;
if (sourceSelectionError) {
setError(sourceSelectionError);
return;
}
setBusy(true);
setError(undefined);
try {
if (
taskId === OBSTACLE_TASK_ID &&
sensorMode === 'multi_ring_raycast' &&
!metadata?.sensorModes?.includes(sensorMode)
)
throw new Error('训练服务不支持multi_ring_raycast,请升级服务');
const ids =
device === 'gpu'
? gpuIds
@@ -191,6 +367,43 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
: [];
if (ids.some((id) => !Number.isInteger(id) || id < 0))
throw new Error('GPU 编号必须是非负整数');
for (const [values, schema] of [
[terrainParams, metadata?.terrainParameters],
[sensorCfg, metadata?.sensorParameters],
] as const) {
for (const [key, value] of Object.entries(values)) {
const bounds = schema?.[key];
if (
!bounds ||
!Number.isFinite(value) ||
value < bounds.min ||
value > bounds.max ||
(bounds.integer && !Number.isInteger(value))
)
throw new Error(`参数 ${key} 超出允许范围`);
}
}
if ((terrainParams.obstacle_height_min ?? 0.2) > (terrainParams.obstacle_height_max ?? 0.6))
throw new Error('障碍物最小高度不能超过最大高度');
if ((sensorCfg.safetyDistance ?? 0.5) >= (sensorCfg.maxDistance ?? 4))
throw new Error('安全距离必须小于探测距离');
if (terrainPreset === 'custom_boxes') {
if (
sceneDirty ||
!syncedScene ||
syncedScene !== JSON.stringify(sceneMaps) ||
!compileScene ||
!customTerrainBoxes
)
throw new Error('自定义地图未同步或场景/坐标已更改,请重新同步');
validateCustomTerrain(customTerrainBoxes);
const current = compileScene({
spawn: [customTerrainBoxes.spawn[0], customTerrainBoxes.spawn[1]],
target: [customTerrainBoxes.target[0], customTerrainBoxes.target[1]],
});
if (JSON.stringify(current) !== JSON.stringify(customTerrainBoxes))
throw new Error('已编译碰撞场景已过时,请重新同步');
}
const next = await new LocalTrainingClient(endpoint, token).start({
taskId,
numEnvs,
@@ -200,7 +413,13 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
device,
gpuIds: ids,
wandbMode,
rewardPresetId: rewardPresetId || undefined,
rewardPresetId: taskId === 'Unitree-Go2-Flat' ? rewardPresetId || undefined : undefined,
...(pretrainedSourceId ? { pretrainedSourceId } : {}),
...(terrainPreset ? { terrainPreset, terrainParams } : {}),
...(terrainPreset === 'custom_boxes' ? { customTerrainBoxes } : {}),
...(taskId === OBSTACLE_TASK_ID
? { sensorType: 'raycast' as const, sensorCfg: { ...sensorCfg, sensorMode } }
: {}),
});
setJob(next);
try {
@@ -231,7 +450,11 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
setBusy(true);
setError(undefined);
try {
onPolicyReady(await new LocalTrainingClient(endpoint, token).downloadPolicy(job.id));
if (job.taskId !== 'Unitree-Go2-Flat' && !job.deployment)
throw new Error('该任务缺少浏览器部署契约');
const deployment = job.deployment ? validatePolicyDeployment(job.deployment) : undefined;
const file = await new LocalTrainingClient(endpoint, token).downloadPolicy(job.id);
await onPolicyReady(file, deployment);
} catch (value) {
setError(errorText(value));
} finally {
@@ -249,7 +472,12 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
aria-label="本地训练服务地址"
className="field h-7 min-w-0 flex-1 px-2 text-xs text-text-primary"
value={endpoint}
onChange={(event) => setEndpoint(event.target.value)}
disabled={busy}
onChange={(event) => {
if (busy) return;
connected();
setEndpoint(event.target.value);
}}
/>
<Button
icon={<Link className="h-3.5 w-3.5" />}
@@ -268,7 +496,12 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
autoComplete="off"
className="field h-7 w-full px-2 text-xs text-text-primary"
value={token}
onChange={(event) => setToken(event.target.value)}
disabled={busy}
onChange={(event) => {
if (busy) return;
connected();
setToken(event.target.value);
}}
/>
</label>
<div className="mt-2 flex items-center justify-between rounded-md border border-border bg-surface px-2 py-1.5 text-[10px] text-text-tertiary">
@@ -288,21 +521,166 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
打开自调参 Agent 工作台
</Button>
{server?.ready && !job && (
<div className="mt-3 space-y-2">
<fieldset disabled={busy} className="mt-3 space-y-2">
<Field label="训练任务">
<Select
aria-label="训练任务"
className="w-full"
value={taskId}
onChange={(event) => setTaskId(event.target.value)}
onChange={(event) => selectTask(event.target.value)}
>
{server.tasks.map((task) => (
<option key={task} value={task}>
{task}
{server.taskMetadata?.find((item) => item.id === task)?.name ?? task}
</option>
))}
</Select>
</Field>
<PretrainedSourceSelect
sources={server.pretrainedSources}
taskId={taskId}
value={pretrainedSourceId}
onChange={setPretrainedSourceId}
disabled={busy || uploading}
upload={{
endpoint,
token,
revision: connectionRevision,
enabled: Boolean(server.pretrainedUpload?.enabled),
onBusyChange: setUploading,
onUploaded: (source) => {
setServer(
(current) =>
current && {
...current,
pretrainedSources: [
...(current.pretrainedSources ?? []).filter((s) => s.id !== source.id),
source,
],
},
);
setPretrainedSourceId(source.id);
},
}}
/>
{metadata && (
<>
<Field label="训练地形">
<Select
aria-label="训练地形"
value={terrainPreset}
onChange={(e) => {
setTerrainPreset(e.target.value);
setCustomTerrainBoxes(undefined);
setSyncedScene(undefined);
setTerrainParams({});
}}
>
{taskId !== OBSTACLE_TASK_ID && <option value="">原任务默认地形</option>}
{metadata.terrainPresets.map((preset) => (
<option key={preset} value={preset}>
{TERRAIN_LABELS[preset] ?? preset}
</option>
))}
</Select>
</Field>
<Button disabled={busy} onClick={syncMap}>
同步当前场景地图
</Button>
<p className="text-[10px] text-text-tertiary">
从全部已应用实例的实际碰撞几何编译世界AABB;旋转障碍会膨胀,底板标准化为z=[-0.2,0],出生高度标准化为0.32m。仅保证训练与浏览器使用相同boxes,不等于原OBB。mesh/hfield、地下结构、混合摩擦明确拒绝。
</p>
{terrainPreset === 'custom_boxes' && customTerrainBoxes && (
<>
<div className="grid grid-cols-2 gap-2">
{(['spawn', 'target'] as const).flatMap((key) =>
[0, 1].map((i) => (
<NumberField
key={`${key}${i}`}
label={`${key === 'spawn' ? '出生' : '目标'} ${i === 0 ? 'X' : 'Y'}`}
value={customTerrainBoxes[key][i]}
min={-12}
max={12}
step={0.1}
onChange={(value) => {
setSyncedScene(undefined);
setCustomTerrainBoxes(
(old) =>
old && {
...old,
[key]: old[key].map((v, j) => (i === j ? value : v)),
},
);
}}
/>
)),
)}
</div>
<p>
此处起终点仅用于固定评估和部署初始演示,训练会在同一连通自由区域内逐episode随机采样。参考点须保留0.5m圆形安全区;修改后请重新同步。
</p>
{syncedScene === JSON.stringify(sceneMaps) && !sceneDirty && (
<p role="status">
已将视口中 {customTerrainBoxes.actualObstacleCount}{' '}
个自定义障碍物编译为训练地图布局
</p>
)}
</>
)}
{terrainPreset && terrainPreset !== 'custom_boxes' && (
<div className="grid grid-cols-2 gap-2">
{Object.entries(metadata.terrainParameters).map(([key, bounds]) => (
<NumberField
key={key}
label={PARAMETER_LABELS[key] ?? key}
value={terrainParams[key] ?? bounds.default}
min={bounds.min}
max={bounds.max}
step={bounds.integer ? 1 : 0.01}
onChange={(value) => setTerrainParams((old) => ({ ...old, [key]: value }))}
/>
))}
</div>
)}
{['rough', 'wave', 'pyramid_stairs'].includes(terrainPreset) && (
<p>训练专用 box 离散近似布局,不等于编辑器高度场。</p>
)}
{taskId === OBSTACLE_TASK_ID && (
<details open>
<summary>避障传感器高级设置</summary>
<Field label="传感器模式">
<Select
aria-label="传感器模式"
value={sensorMode}
onChange={(event) => setSensorMode(event.target.value as typeof sensorMode)}
>
<option value="single_ring_raycast">水平32射线 / 81维(默认)</option>
<option
value="multi_ring_raycast"
disabled={!metadata?.sensorModes?.includes('multi_ring_raycast')}
>
三层48射线 / 97维(非高程图)
</option>
</Select>
</Field>
{Object.entries(metadata.sensorParameters).map(([key, bounds]) => (
<NumberField
key={key}
label={PARAMETER_LABELS[key] ?? key}
value={sensorCfg[key] ?? bounds.default}
min={bounds.min}
max={bounds.max}
step={0.01}
onChange={(value) => setSensorCfg((old) => ({ ...old, [key]: value }))}
/>
))}
</details>
)}
{!metadata.browserCompatible && (
<p>此任务可训练,但浏览器不支持其观测契约,不能一键部署。</p>
)}
</>
)}
<div className="grid grid-cols-2 gap-2">
<NumberField
label="并行环境"
@@ -358,17 +736,20 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
</div>
<Field label="奖励配置">
<Select
disabled={taskId !== 'Unitree-Go2-Flat'}
aria-label="奖励配置 preset"
className="w-full"
value={rewardPresetId}
onChange={(event) => setRewardPresetId(event.target.value)}
>
<option value="">仓库默认奖励</option>
{presets.map((preset) => (
<option key={preset.id} value={preset.id}>
{preset.name}
</option>
))}
{presets
.filter((preset) => preset.taskId === 'Unitree-Go2-Flat')
.map((preset) => (
<option key={preset.id} value={preset.id}>
{preset.name}
</option>
))}
</Select>
</Field>
<Field label="实验记录">
@@ -387,7 +768,7 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
variant="primary"
className="w-full"
icon={<Play className="h-3.5 w-3.5" />}
disabled={busy}
disabled={busy || uploading || Boolean(sourceSelectionError)}
onClick={() => void start()}
>
发起本地训练
@@ -396,7 +777,7 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
训练使用本地 mjlab
任务资产,不会把浏览器中的模型上传到网络。服务一次只运行一个训练任务。
</p>
</div>
</fieldset>
)}
{job && (
<div className="mt-3 rounded-lg border border-border bg-surface p-2.5">
@@ -416,11 +797,26 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
{stateLabel(job.state)}
</Badge>
</div>
<PretrainedIdentity source={job.pretrained} />
{job.taskId === 'Unitree-Go2-Rough' && (
<p>234维 Rough 策略仅支持后端评测,浏览器不可加载。</p>
)}
{job.deployment?.terrain && (
<p className="text-[10px] text-text-tertiary">
导入将替换当前物理地图并启动配套策略;
{job.deployment.terrain.approximation ? '训练专用近似布局' : '配套碰撞布局'}
。请先保存场景。
</p>
)}
<ProgressBar value={job.progress} label="训练进度" />
<div className="mt-2">
<PropertyRow label="迭代" value={`${job.iteration} / ${job.maxIterations}`} />
<PropertyRow label="状态" value={job.message} />
{trainingLosses(job.logs).map(({ label, value }) => (
<PropertyRow key={label} label={label} value={value} />
))}
</div>
<TrainingMetricsPanel key={job.id} jobId={job.id} logs={job.logs} />
{job.logs.length > 0 && (
<details className="mt-2">
<summary className="cursor-pointer text-[10px] text-text-secondary">最近日志</summary>
@@ -443,7 +839,12 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
) : (
<>
<Button
disabled={busy || !job.artifactReady}
disabled={
busy ||
!job.artifactReady ||
job.taskId === 'Unitree-Go2-Rough' ||
(job.deployment && !job.deployment.browserCompatible)
}
icon={<Download className="h-3.5 w-3.5" />}
onClick={() => void importResult()}
>
@@ -492,11 +893,13 @@ function NumberField({
min,
max,
onChange,
step = 1,
}: {
label: string;
value: number;
min: number;
max: number;
step?: number;
onChange(value: number): void;
}) {
return (
@@ -504,6 +907,7 @@ function NumberField({
<input
aria-label={label}
type="number"
step={step}
className="field h-7 w-full px-2 text-xs text-text-primary"
value={value}
min={min}
@@ -513,3 +917,27 @@ function NumberField({
</Field>
);
}
const TERRAIN_LABELS: Record<string, string> = {
custom_boxes: '自定义场景碰撞布局(AABB近似)',
plane: '平地',
discrete_obstacles: '离散障碍物',
rough: '崎岖地面',
pyramid_stairs: '金字塔台阶',
wave: '波浪地形',
};
const PARAMETER_LABELS: Record<string, string> = {
size: '地图尺寸 m',
obstacle_count: '障碍物数量',
obstacle_height_min: '最小障碍高度 m',
obstacle_height_max: '最大障碍高度 m',
spacing: '障碍物间距 m',
friction: '地面摩擦',
roughness: '崎岖高度 m',
step_height: '台阶高度 m',
wave_amplitude: '波浪幅度 m',
fov: '感知角 FOV',
maxDistance: '探测距离 m',
safetyDistance: '安全距离 m',
avoidanceWeight: '避障权重',
};
@@ -0,0 +1,358 @@
import { fireEvent, render, screen, waitFor } from '@testing-library/react';
import { beforeEach, expect, it, vi } from 'vitest';
import { LocalTrainingPanel } from './LocalTrainingPanel';
import { TuningApp } from '../tuning/TuningApp';
import { useTuningStore } from '../tuning/tuningStore';
import type { PretrainedSource } from './types';
import { PretrainedSourceSelect } from './PretrainedSourceSelect';
const source: PretrainedSource = {
id: 'base',
label: '用户基础行走策略',
ready: true,
compatibleTasks: ['Unitree-Go2-Flat', 'Unitree-Go2-ObstacleAvoidance'],
observationSizes: [47, 81, 97],
initialization: {
sourceId: 'a'.repeat(64),
registeredId: 'base',
label: '用户基础行走策略',
manifest: {
source_iteration: 10000,
source_actor_dim: 47,
normalization: 'preserve-source-count/unit-new-features',
artifacts: Object.fromEntries(
['checkpoint', 'onnx', 'env', 'agent'].map((key) => [
key,
{ name: key === 'checkpoint' ? 'model_10000.pt' : key, sha256: 'b'.repeat(64), bytes: 1 },
]),
) as NonNullable<PretrainedSource['initialization']>['manifest']['artifacts'],
},
},
};
beforeEach(() => {
vi.restoreAllMocks();
vi.unstubAllGlobals();
localStorage.clear();
sessionStorage.clear();
useTuningStore.setState({
sessionId: undefined,
sessions: [],
capability: undefined,
connectionState: 'idle',
error: undefined,
});
});
const json = (value: unknown, status = 200) =>
new Response(JSON.stringify(value), { status, headers: { 'Content-Type': 'application/json' } });
it('普通训练可选来源、显示resolved checkpoint/SHA,任务切换保留选择,服务器错误不退回随机', async () => {
const requests: Record<string, unknown>[] = [];
vi.stubGlobal(
'fetch',
vi.fn((url: string, init?: RequestInit) => {
if (url.endsWith('/health'))
return Promise.resolve(
json({
ready: true,
tasks: ['Unitree-Go2-Flat', 'Unitree-Go2-Rough'],
pretrainedSources: [source],
}),
);
if (url.endsWith('/presets')) return Promise.resolve(json({ presets: [] }));
requests.push(JSON.parse(String(init?.body)) as Record<string, unknown>);
return Promise.resolve(json({ error: '基础策略快照SHA不匹配' }, 400));
}),
);
render(<LocalTrainingPanel onPolicyReady={vi.fn()} />);
fireEvent.change(screen.getByLabelText('训练服务访问令牌'), { target: { value: 'token' } });
fireEvent.click(screen.getByRole('button', { name: /^连接$/ }));
const select = await screen.findByLabelText('基础策略');
fireEvent.change(select, { target: { value: 'base' } });
expect(screen.getByText(/model_10000.pt/)).toBeVisible();
expect(screen.getByText(/checkpoint SHA256/)).toHaveTextContent('b'.repeat(64));
fireEvent.change(screen.getByLabelText('训练任务'), { target: { value: 'Unitree-Go2-Rough' } });
expect(select).toHaveValue('base');
expect(screen.getByRole('button', { name: '发起本地训练' })).toBeDisabled();
expect(screen.getByRole('option', { name: /任务不兼容/ })).toBeDisabled();
fireEvent.change(screen.getByLabelText('训练任务'), { target: { value: 'Unitree-Go2-Flat' } });
fireEvent.change(select, { target: { value: 'base' } });
fireEvent.click(screen.getByRole('button', { name: '发起本地训练' }));
expect(await screen.findByRole('alert')).toHaveTextContent('SHA不匹配');
expect(requests).toHaveLength(1);
expect(requests[0].pretrainedSourceId).toBe('base');
expect(requests[0]).not.toHaveProperty('pretrainedCheckpoint');
expect(select).toHaveValue('base');
});
it('自调参面板选择来源只提交注册ID,切换任务保留来源且保持approval模式', async () => {
const requests: Record<string, unknown>[] = [];
vi.stubGlobal(
'fetch',
vi.fn((url: string, init?: RequestInit) => {
if (url.endsWith('/capabilities'))
return Promise.resolve(
json({ ready: true, configured: true, model: 'stub', pretrainedSources: [source] }),
);
if (init?.method === 'POST') {
requests.push(JSON.parse(String(init.body)) as Record<string, unknown>);
return Promise.resolve(json({ error: '基础策略快照文件失效' }, 400));
}
return Promise.resolve(json({ sessions: [] }));
}),
);
render(<TuningApp />);
fireEvent.change(screen.getByLabelText('访问令牌(仅当前标签页)'), {
target: { value: 'token' },
});
fireEvent.click(screen.getByRole('button', { name: '连接/刷新' }));
await screen.findByRole('option', { name: /用户基础行走策略/ });
const select = screen.getByLabelText('基础策略');
fireEvent.change(select, { target: { value: 'base' } });
fireEvent.change(screen.getByLabelText('调参任务'), {
target: { value: 'Unitree-Go2-ObstacleAvoidance' },
});
expect(select).toHaveValue('base');
fireEvent.change(select, { target: { value: 'base' } });
fireEvent.click(screen.getByRole('button', { name: '启动自调参' }));
await waitFor(() => expect(requests).toHaveLength(1));
expect(requests[0]).toMatchObject({
pretrainedSourceId: 'base',
mode: 'approval',
taskId: 'Unitree-Go2-ObstacleAvoidance',
});
expect(await screen.findByText('基础策略快照文件失效')).toBeVisible();
});
it('缺checkpoint注册条目显式错误,不作为可用来源', () => {
render(
<PretrainedSourceSelect
taskId="Unitree-Go2-Flat"
value=""
onChange={vi.fn()}
sources={[
{
id: 'bad',
label: '错误来源',
ready: false,
compatibleTasks: [],
error: '缺少.pt,请配置匹配checkpoint',
},
]}
/>,
);
expect(screen.getByRole('alert')).toHaveTextContent('缺少.pt');
expect(screen.getByRole('option', { name: /验证失败/ })).toBeDisabled();
});
const sourceA: PretrainedSource = { ...source, id: 'a'.repeat(64) };
const sourceB: PretrainedSource = {
...source,
id: 'c'.repeat(64),
// Same registration alias/label is deliberately retained; content B is not A.
initialization: { ...source.initialization!, sourceId: 'c'.repeat(64) },
};
const invalidCatalogs: [string, PretrainedSource[] | undefined][] = [
['同别名A替换为B', [sourceB]],
['A变为not-ready', [{ ...sourceA, ready: false }, sourceB]],
['目录缺项', undefined],
['A变为任务不兼容', [{ ...sourceA, compatibleTasks: ['Unitree-Go2-Rough'] }, sourceB]],
];
for (const panel of ['普通训练', '自调参'] as const) {
for (const [scenario, invalidCatalog] of invalidCatalogs) {
for (const resolution of ['明确从头训练', '明确选择B'] as const) {
it(`${panel}真实刷新:${scenario}保留A并阻止请求,${resolution}后恢复`, async () => {
let catalog: PretrainedSource[] | undefined = [sourceA];
const requests: Record<string, unknown>[] = [];
vi.stubGlobal(
'fetch',
vi.fn((url: string, init?: RequestInit) => {
if (url.endsWith('/health'))
return Promise.resolve(
json({
ready: true,
tasks: ['Unitree-Go2-Flat'],
pretrainedSources: catalog,
}),
);
if (url.endsWith('/capabilities'))
return Promise.resolve(
json({
ready: true,
configured: true,
model: 'stub',
pretrainedSources: catalog,
}),
);
if (url.endsWith('/presets')) return Promise.resolve(json({ presets: [] }));
if (init?.method === 'POST') {
requests.push(JSON.parse(String(init.body)) as Record<string, unknown>);
return Promise.resolve(json({ error: '测试截获请求,未启动训练' }, 400));
}
return Promise.resolve(json({ sessions: [] }));
}),
);
const ordinary = panel === '普通训练';
render(ordinary ? <LocalTrainingPanel onPolicyReady={vi.fn()} /> : <TuningApp />);
fireEvent.change(
screen.getByLabelText(ordinary ? '训练服务访问令牌' : '访问令牌(仅当前标签页)'),
{ target: { value: 'token' } },
);
const refresh = screen.getByRole('button', { name: ordinary ? /^连接$/ : '连接/刷新' });
fireEvent.click(refresh);
await screen.findByRole('option', { name: /用户基础行走策略/ });
const select = screen.getByLabelText('基础策略');
fireEvent.change(select, { target: { value: sourceA.id } });
const start = screen.getByRole('button', {
name: ordinary ? '发起本地训练' : '启动自调参',
});
await waitFor(() => expect(start).toBeEnabled());
catalog = invalidCatalog;
fireEvent.click(refresh);
await screen.findByText(/所选基础策略已失效:/);
await waitFor(() => expect(refresh).toBeEnabled());
expect(select).toHaveValue(sourceA.id);
expect(start).toBeDisabled();
fireEvent.click(start);
expect(requests).toHaveLength(0);
if (resolution === '明确选择B') {
catalog = [sourceB];
fireEvent.click(refresh);
await screen.findByRole('option', { name: /用户基础行走策略.*已验证/ });
await waitFor(() => expect(refresh).toBeEnabled());
expect(select).toHaveValue(sourceA.id);
expect(start).toBeDisabled();
expect(requests).toHaveLength(0);
fireEvent.change(select, { target: { value: sourceB.id } });
} else fireEvent.change(select, { target: { value: '' } });
expect(screen.queryByText(/所选基础策略已失效:/)).not.toBeInTheDocument();
await waitFor(() => expect(start).toBeEnabled());
fireEvent.click(start);
await waitFor(() => expect(requests).toHaveLength(1));
if (resolution === '明确选择B') expect(requests[0].pretrainedSourceId).toBe(sourceB.id);
else expect(requests[0]).not.toHaveProperty('pretrainedSourceId');
if (!ordinary) expect(requests[0].mode).toBe('approval');
});
}
}
}
const uploadCapability = {
enabled: true,
templateId: 'go2-legacy47-v1',
formats: { pt: 256 * 1024 ** 2, onnx: 64 * 1024 ** 2 },
endpoint: '/api/training/pretrained-sources/upload',
};
for (const ordinary of [true, false]) {
for (const outcome of ['pt', 'onnx', 'failure', 'cancel', 'task', 'connection'] as const) {
it(`${ordinary ? '普通' : '自调参'}上传${outcome}:阻止未完成启动并隔离旧epoch,保留旧选择`, async () => {
let finish!: (response: Response) => void;
const uploads: { url: string; init?: RequestInit }[] = [];
const starts: Record<string, unknown>[] = [];
const format = outcome === 'onnx' ? 'onnx' : 'pt';
const uploaded: PretrainedSource = {
...sourceB,
label: `single.${format}`,
initialization: {
...sourceB.initialization!,
label: `single.${format}`,
manifest: {
...sourceB.initialization!.manifest,
sourceFormat: format,
source_iteration: format === 'onnx' ? null : 10000,
contract: 'go2-legacy47-v1',
derived_fields: { normalizer_count: { policy: 'synthetic', value: 1000000 } },
artifacts: {
checkpoint: source.initialization!.manifest.artifacts.checkpoint,
upload: { name: `upload.${format}`, sha256: 'd'.repeat(64), bytes: 3 },
},
},
},
};
vi.stubGlobal(
'fetch',
vi.fn((url: string, init?: RequestInit) => {
if (url.includes('/pretrained-sources/upload?')) {
uploads.push({ url, init });
return new Promise<Response>((resolve) => {
finish = resolve;
});
}
if (url.endsWith('/health') || url.endsWith('/capabilities'))
return Promise.resolve(
json({
ready: true,
configured: true,
tasks: ['Unitree-Go2-Flat', 'Unitree-Go2-Rough'],
pretrainedSources: [sourceA],
pretrainedUpload: uploadCapability,
}),
);
if (url.endsWith('/presets')) return Promise.resolve(json({ presets: [] }));
if (init?.method === 'POST') {
starts.push(JSON.parse(String(init.body)) as Record<string, unknown>);
return Promise.resolve(json({ error: '测试禁止真实训练' }, 400));
}
return Promise.resolve(json({ sessions: [] }));
}),
);
render(ordinary ? <LocalTrainingPanel onPolicyReady={vi.fn()} /> : <TuningApp />);
fireEvent.change(
screen.getByLabelText(ordinary ? '训练服务访问令牌' : '访问令牌(仅当前标签页)'),
{ target: { value: 'token' } },
);
const connect = () =>
fireEvent.click(screen.getByRole('button', { name: ordinary ? /^连接$/ : '连接/刷新' }));
connect();
await screen.findByRole('option', { name: /用户基础行走策略/ });
fireEvent.change(screen.getByLabelText('基础策略'), { target: { value: sourceA.id } });
const input = screen.getByLabelText('选择基础策略文件');
expect(input).toBeDisabled();
fireEvent.click(screen.getByLabelText('确认Go2 legacy47模板'));
const file = new File(['abc'], `single.${format}`);
fireEvent.change(input, { target: { files: [file] } });
await waitFor(() => expect(uploads).toHaveLength(1));
expect(uploads[0].init?.body).toBe(file);
expect(uploads[0].url).toContain('template=go2-legacy47-v1');
const startName = ordinary ? '发起本地训练' : '启动自调参';
expect(screen.getByRole('button', { name: startName })).toBeDisabled();
expect(starts).toHaveLength(0);
if (outcome === 'cancel') fireEvent.click(screen.getByRole('button', { name: '取消上传' }));
if (outcome === 'task')
fireEvent.change(screen.getByLabelText(ordinary ? '训练任务' : '调参任务'), {
target: { value: ordinary ? 'Unitree-Go2-Rough' : 'Unitree-Go2-ObstacleAvoidance' },
});
if (outcome === 'connection') {
fireEvent.change(screen.getByLabelText(ordinary ? '本地训练服务地址' : '训练服务地址'), {
target: { value: 'http://127.0.0.1:9999' },
});
connect();
await screen.findByRole('option', { name: /用户基础行走策略/ });
}
finish(outcome === 'failure' ? json({ error: '不支持该模型' }, 400) : json(uploaded, 201));
if (outcome === 'pt' || outcome === 'onnx') {
await waitFor(() => expect(screen.getByLabelText('基础策略')).toHaveValue(sourceB.id));
expect(screen.getByText(/原文件 SHA256/)).toHaveTextContent('d'.repeat(64));
if (outcome === 'onnx')
expect(screen.getByText(/ONNX统计count合成/)).toHaveTextContent('1000000');
fireEvent.click(screen.getByRole('button', { name: startName }));
await waitFor(() => expect(starts).toHaveLength(1));
expect(starts[0].pretrainedSourceId).toBe(sourceB.id);
if (!ordinary) expect(starts[0].mode).toBe('approval');
} else {
if (outcome === 'failure') {
await screen.findByText(/不支持该模型/);
fireEvent.change(input, { target: { files: [file] } });
await waitFor(() => expect(uploads).toHaveLength(2));
fireEvent.click(screen.getByRole('button', { name: '取消上传' }));
finish(json(uploaded, 201));
}
await waitFor(() => expect(screen.getByLabelText('基础策略')).toHaveValue(sourceA.id));
expect(uploads[0].init?.signal?.aborted).toBe(outcome !== 'failure');
expect(starts).toHaveLength(0);
}
});
}
}
@@ -0,0 +1,123 @@
import { Select } from '../components/ui';
import type { PretrainedInitialization, PretrainedSource } from './types';
import { pretrainedSelectionError } from './pretrainedSelection';
import { PretrainedUpload, type PretrainedUploadConnection } from './PretrainedUpload';
export function PretrainedIdentity({ source }: { source?: PretrainedInitialization }) {
if (!source)
return <p className="text-xs text-text-tertiary">初始化:随机新策略(未选择基础策略)</p>;
return (
<div className="break-all text-xs text-text-secondary" aria-label="基础策略身份">
<p>
初始化:{source.label} · {source.manifest.artifacts.checkpoint.name}
</p>
<p>checkpoint SHA256:{source.manifest.artifacts.checkpoint.sha256}</p>
{source.manifest.artifacts.upload ? (
<>
<p>
上传格式:{source.manifest.sourceFormat} · 原文件:{source.label}
</p>
<p>原文件 SHA256:{source.manifest.artifacts.upload.sha256}</p>
<p>
模板:{source.manifest.contract}
(用户确认缺失的物理语义);继承actor权重,非完整resume。
</p>
{source.manifest.sourceFormat === 'onnx' && (
<p>
ONNX统计count合成:{source.manifest.derived_fields?.normalizer_count?.value ?? '未知'}
; 探索std使用新训练默认,critic/optimizer重新初始化。
</p>
)}
</>
) : (
<p>ONNX SHA256:{source.manifest.artifacts.onnx?.sha256 ?? '未知'}</p>
)}
<p>
来源迭代 {source.manifest.source_iteration ?? '未知'}
;新训练从0开始,critic/optimizer重新初始化;同trial续训保留自身checkpoint。
</p>
</div>
);
}
export function PretrainedSourceSelect({
sources = [],
taskId,
value,
onChange,
disabled = false,
upload,
}: {
sources?: PretrainedSource[];
taskId: string;
value: string;
onChange(value: string): void;
disabled?: boolean;
upload?: PretrainedUploadConnection;
}) {
const selected = sources.find((source) => source.id === value);
const selectionError = pretrainedSelectionError(sources, taskId, value);
return (
<div className="space-y-1">
<label className="block text-xs text-text-secondary">
基础策略
<Select
aria-label="基础策略"
className="w-full"
value={value}
disabled={disabled}
onChange={(event) => onChange(event.target.value)}
>
<option value="">不选择(随机初始化)</option>
{value && !selected && (
<option value={value} disabled>
所选基础策略已失效({value})
</option>
)}
{sources.map((source) => (
<option
key={source.id}
value={source.id}
disabled={!source.ready || !source.compatibleTasks.includes(taskId)}
>
{source.label}
{!source.ready
? '(验证失败)'
: !source.compatibleTasks.includes(taskId)
? '(任务不兼容)'
: '(已验证)'}
</option>
))}
</Select>
</label>
{upload && (
<PretrainedUpload
key={JSON.stringify([upload.endpoint, upload.token, upload.revision, taskId])}
connection={upload}
disabled={disabled}
/>
)}
{!sources.length && (
<p className="text-xs">尚无基础策略,可直接上传单个文件;不选择则随机初始化。</p>
)}
{sources
.filter((source) => source.error)
.map((source) => (
<p role="alert" key={source.id}>
{source.label}:{source.error}
</p>
))}
{selected && (
<p className="text-xs">
兼容:{selected.compatibleTasks.join(' / ')};观测 {selected.observationSizes?.join('/')}{' '}
→ 12动作。只继承actor权重,新训练从0开始。
</p>
)}
{selectionError ? (
<p role="alert">{selectionError}</p>
) : (
<PretrainedIdentity source={selected?.initialization} />
)}
</div>
);
}
@@ -0,0 +1,111 @@
import { useEffect, useRef, useState } from 'react';
import { Button } from '../components/ui';
import { LocalTrainingClient } from './LocalTrainingClient';
import type { PretrainedSource } from './types';
export interface PretrainedUploadConnection {
endpoint: string;
token: string;
revision: number;
enabled: boolean;
onUploaded(source: PretrainedSource): void;
onBusyChange(busy: boolean): void;
}
/** Remounted for each connection/task epoch by the shared source selector. */
export function PretrainedUpload({
connection,
disabled,
}: {
connection: PretrainedUploadConnection;
disabled: boolean;
}) {
const [confirmed, setConfirmed] = useState(false);
const [pending, setPending] = useState(false);
const [message, setMessage] = useState('');
const [error, setError] = useState('');
const request = useRef<AbortController | null>(null);
const onBusyChange = connection.onBusyChange;
useEffect(
() => () => {
request.current?.abort();
onBusyChange(false);
},
[onBusyChange],
);
const upload = async (file: File) => {
if (!confirmed || disabled || pending || !connection.enabled) return;
const controller = new AbortController();
request.current = controller;
setPending(true);
onBusyChange(true);
setError('');
setMessage(`正在上传并验证 ${file.name},请稍候…`);
try {
const source = await new LocalTrainingClient(
connection.endpoint,
connection.token,
).uploadPretrained(file, 'go2-legacy47-v1', controller.signal);
if (controller.signal.aborted) return;
if (!source.ready || !source.initialization)
throw new Error(source.error ?? '上传来源未通过验证');
connection.onUploaded(source);
setMessage(`已选择 ${source.label};仅继承策略权重,不是完整训练resume。`);
} catch (value) {
if (!controller.signal.aborted) {
setMessage('');
setError(
`${value instanceof Error ? value.message : String(value)};保留原基础策略选择,未启动训练。`,
);
}
} finally {
if (!controller.signal.aborted) {
request.current = null;
setPending(false);
onBusyChange(false);
}
}
};
const cancel = () => {
request.current?.abort();
request.current = null;
setPending(false);
onBusyChange(false);
setMessage('已取消等待,保留原基础策略选择;服务若已完成验证,可刷新目录查看。');
setError('');
};
return (
<div className="space-y-2 rounded border border-border p-2 text-xs">
<p>上传基础策略:单个.pt(≤256MiB)或.onnx(≤64MiB),无需服务器路径或配套文件。</p>
<label className="flex gap-2">
<input
type="checkbox"
aria-label="确认Go2 legacy47模板"
checked={confirmed}
disabled={disabled || pending || !connection.enabled}
onChange={(event) => setConfirmed(event.target.checked)}
/>
按Go2 legacy47观测及FL/FR/RL/RR关节顺序解释文件;缺失的物理语义由我确认。
</label>
<p>
仅支持47维Go2行走actor;ONNX仅继承推理网络,统计count合成,探索/价值网络/优化器重新初始化。
</p>
<input
type="file"
accept=".pt,.onnx"
aria-label="选择基础策略文件"
disabled={!confirmed || disabled || pending || !connection.enabled}
onChange={(event) => {
const file = event.target.files?.[0];
event.target.value = ''; // Same-file retry must still dispatch change.
if (file) void upload(file);
}}
/>
{!connection.enabled && <p>请先连接支持单文件上传的训练服务。</p>}
{pending && <Button onClick={cancel}>取消上传</Button>}
{message && <p role="status">{message}</p>}
{error && <p role="alert">{error}</p>}
</div>
);
}
@@ -0,0 +1,63 @@
import { TrainingMetricHistory } from './TrainingMetricHistory';
const header = (i: number) => `Learning iteration ${i} / 1000`;
it('解析多个迭代全部指标、重复poll去重、同迭代分批补全', () => {
const h = new TrainingMetricHistory();
const logs = [
header(1),
'Mean value loss: 1e-3',
'Mean surrogate loss: -0.1',
header(2),
'Mean value loss: 2',
];
expect(h.update('a', logs)).toBe(true);
expect(h.update('a', logs)).toBe(false);
expect(h.series()[0].points.map((p) => [p.step, p.value])).toEqual([
[1, 0.001],
[2, 2],
]);
expect(
h.update('a', [...logs, 'Mean entropy loss: -3', 'Mean reward: 4', 'Mean episode length: 50']),
).toBe(true);
expect(h.series()).toHaveLength(5);
expect(h.series().find((s) => s.tag === '平均奖励')!.points[0]).toEqual({
step: 2,
value: 4,
wallTime: 0,
});
});
it('滚动截断按重叠上下文补全,无header/overlap不猜迭代,忽略非有限/缺失/其它数字', () => {
const h = new TrainingMetricHistory();
h.update('a', [header(3), 'Mean value loss: 3']);
h.update('a', ['Mean value loss: 3', 'Mean reward: 7']);
expect(h.series()[1].points[0].step).toBe(3);
h.update('a', ['Mean entropy loss: 5']);
expect(h.series()).toHaveLength(2);
h.update('a', [
header(4),
'Mean value loss: NaN',
'Mean reward: Infinity',
'Mean reward: 1e999',
'Mean reward: 3 ms',
'Total timesteps: 40',
'Iteration time: 2',
]);
expect(h.series()[0].points).toHaveLength(1);
h.update('b', ['Mean reward: 1']);
expect(h.series()).toEqual([]);
h.update('b', [header(0), 'Mean reward: .5']);
expect(h.series()[0].points[0].step).toBe(0);
});
it('容量和去重索引有界,旧重发不挤掉最新迭代,job切换清空', () => {
const h = new TrainingMetricHistory();
for (let i = 0; i < 650; i++) h.update('a', [header(i), `Mean reward: ${i}`]);
const points = h.series()[0].points;
expect(points).toHaveLength(500);
expect(points[0].step).toBe(150);
expect(points.at(-1)!.step).toBe(649);
h.update('a', [header(1), 'Mean reward: 999']);
expect(h.series()[0].points).toEqual(points);
h.update('new-job', [header(0), 'Mean value loss: 1']);
expect(h.series()[0].points).toHaveLength(1);
expect(() => new TrainingMetricHistory(1)).toThrow();
});
@@ -0,0 +1,93 @@
import type { ScalarSeries } from './types';
export const TRAINING_METRICS = {
value: '价值损失',
surrogate: '策略损失',
entropy: '熵损失',
reward: '平均奖励',
episodeLength: '平均回合长度',
} as const;
type Metric = keyof typeof TRAINING_METRICS;
const METRIC_PATTERN =
/^\s*Mean (value loss|surrogate loss|entropy loss|reward|episode length):\s*([-+]?(?:\d+(?:\.\d*)?|\.\d+)(?:e[-+]?\d+)?)\s*$/i;
const KEYS: Record<string, Metric> = {
'value loss': 'value',
'surrogate loss': 'surrogate',
'entropy loss': 'entropy',
reward: 'reward',
'episode length': 'episodeLength',
};
/** Bounded iteration merge. Headerless lines are accepted only with proven suffix/prefix overlap.
* On reconnect without overlap, orphan scalars are skipped rather than guessed onto a new step. */
export class TrainingMetricHistory {
private readonly rows = new Map<number, Partial<Record<Metric, number>>>();
private previous: { line: string; iteration?: number }[] = [];
private jobId?: string;
constructor(private readonly capacity = 500) {
if (!Number.isInteger(capacity) || capacity < 2) throw new Error('指标容量必须至少为2');
}
update(jobId: string, logs: readonly string[]): boolean {
let changed = false;
if (this.jobId !== jobId) {
this.rows.clear();
this.previous = [];
this.jobId = jobId;
changed = true;
}
const lines = logs
.flatMap((line) => line.split('\n'))
.slice(-1000)
// rsl_rl uses ANSI SGR colors around iteration headers.
// eslint-disable-next-line no-control-regex
.map((line) => line.replace(/\u001b\[[0-9;]*m/g, '').trim());
let iteration: number | undefined;
for (let overlap = Math.min(lines.length, this.previous.length); overlap > 0; overlap--) {
const start = this.previous.length - overlap;
if (lines.slice(0, overlap).every((line, i) => this.previous[start + i].line === line)) {
iteration = this.previous[start].iteration;
break;
}
}
const contexts: typeof this.previous = [];
for (const line of lines) {
const header = /\bLearning iteration\s+(\d+)\s*\/\s*\d+/i.exec(line);
if (header) {
const step = Number(header[1]);
iteration = Number.isSafeInteger(step) ? step : undefined;
}
contexts.push({ line, iteration });
const scalar = METRIC_PATTERN.exec(line);
if (iteration === undefined || !scalar) continue;
const value = Number(scalar[2]);
if (!Number.isFinite(value)) continue;
const key = KEYS[scalar[1].toLowerCase()];
if (!this.rows.has(iteration) && this.rows.size >= this.capacity) {
const oldest = Math.min(...this.rows.keys());
if (iteration < oldest) continue;
this.rows.delete(oldest);
}
const row = this.rows.get(iteration) ?? {};
if (row[key] !== value) {
row[key] = value;
changed = true;
}
this.rows.set(iteration, row);
}
this.previous = contexts;
return changed;
}
series(): ScalarSeries[] {
const rows = [...this.rows].sort(([a], [b]) => a - b);
return Object.entries(TRAINING_METRICS)
.map(([key, tag]) => ({
tag,
points: rows.flatMap(([step, row]) =>
row[key as Metric] === undefined
? []
: [{ step, wallTime: 0, value: row[key as Metric]! }],
),
}))
.filter((series) => series.points.length > 0);
}
}
@@ -0,0 +1,36 @@
import { fireEvent, render, screen } from '@testing-library/react';
import { TrainingMetricsPanel } from './TrainingMetricsPanel';
import { ScalarChart } from '../components/charts/ScalarChart';
vi.mock('../components/charts/ScalarChart', () => ({
ScalarChart: vi.fn(({ title }: { title: string }) => (
<div data-testid="metric-chart">{title}曲线</div>
)),
}));
it('折叠不挂载图表,价值/策略/综合独立缩放,job切换清空历史', () => {
const logs = [
'Learning iteration 1 / 10',
'Mean value loss: 1',
'Mean surrogate loss: -2',
'Mean entropy loss: -3',
'Mean reward: 4',
'Mean episode length: 50',
];
const view = render(<TrainingMetricsPanel jobId="a" logs={logs} />);
expect(screen.queryByTestId('metric-chart')).not.toBeInTheDocument();
fireEvent.click(screen.getByRole('button', { name: /训练指标趋势/ }));
expect(screen.getByText('价值损失曲线')).toBeInTheDocument();
vi.mocked(ScalarChart).mockClear();
view.rerender(<TrainingMetricsPanel jobId="a" logs={[...logs]} />);
expect(ScalarChart).not.toHaveBeenCalled();
fireEvent.click(screen.getByRole('tab', { name: '策略损失' }));
expect(screen.getByText('策略损失曲线')).toBeInTheDocument();
fireEvent.click(screen.getByRole('tab', { name: '综合' }));
expect(screen.getAllByTestId('metric-chart')).toHaveLength(5);
expect(screen.getByText('平均回合长度曲线')).toBeInTheDocument();
view.rerender(<TrainingMetricsPanel jobId="b" logs={[]} />);
expect(screen.queryByTestId('metric-chart')).not.toBeInTheDocument();
expect(screen.getByText(/尚无带迭代编号/)).toBeInTheDocument();
fireEvent.click(screen.getByRole('button', { name: /训练指标趋势/ }));
expect(screen.queryByRole('tabpanel')).not.toBeInTheDocument();
});
@@ -0,0 +1,81 @@
import { memo, useEffect, useState } from 'react';
import { ScalarChart } from '../components/charts/ScalarChart';
import { TrainingMetricHistory, TRAINING_METRICS } from './TrainingMetricHistory';
import type { ScalarSeries } from './types';
const MetricChart = memo(function MetricChart({ series }: { series: ScalarSeries }) {
return (
<div>
<p className="mb-1 text-xs text-text-secondary">
{series.tag} · 最新原值 {series.points.at(-1)!.value.toPrecision(4)}
</p>
<ScalarChart series={[series]} smoothing={0.4} title={series.tag} xLabel="Iteration" />
</div>
);
});
export const TrainingMetricsPanel = memo(function TrainingMetricsPanel({
jobId,
logs,
}: {
jobId: string;
logs: readonly string[];
}) {
const [history] = useState(() => new TrainingMetricHistory());
const [series, setSeries] = useState<ScalarSeries[]>([]);
const [open, setOpen] = useState(false);
const [tab, setTab] = useState('value');
useEffect(() => {
if (history.update(jobId, logs)) setSeries(history.series());
}, [history, jobId, logs]);
const selected = series.filter(
(item) =>
tab === 'all' ||
item.tag === (tab === 'value' ? TRAINING_METRICS.value : TRAINING_METRICS.surrogate),
);
return (
<section className="mt-3 min-w-0 rounded-lg border border-border">
<button
type="button"
className="w-full p-2 text-left text-xs text-text-primary"
aria-expanded={open}
onClick={() => setOpen(!open)}
>
{open ? '▾' : '▸'} 训练指标趋势
</button>
{open && (
<div className="min-w-0 space-y-2 p-2 pt-0">
<div role="tablist" aria-label="训练指标" className="flex gap-2">
{[
['value', '价值损失'],
['surrogate', '策略损失'],
['all', '综合'],
].map(([id, label]) => (
<button
key={id}
type="button"
role="tab"
aria-selected={tab === id}
className={`rounded px-2 py-1 text-xs ${tab === id ? 'bg-accent/10 text-accent' : 'text-text-secondary'}`}
onClick={() => setTab(id)}
>
{label}
</button>
))}
</div>
<div role="tabpanel" aria-label="训练指标曲线" className="space-y-2">
{selected.length ? (
selected.map((item) => <MetricChart key={item.tag} series={item} />)
) : (
<p className="text-xs text-text-tertiary">尚无带迭代编号的指标日志</p>
)}
</div>
<p className="text-[10px] text-text-tertiary">
各指标独立纵轴;EMA 0.4
仅用于曲线,悬停显示原值。最多保留最近500个有指标的迭代,重连仅恢复服务端日志尾部。
</p>
</div>
)}
</section>
);
});
@@ -0,0 +1,21 @@
import type { PretrainedSource } from './types';
/** A nonempty selection is an initialization intent, never a random-training fallback. */
export function pretrainedSelectionError(
sources: readonly PretrainedSource[] | undefined,
taskId: string,
selectedId: string,
): string | undefined {
if (!selectedId) return undefined;
const source = sources?.find((item) => item.id === selectedId);
const reason = !source
? '目录缺少该内容ID'
: !source.ready
? '来源未通过验证'
: !source.compatibleTasks.includes(taskId)
? '来源与当前任务不兼容'
: undefined;
return reason
? `所选基础策略已失效:${reason}。请明确选择其他有效基础策略,或选择“不选择(随机初始化)”;不会自动退回随机初始化。`
: undefined;
}
@@ -0,0 +1,18 @@
import { expect, it } from 'vitest';
import { trainingLosses } from './trainingLosses';
it('提取最新有限PPO损失,兼容无指标旧日志', () => {
expect(
trainingLosses([
'Mean value loss: 1.2',
'Mean surrogate loss: -0.03',
'Mean entropy loss: 1e-3',
'Mean value loss: 0.1',
'Mean value loss: Infinity',
]),
).toEqual([
{ label: '价值损失', value: 0.1 },
{ label: '策略损失', value: -0.03 },
{ label: '熵损失', value: 0.001 },
]);
expect(trainingLosses(['Starting training'])).toEqual([]);
});
@@ -0,0 +1,14 @@
/** rsl_rl console scalars, newest finite value per loss; old servers need no new endpoint. */
export function trainingLosses(logs: readonly string[]): { label: string; value: number }[] {
const values = new Map<string, number>();
for (const line of logs) {
const match = /Mean (value|surrogate|entropy) loss:\s*([-+\d.eE]+)/.exec(line);
if (match && Number.isFinite(Number(match[2]))) values.set(match[1], Number(match[2]));
}
const labels: Record<string, string> = {
value: '价值损失',
surrogate: '策略损失',
entropy: '熵损失',
};
return Array.from(values, ([key, value]) => ({ label: labels[key], value }));
}
+85 -10
View File
@@ -1,13 +1,75 @@
import type { PolicyDeployment, TrainingTerrain } from '../rl/deployment';
export interface TrainingParameter {
min: number;
max: number;
default: number;
integer?: boolean;
}
export interface TrainingTaskMetadata {
id: string;
name: string;
browserCompatible: boolean;
terrainPresets: string[];
terrainParameters: Record<string, TrainingParameter>;
sensorTypes: string[];
sensorModes?: string[];
sensorParameters: Record<string, TrainingParameter>;
mapSyncScope: string;
}
export type TrainingJobState = 'queued' | 'running' | 'succeeded' | 'failed' | 'cancelled';
export type TrainingDevice = 'cpu' | 'gpu';
export type WandbMode = 'offline' | 'online' | 'disabled';
export interface PretrainedInitialization {
sourceId: string;
registeredId: string;
label: string;
manifest: {
source_iteration: number | null;
sourceFormat?: 'pt' | 'onnx';
contract?: string;
derived_fields?: {
normalizer_count?: { policy: string; value: number };
exploration_std?: { policy: string; value: number };
};
source_actor_dim: number;
normalization: string;
artifacts: { checkpoint: PretrainedArtifact } & Partial<
Record<'upload' | 'onnx' | 'env' | 'agent', PretrainedArtifact>
>;
};
}
export interface PretrainedArtifact {
name: string;
sha256: string;
bytes: number;
}
export interface PretrainedUploadCapability {
enabled: boolean;
templateId: 'go2-legacy47-v1';
formats: { pt: number; onnx: number };
endpoint: string;
}
export interface PretrainedSource {
id: string;
label: string;
ready: boolean;
compatibleTasks: string[];
observationSizes?: number[];
initialization?: PretrainedInitialization;
error?: string;
}
export interface TrainingServerInfo {
pretrainedSources?: PretrainedSource[];
pretrainedUpload?: PretrainedUploadCapability;
version: string;
ready: boolean;
trainerRoot: string;
python: string;
tasks: string[];
taskMetadata?: TrainingTaskMetadata[];
activeJobId?: string;
resourceOwner?: string;
tuning?: TuningCapability;
@@ -15,6 +77,11 @@ export interface TrainingServerInfo {
}
export interface TrainingRequest {
customTerrainBoxes?: TrainingTerrain;
terrainPreset?: string;
terrainParams?: Record<string, number>;
sensorType?: 'raycast';
sensorCfg?: Partial<import('../rl/deployment').ObstacleSensorConfig>;
taskId: string;
numEnvs: number;
maxIterations: number;
@@ -24,9 +91,12 @@ export interface TrainingRequest {
gpuIds: number[];
wandbMode: WandbMode;
rewardPresetId?: string;
pretrainedSourceId?: string;
}
export interface TrainingJob {
pretrained?: PretrainedInitialization;
deployment?: PolicyDeployment;
id: string;
state: TrainingJobState;
taskId: string;
@@ -60,16 +130,11 @@ export interface RewardConfiguration {
params: Record<string, number>;
}
export interface ObjectiveWeights {
velocity_tracking: number;
action_smoothness: number;
posture_stability: number;
fall_avoidance: number;
foot_slip: number;
energy: number;
}
export type ObjectiveWeights = Record<string, number>;
export interface TuningCapability {
pretrainedSources?: PretrainedSource[];
pretrainedUpload?: PretrainedUploadCapability;
ready: boolean;
configured: boolean;
apiKeyConfigured: boolean;
@@ -79,7 +144,12 @@ export interface TuningCapability {
}
export interface TuningCreateRequest {
taskId: 'Unitree-Go2-Flat';
pretrainedSourceId?: string;
taskId: 'Unitree-Go2-Flat' | 'Unitree-Go2-ObstacleAvoidance';
taskConfig?: Pick<
TrainingRequest,
'terrainPreset' | 'terrainParams' | 'sensorType' | 'sensorCfg' | 'customTerrainBoxes'
> & { seed?: number };
mode: TuningMode;
runName: string;
numEnvs: number;
@@ -160,7 +230,11 @@ export interface TuningSession {
mode: TuningMode;
createdAt: string;
updatedAt: string;
config: TuningCreateRequest & { rungs: number[]; promote: number[] };
config: TuningCreateRequest & {
rungs: number[];
promote: number[];
pretrained?: PretrainedInitialization;
};
objectiveWeights: ObjectiveWeights;
message: string;
currentTrialId?: string;
@@ -194,6 +268,7 @@ export interface TuningMetricsResponse {
export interface RewardPreset {
id: string;
taskId: string;
name: string;
sessionId: string;
trialId: string;
@@ -25,6 +25,7 @@ import {
formatMetric,
mergeRewardPatch,
OBJECTIVE_META,
OBSTACLE_OBJECTIVE_META,
PARAMETER_BY_PATH,
rewardConfigurationDiff,
} from './domain';
@@ -223,7 +224,9 @@ function ExpectedImpact({ proposal }: { proposal: TuningProposal }) {
return (
<div className="grid gap-1 sm:grid-cols-2">
{values.map(([key, value]) => {
const label = OBJECTIVE_META.find((item) => item.key === key)?.label ?? key;
const label =
[...OBJECTIVE_META, ...OBSTACLE_OBJECTIVE_META].find((item) => item.key === key)?.label ??
key;
return (
<div key={key} className="rounded border border-border bg-app px-2 py-1.5">
<p className="text-[8px] uppercase tracking-wider text-text-tertiary">{label}</p>

Some files were not shown because too many files have changed in this diff Show More