feat(training): release V0.9.1 避障训练与基础策略迁移
This commit is contained in:
@@ -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文件
|
||||
@@ -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、缺项及兼容性变更的真实组件刷新回归,验证失效时零启动请求及用户明确选择后的请求体。
|
||||
|
||||
Generated
+2
-2
@@ -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
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "mujoco-web-platform",
|
||||
"version": "0.8.3",
|
||||
"version": "0.9.1",
|
||||
"description": "基于 MuJoCo WebAssembly 的本地机器人仿真与控制平台",
|
||||
"private": true,
|
||||
"type": "module",
|
||||
|
||||
@@ -0,0 +1,188 @@
|
||||
# Go2 前视射线避障部署契约 v1
|
||||
|
||||
## API
|
||||
|
||||
默认任务白名单:`Unitree-Go2-Flat`、`Unitree-Go2-Rough`、`Unitree-Go2-ObstacleAvoidance`。
|
||||
`GET /api/training/health` 的 `tasks` 保持字符串列表,新增 `taskMetadata`:每项含 `id/name/browserCompatible/terrainPresets/terrainParameters/sensorTypes/sensorParameters/mapSyncScope`。参数元数据含 `min/max/default`,计数含 `integer:true`。
|
||||
|
||||
`POST /api/training/jobs` 保留原参数,并接受:
|
||||
|
||||
```json
|
||||
{
|
||||
"taskId": "Unitree-Go2-ObstacleAvoidance",
|
||||
"numEnvs": 1024,
|
||||
"maxIterations": 1000,
|
||||
"seed": 42,
|
||||
"runName": "obstacle-navigation",
|
||||
"device": "gpu",
|
||||
"gpuIds": [0],
|
||||
"wandbMode": "offline",
|
||||
"terrainPreset": "discrete_obstacles",
|
||||
"terrainParams": {
|
||||
"size": 12,
|
||||
"obstacle_count": 24,
|
||||
"obstacle_height_min": 0.2,
|
||||
"obstacle_height_max": 0.6,
|
||||
"spacing": 1.2,
|
||||
"friction": 0.8
|
||||
},
|
||||
"sensorType": "raycast",
|
||||
"sensorCfg": { "fov": 90, "maxDistance": 4, "safetyDistance": 0.5, "avoidanceWeight": 2 }
|
||||
}
|
||||
```
|
||||
|
||||
只实现 `raycast`,**不是 RGB 或真实深度相机**;`camera_depth` 明确拒绝。`sensorCfg.type` 可选但只能是 `raycast`。默认观测数量32;显式multi模式固定48,不接受任意数量(见文末多层契约)。预设与字段严格白名单校验,布尔、非有限、超范围值拒绝。地形 `size` 8–24m、障碍物数1–100,实际数量受 `spacing` 容量限制。全部可选参数与范围见 `task_config.py`。普通训练的奖励 preset 仍仅允许 Flat;Obstacle专属自调参见文末。
|
||||
|
||||
创建和查询 job 返回 `deployment`,即该作业的确定性部署配置;列表中的 job 也携带此字段。不新增端点;ONNX 下载仍为 `/api/training/jobs/{id}/artifacts/policy.onnx`。自定义配置由服务写入作业目录 `training_config.json`,以参数数组 `--task-config <server-owned-path>` 传入训练器,绝不使用 shell。训练器再次校验配置与seed。
|
||||
|
||||
`policy.onnx` 的 `platform_deployment` metadata 字符串是同一对象的JSON编码,旁边同时导出 `deployment.json`。ONNX中包含Actor观测归一化,不要在前端再套一层运行均值/方差。其他标准 mjlab metadata 保留。旧 Rough 实际 Actor 为234维(47+17×11 downward height scan),`browserCompatible:false`;前端必须禁用一键加载,不能按 Flat 处理。没有新配置的旧 Flat/tuning 保持原环境行为。
|
||||
|
||||
## 地图:精确碰撞布局,不是同名编辑器高度场
|
||||
|
||||
`deployment.terrain` 为 `boxes-v1`:`size/friction/boxes/spawn/spawnQuaternion/target/actualObstacleCount/approximation`。
|
||||
|
||||
- 单块正方形地图,世界原点在地图中心,+X前、+Y左、+Z上,米/弧度。每个 box 的 `pos` 是世界中心,`size` **为半尺寸**,`yaw=0`。底板占 `z=[-0.2,0]`。无额外无限地面/边墙。
|
||||
- 直接用导出的boxes构造碰撞几何,**不要重新运行浏览器的同名地形生成器**。`rough/wave` 为16×16有界box离散化,`pyramid_stairs` 为四层box;`approximation:true` 必须在UI标识训练专用近似布局。
|
||||
- 障碍物坐标由 `random.Random(seed)` 对确定性格点洗牌;宽深均0.4m、高度按范围随机,实际数量受spacing与地图容量限制。两端保留平坦出生/目标条带。无需在JS复刻Python RNG,布局本身是权威结果。最多257个boxes。
|
||||
- `spawn=[-size/2+1,0,0.32]`、四元数 `wxyz=[1,0,0,0]`、目标 `[size/2-1,0]` 是部署及固定评估的参考起终点。训练时每个episode在同一连通自由区域内重新采样起终点(障碍/边界净空>0.55m、间距>=2m)并随机化初始yaw,初始速度为0;不会跨不可通行墙体抽取目标。MuJoCo地形原点仍为参考spawn,随机复位直接写入世界坐标。
|
||||
- mjlab采用1×1 patch,同一张地图用于多个互相独立的Warp环境世界;env origin就是出生点而非地图中心。目标在每个env中用 `env_origin.xy+[size-2,0]`,不是错误地再加地图中心。
|
||||
- terrain摩擦为 `[friction,0.005,0.0001]`,priority=1、condim=3;足端startup滑动摩擦固定为相同值。浏览器应保持匹配的摩擦及足端接触配置。
|
||||
- “同步当前场景地图”导出已应用场景的 `custom_boxes` 权威布局,细则见下节;预设生成模式仍保持原行为。
|
||||
|
||||
## 已应用场景 → custom_boxes
|
||||
|
||||
请求使用 `terrainPreset: "custom_boxes"`、`customTerrainBoxes: <完整boxes-v1对象>`;`terrainParams` 可省略,若提供则只能是与布局完全一致的 `{size, friction}`。服务和训练入口均重新验证,不读取客户端路径或XML,不运行地形RNG、不降级预设。作业目录task JSON、deployment JSON及ONNX metadata中的terrain使用同一布局。health的terrainPresets公布能力。
|
||||
|
||||
- 浏览器从成功编译的全部sceneMaps静态碰撞 `geom_xpos/geom_xmat/geom_size` 提取,包括工程地图路径、重复实例和嵌套body变换;不使用Three装饰包围盒、编辑器RNG或不完整XML。动态机器人、无碰撞装饰排除;未声明地图身份的其他静态碰撞报错,不能漏障碍。
|
||||
- box世界半尺寸为 `abs(R)*halfsize`;sphere、capsule、ellipsoid、cylinder使用解析精确世界AABB。**旋转box与非box转换会扩大碰撞占据区域**。mesh/hfield明确拒绝;项目地图原有include导入限制不变,已经编译模型的include变换不重新解析。
|
||||
- 固定floor为中心 `[0,0,-.1]`、半尺寸 `[size/2,size/2,.1]`。仅明确命名floor/ground/flat的朝上水平z=0 plane,以及顶面z≈0的水平支撑box归并。地下/坑底/非零高度或倾斜plane拒绝,不能静默填坑。其他障碍原坐标保留。底板覆盖整个正方形范围,可能填补支撑面间空白;UI明确说明标准化,不保证与原场景几何同构。
|
||||
- 最多256障碍+1底板,超量报错不截断;世界原点不变,尺寸8–24m,取覆盖应用地图/几何的正方形,不clamp或平移几何。XY必须在floor内,Z界限[-.2,12],所有半尺寸严格>0(包括拒绝负零),有限数值。摩擦必须统一为 `[f,.005,.0001]` 且f为.2–2,混合值报错要求用户先显式统一,不能静默覆盖。
|
||||
- `approximation` 对custom必须为true(AABB及标准底板转换);`actualObstacleCount` 必须等于boxes数减一。字段严格白名单,不信任声明值。floor必须标准。出生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均保持原协议。
|
||||
@@ -0,0 +1,178 @@
|
||||
# Go2 基础策略迁移(普通训练、自调参、CLI)
|
||||
|
||||
普通训练和自调参均可选择经服务验证的基础策略。浏览器持续导航不再因20秒截止而停用;训练/三seed评估协议不变。不新增高程图,不声称短训已学会复杂障碍导航。
|
||||
|
||||
## 前端直接上传单个文件(推荐)
|
||||
|
||||
已安装训练依赖后,在仓库根目录执行一条默认启动命令:
|
||||
|
||||
```bash
|
||||
.venv/bin/python training_server/server.py
|
||||
```
|
||||
|
||||
1. 打开普通训练面板或自调参工作台,将终端显示的令牌填入访问令牌,连接`http://127.0.0.1:8765`。
|
||||
2. 在“基础策略”勾选“确认Go2 legacy47模板”,点击“选择基础策略文件”,选择单个`.pt`或`.onnx`;无需注册JSON、服务器文件路径、YAML或其他sidecars。
|
||||
3. 等待上传/验证结束;完成后自动选中该文件的内容ID,检查文件名、格式、原文件SHA及初始化说明。上传中无法启动;可取消等待,同一文件可再次选择重试。
|
||||
4. 普通训练选择Flat47或Obstacle81/97,再点击“发起本地训练”;自调参保留逐轮审批/全自动选择,点击“启动自调参”。无DeepSeek配置时必须明确勾选Optuna fallback才能启动自调参;上传本身不调用Agent。
|
||||
5. 失败、取消、切换任务或更换连接不会把旧基础策略默默清空;失效/不兼容选择会阻止启动。若要从头训练,明确选择“不选择(随机初始化)”。
|
||||
|
||||
.pt继承actor及其归一化/探索参数;ONNX只继承确定性推理网络,缺失的count采用合成1,000,000、探索std采用新训练默认1.0,两者critic/optimizer均全新、iteration0。不是恢复原PPO训练。旧管理员注册功能仅为可选高级入口,见下文。
|
||||
|
||||
### HTTP与安全契约
|
||||
|
||||
默认启动服务即可上传,**不需要管理员JSON、服务器路径或邻接文件**。旧注册来源仍兼容,见下文“旧CLI/注册来源”。HTTP接口沿用Host/Origin检查及Bearer令牌:
|
||||
|
||||
`POST /api/training/pretrained-sources/upload?format=pt|onnx&template=go2-legacy47-v1&name=<URL编码显示名>`
|
||||
|
||||
- body为原始文件字节,`Content-Type: application/octet-stream`,唯一`Content-Length`。不是JSON/base64/multipart/ZIP;普通JSON请求仍限制128KiB。
|
||||
- 必须显式确认模板,不能只凭47维shape自动证明物理语义。模板按Go2 FL/FR/RL/RR各hip/thigh/calf解释关节,以及本文legacy47观测、50Hz、PD/default/action_scale。文件缺失的坐标系/噪声等语义来自用户确认,不声称验证了源env.yaml。
|
||||
- `.pt`上限256MiB、ONNX上限64MiB;单个上传/解析槽,接收总限60秒、单次阻塞最多10秒;累计快照最多32个/2GiB(包括旧来源),额外预留16MiB派生产物。临时目录由服务随机生成,显示名去路径/控制字符。接收中断/超时/校验失败清理临时文件,不登记可选来源。崩溃遗留临时目录不列入目录且计入容量,需本机维护者清理。
|
||||
- CPU子进程解析:60秒墙钟/40秒CPU、8GiB虚拟地址、16MiB输出文件、64文件描述符、无core dump、CPU单线程。`.pt`只用`weights_only=True`,检查actor完整键/shape/dtype/有限值、normalizer var/std/count、探索std及已知嵌入语义metadata;缺actor的ZIP/任意checkpoint拒绝。不是针对本机同账户恶意写入或原生库漏洞的完整沙箱。
|
||||
- 单ONNX只支持opset17/18、自包含float32 `[1,47]→[1,12]` 的精确9节点链:Sub、Div、Gemm/ELU/Gemm/ELU/Gemm/ELU/Gemm。校验每条连边、算子属性、initializer、I/O、关节/观测/PD/action metadata,拒绝外部data/custom op/支路/重排/额外算子;随后CPU ORT对48个随机+物理probe对照重建actor(atol=rtol=2e-5)。不承诺任意ONNX可恢复PPO。
|
||||
- `.pt`保留原count和探索std;ONNX只有确定性网络/mean/denominator,以已知模板`epsilon=.01`推导`std=denominator-.01>0`及`var=std²`,**合成count=1,000,000**,探索std使用经检查的目标默认1.0。两者critic/optimizer fresh、iteration0;ONNX source_iteration未知,不能称为恢复原.pt训练。
|
||||
- 成功201返回与`health.pretrainedSources[]`同形的单个source record。ID为SHA256(`template:format:原文件SHA256`),文件名不参与;同时保存原文件SHA及服务派生`actor.pt` SHA。manifest明确`sourceFormat`、`derived_fields`、`template_confirmation`与`verified_facts`,不把派生artifact伪装成用户原checkpoint。
|
||||
- 后续普通job和tuning session继续只传`pretrainedSourceId`,复用原binding/immutable SHA/CLI初始化协议。受控`actor.pt`加`--pretrained-upload-manifest`加载,不扫描邻接文件;81/97首层新增列为零、统计0/1/1且可学习。rung仍只resume自身checkpoint,保存的已更新count不会重置为1e6。
|
||||
- 上传来源及manifest原子目录发布、只读快照;无注册参数重启仍能列出。重复上传返回旧record,不改label/artifact或旧job绑定;失效快照保持显式不可用,不退回随机初始化。不要删除仍被历史session引用的快照。
|
||||
|
||||
两面板共用PretrainedSourceSelect/LocalTrainingClient及上传控件,`accept=".pt,.onnx"`;成功后合并服务返回的已验证record并选择内容ID,连接/任务epoch变化会取消旧请求并忽略晚到响应,不能把旧endpoint结果写入新连接。上传状态显示上传并验证中(非百分比进度);取消只停止浏览器等待,服务可能已经完成验证,可刷新查看。
|
||||
|
||||
### 本轮真实CPU证据
|
||||
|
||||
两个真实文件分别在空临时根以单文件上传成功,均未要求或读取相邻`.pt`/ONNX/YAML。双方重建actor对原ONNX的48组eval最大误差均为`1.9073486e-6`。真实ONNX导入及81/97扩展、非零有限梯度、保存/恢复已更新统计通过;不是PPO收敛或安全行走证明。
|
||||
|
||||
| 来源 | 更新batch | 更新率 | 物理probe最大动作漂移 |
|
||||
| ------------------- | -------------: | ---------: | --------------------: |
|
||||
| pt,count=983138304 | 4096(UI默认) | 4.16623e-6 | 2.62260e-6 |
|
||||
| pt | 16384(上限) | 1.66647e-5 | 1.04904e-5 |
|
||||
| onnx,合成count=1e6 | 4096 | .00407929 | .00254631 |
|
||||
| onnx | 16384 | .01611989 | .01006627 |
|
||||
|
||||
以上为一次正常统计更新的实测,后续PPO可能退化,不能用更新前数值等价保证训练稳定性。
|
||||
|
||||
双面板×真实.pt/ONNX均经独立默认服务空上传根的浏览器文件选择验证。额外CPU真实runner使用4环境:两来源初始化对MuJoCo实际观测的原ONNX最大误差1.19209e-7;ONNX-derived仅运行1次PPO iteration,新增81维输入列更新为非零,导出对runner误差4.76837e-7,同trial恢复完整actor/optimizer且count保持1000096,未重设为1e6。该检查仅证明链路,不声称导航收敛。
|
||||
|
||||
## 旧CLI/注册来源使用(仍需配套文件;不适用于单文件上传)
|
||||
|
||||
先激活仓库 `.venv`,安装 `training_server/rl/requirements.txt`。CPU 身份校验依赖 `onnxruntime==1.29.0`。
|
||||
|
||||
```bash
|
||||
# SOURCE 是管理员认可的本地训练产物目录,不是浏览器传入的任意路径。
|
||||
python training_server/rl/scripts/train.py Unitree-Go2-ObstacleAvoidance \
|
||||
--pretrained-checkpoint "$SOURCE/model_10000.pt" \
|
||||
--pretrained-allowed-roots "[\"$SOURCE\"]" \
|
||||
--agent.logger tensorboard \
|
||||
--output-dir "$OUTPUT"
|
||||
```
|
||||
|
||||
目录必须有 `params/env.yaml`、`params/agent.yaml`、`policy.onnx` 和明确选定的 `.pt`。可用 `--pretrained-onnx` 指定允许根内的配对 ONNX;没有对应 `.pt` 拒绝,不把 ONNX 当成 PPO resume。**不扫描/按 mtime 自动认定 model_10000 对应 ONNX**;当前实际对应关系由 CPU ORT 数值验证建立。
|
||||
|
||||
- `--pretrained-checkpoint` 与 `--resume-checkpoint` / `--agent.resume` 互斥。
|
||||
- warm-start:新 trial 的 iteration=0;只迁移 actor,critic 和 optimizer 保持新初始化。预算仍是新训练预算。
|
||||
- resume:既有 runner.load 和 explicit-resume 的 `current+1`、`max_iterations-current` 行为完全不变;同 trial successive-halving 不应再次传 pretrained 参数。
|
||||
- 原 Flat legacy47 支持迁移;81/97 支持扩展。Rough 普通训练/原 resume 不受影响,但本轮**不支持从 legacy47 warm-start 到含 height scan 的 Rough**,会明确拒绝。
|
||||
- 校验发生在模拟器分配前;创建环境后另核对编译后的 joint 顺序、PD、默认位姿、观测顺序与动作 scale。成功初始化时在输出目录写 `initialization.json`。
|
||||
|
||||
## 严格迁移契约
|
||||
|
||||
源 actor 必须是 RSL `MLPModel`:47→512→256→128→12、ELU、经验归一化、Gaussian scalar std。所有 state key、dtype、shape、有限值及 std/var 一致性均验证;没有 `strict=False`。观测语义顺序为:
|
||||
|
||||
`base_ang_vel(3), projected_gravity(3), command(3), phase(2), joint_pos(12), joint_vel(12), actions(12)`。
|
||||
|
||||
校验 term 函数名、参数、噪声/缩放/裁剪/延迟/history、动作配置、机器人 spec 名、默认关节状态、actuator/armature/collision 配置、decimation=4 与 dt=.005。观测 corruption 开关和命令采样分布可以因导航任务改变,但基础观测的坐标/单位/顺序不变。phase 为 .6s sin/cos,命令 norm<.1 时清零;本地与源 phase 实现只差尾空行,robot XML 逐字节相同(本次只读比对)。并非宣称训练分布和导航地形完全相同。
|
||||
|
||||
迁移覆盖:
|
||||
|
||||
| 参数 | 行为 |
|
||||
| ---------------------------------------------------- | ----------------------------------------------- |
|
||||
| `mlp.0.weight[:, :47]` | 原值复制 |
|
||||
| `mlp.0.weight[:, 47:]` | 零初始化,仍 requires_grad,可由 PPO 更新 |
|
||||
| 所有 MLP bias、其余 weight、`distribution.std_param` | 原值复制 |
|
||||
| normalizer mean/var/std 前47维及标量 count | 逐值复制 |
|
||||
| normalizer 新维 | mean=0、var=1、std=1 |
|
||||
| 源 critic/optimizer/iter | 不加载,保持全新目标 critic/optimizer/iteration |
|
||||
|
||||
归一化策略经批准为 **`preserve-source-count/unit-new-features`**。真实源 count=983138304,RSL 更新率为 batch_size / 更新后 count。清 count 会让首批覆盖旧47统计,故禁止清零。新 ray∈[0,1]、heading∈[-1,1]、distance∈[0,1] 原本有界;使用近 identity 初始化(实际除以 std+.01),继承原算法缓慢更新,**不是完全冻结**。不引入 split normalizer、不改变 checkpoint 格式。
|
||||
|
||||
初始化等价在 eval / 未更新统计时验证。统计更新后不是要求动作数学上不变,而是必须与原47 actor执行相同统计更新后的动作一致,且变化无突变。48组含分布外随机 gravity 的探针实测最大变化9.06e-5;真实4env首batch变化0。
|
||||
|
||||
## 可选旧管理员注册来源并在两种面板使用
|
||||
|
||||
注册JSON由管理员在本机配置,不是HTTP请求。以下直接使用本次用户提供的源目录(只读,不改原文件):
|
||||
|
||||
```bash
|
||||
source .venv/bin/activate
|
||||
export SOURCE=/home/cen/Embodied_Workspace/unitree_rl_mjlab/logs/rsl_rl/go2_velocity/2026-08-25_Go2
|
||||
export SOURCE_CONFIG=/tmp/go2-pretrained-sources.json
|
||||
python - <<'PYCONFIG'
|
||||
import json, os
|
||||
from pathlib import Path
|
||||
root = Path(os.environ["SOURCE"])
|
||||
Path(os.environ["SOURCE_CONFIG"]).write_text(json.dumps({
|
||||
"allowedRoots": [str(root)],
|
||||
"sources": [{"id": "go2-base", "label": "Go2已训练基础行走策略",
|
||||
"checkpoint": str(root / "model_10000.pt"),
|
||||
"onnx": str(root / "policy.onnx")}]
|
||||
}, ensure_ascii=False, indent=2))
|
||||
PYCONFIG
|
||||
python training_server/server.py --trainer-python "$VIRTUAL_ENV/bin/python" \
|
||||
--pretrained-sources "$SOURCE_CONFIG" \
|
||||
--tuning-data-root "$HOME/.local/share/mujoco-go2-training"
|
||||
```
|
||||
|
||||
1. 复制服务显示的访问令牌,在普通训练面板连接服务,或在自调参工作台连接同一服务。
|
||||
2. 在“基础策略”选择已验证条目,检查`model_10000.pt`及checkpoint/ONNX完整SHA、兼容任务和观测维度。两面板默认不选,保留原随机训练行为。
|
||||
3. 选择Flat47或Obstacle81/97;Rough显示不兼容并由服务拒绝。不兼容/失效来源须显式处理,不可静默清空后退回随机初始化。
|
||||
4. 普通训练点击“发起本地训练”;自调参选择Approval/Automatic后启动。基础策略不是LLM可调字段;每个新trial(含baseline)使用同一快照与相同session seed的新critic/optimizer初始化协议。后续rung只加载该trial自己的checkpoint,不覆盖成原基础权重。
|
||||
|
||||
允许根只能来自管理员JSON,不能来自浏览器;最多32注册条目。服务对路径各级使用dir_fd/O_NOFOLLOW,拒绝symlink、穿越、FIFO、目录与超限文件(pt256MiB/ONNX64MiB/env2MiB/agent128KiB)。CPU子进程只使用受限YAML与`torch.load(weights_only=True)`验证形状、语义和ONNX数值身份。未知object/apply、重复键、不匹配配置或缺.pt均拒绝;不逆向ONNX恢复PPO,也不按mtime猜checkpoint。
|
||||
|
||||
验证后原始四件产物复制到`<tuning-data-root>/pretrained_sources/<source_id>/`内容寻址只读快照。浏览器提交的`pretrainedSourceId`是目录返回的内容SHA身份,不是管理员可读别名或文件路径;同一别名重新注册了不同文件时,旧UI的内容ID会拒绝,必须刷新。注册之后原文件变化不改变已有job/session;每次启动trial都重新核验快照SHA,CLI还核验期望source_id。勿删除有历史session依赖的快照。管理员仍须保护数据根;这是常规文件/反序列化防护,不是任意ONNX计算或同账户恶意写入的完整沙箱。
|
||||
|
||||
session的SQLite config持久化完整无路径来源描述,旧session无该字段仍走原默认。服务重启会将queued/running/evaluating/paused/awaiting_approval标记interrupted并记录原state,不自动训练或调用Agent;必须用户显式恢复。恢复只清理不完整trial,保留完整评估/来源/约束版本/mode。遗留pending proposal以`service_restart/recovery_invalidated`审计拒绝(非用户拒绝),重新提案后Approval仍需审批;已批准历史不改。来源快照丢失/损坏时恢复排队前拒绝。连续resume请求不会启动两个worker。
|
||||
|
||||
源manifest/迁移覆盖写入`initialization.json`;source-bound rung要求父checkpoint的来源记录匹配。最佳ONNX的`pretrained_initialization`metadata仅含SHA/source ID、源迭代、normalizer协议与initialIteration=0,不含本地路径或任意checkpoint infos。
|
||||
|
||||
## 持续点击导航
|
||||
|
||||
仅Obstacle81/97启用配套地图点击导航,不给普通Flat47增加非同构导航地图模式。浏览器没有20秒交互截止,但跌倒/越界仍安全停止;需用户重置并重新启用。设定新目标不重新加载策略、不改变单inflight/held-action机制,也不会自动启用已停止策略。面板明确显示策略未启用、仿真暂停或安全停止。训练与固定三seed/1000步(20秒)评估不变。
|
||||
|
||||
真实源warm-start81在WASM/ORT中21.02秒仍启用:从x=-5走到2.566;换目标使command从vx=.598/yaw=-.077变为vx=0/yaw=1。25.04秒仍启用,距新目标从2.828m降到2.243m;无teleport、脚本轨迹或constant actor。此为空旷plane单例,不证明复杂障碍泛化或导航收敛。
|
||||
|
||||
## 验证
|
||||
|
||||
```bash
|
||||
python -m unittest discover -s training_server/tests -p test_pretrained.py -v
|
||||
# 可选真实源;不硬编码用户个人目录到源码:
|
||||
GO2_PRETRAINED_SOURCE="$SOURCE" \
|
||||
GO2_PRETRAINED_EVIDENCE_DIR="$EVIDENCE" \
|
||||
python -m unittest discover -s training_server/tests -p test_pretrained.py -v
|
||||
# 可选仅4env×1iteration GPU初始化/优化/导出/重载烟测:
|
||||
GO2_PRETRAINED_GPU_SMOKE=1 MUJOCO_GL=egl \
|
||||
GO2_PRETRAINED_SOURCE="$SOURCE" GO2_PRETRAINED_EVIDENCE_DIR="$EVIDENCE" \
|
||||
python -m unittest discover -s training_server/tests -p test_pretrained.py -v
|
||||
```
|
||||
|
||||
本次真实源:checkpoint SHA `d94eebd23be8b3cc998b493fe056a9e60ceaf4705900320015f30387412a4ffb`;ONNX SHA `80150119e93ce2f656625fc3048ece43ec1281fcdfe9109e07a0157438d0df7c`。
|
||||
|
||||
- 源 ORT vs checkpoint 48 probe 最大误差 `1.9073486328125e-6`。
|
||||
- 47→47/81/97、不同ray/goal:CPU初始化最大误差0;扩展81/97的CPU ONNX导出最大误差 `1.9073486328125e-6`。
|
||||
- 真实4env源ONNX vs 扩展actor最大误差 `3.427267074584961e-7`;首batch基础mean变化 `4.190951585769653e-9`,动作变化0。
|
||||
- 新列梯度最大 `.0273208`;单iteration PPO后新列weight最大 `.00291905`,critic/optimizer得到更新。
|
||||
- 新runner加载保存checkpoint后全部actor张量精确恢复,count=983138404;导出ONNX最大误差 `4.76837158203125e-7`。
|
||||
- 安全失败例:越界路径/逃逸symlink、unsafe YAML、未知观测语义/phase/scale/armature、坏shape/NaN/std、错误actor身份、warmstart+resume混用。
|
||||
|
||||
这些是初始化和优化步骤证据,**不等于导航收敛/到达率验证**。完整命令日志和源manifest在本次交接的 `/tmp/go2-pretrained-core-*/` 工件目录。
|
||||
|
||||
### 浏览器实产物复核(不再次训练)
|
||||
|
||||
本次保留的真实warm-start产物为 `/tmp/go2-pretrained-integration-Yqhm9K/train/policy.onnx`,source SHA及新迭代0可从同目录`initialization.json`核对。下列测试需Vite开发服务(动态导入真实PhysicsAdapter);另一个终端启动后运行:
|
||||
|
||||
```bash
|
||||
npm run dev -- --host 127.0.0.1 --port 4174
|
||||
# 在另一个终端:
|
||||
GO2_PRETRAINED_DEV_URL=http://127.0.0.1:4174 \
|
||||
GO2_PRETRAINED_NAV_POLICY=/tmp/go2-pretrained-integration-Yqhm9K/train/policy.onnx \
|
||||
npx playwright test -c web_platform/playwright.config.ts web_platform/e2e/pretrainedNavigation.spec.ts
|
||||
```
|
||||
|
||||
原始47维Flat UI兼容测试还可设置`GO2_PRETRAINED_FLAT_POLICY`和`GO2_PRETRAINED_FLAT_XML`。后者必须是带12个有效actuator、采用既有`FL_hip`等绑定名称的Go2模型;裸训练XML本来不含actuator,会明确拒绝。此次`/tmp/go2-pretrained-integration-Yqhm9K/go2-flat.xml`由本地`Entity(get_go2_robot_cfg())`导出,保留源PD/armature,仅把actuator标识改为`actuator.target.removesuffix('_joint')`以匹配既有浏览器命名;没有修改生产模型或为Flat新增点击导航。Flat此测试只证明加载/推理兼容;持续行走/换目标证据来自上述真实81维warm-start产物。
|
||||
@@ -4,6 +4,10 @@
|
||||
|
||||
仓库已在 [`rl/`](rl/) 内置 `Unitree-Go2-Flat` 所需的 PPO 训练代码、Go2 模型资产和 ONNX 导出逻辑,不再要求另外克隆 `unitree_rl_mjlab`。`mjlab`、PyTorch 等大型运行依赖仍需安装在本机训练环境中。
|
||||
|
||||
## 自定义任务与地形
|
||||
|
||||
内置新增 `Unitree-Go2-ObstacleAvoidance`(81维前视射线导航),并放行 `Unitree-Go2-Rough` 训练。健康接口提供可配置参数元数据,job 的 `deployment` 返回精确地图布局、传感器及策略契约。首版支持 `plane/discrete_obstacles/rough/pyramid_stairs/wave` 的训练专用box布局;不是任意场景导入,也不实现真实深度相机。旧 Rough 的234维actor不能在当前浏览器一键部署。完整字段、坐标系、观测动作及复现方式见 [避障部署契约](OBSTACLE_AVOIDANCE.md)。
|
||||
|
||||
## 准备训练环境
|
||||
|
||||
使用仓库已有的 `.venv` 项目虚拟环境:
|
||||
@@ -19,8 +23,7 @@ python -m pip install -r training_server/requirements.txt
|
||||
使用已安装训练依赖的 Python 解释器启动服务:
|
||||
|
||||
```bash
|
||||
npm run training-server -- \
|
||||
--trainer-python "$PWD/.venv/bin/python"
|
||||
.venv/bin/python training_server/server.py
|
||||
```
|
||||
|
||||
服务启动时会在终端显示一个随机访问令牌。将该令牌填入前端“访问令牌”字段后再连接。令牌只保存在当前浏览器标签页的 `sessionStorage` 中。自动化启动时可固定令牌:
|
||||
@@ -39,7 +42,15 @@ export MUJOCO_TUNING_AGENT_MODEL='deepseek-v4-flash' # 可省略
|
||||
npm run training-server -- --trainer-python "$PWD/.venv/bin/python"
|
||||
```
|
||||
|
||||
普通训练不要求 DeepSeek key。未配置时健康接口会把 tuning 标记为不可用;只有创建 session 时明确勾选 fallback,才允许 Agent 失败后使用 Optuna 候选,不会静默降级。
|
||||
普通训练与基础策略上传不要求 DeepSeek key。未配置时健康接口会把 tuning 标记为不可用;只有创建 session 时明确勾选 fallback,才允许使用 Optuna 候选,不会静默降级。
|
||||
|
||||
### 从浏览器选择已有策略(推荐)
|
||||
|
||||
普通训练和自调参均可在“基础策略”直接选择单个`.pt`或`.onnx`文件:连接默认服务 → 勾选Go2 legacy47观测/关节语义确认 → 点击“选择基础策略文件” → 等待验证成功后自动选择内容ID并检查文件名/格式/原文件SHA → 发起训练或自调参。不需要`--pretrained-sources`、管理员JSON、服务器路径或sidecars;旧管理员注册是可选高级功能。
|
||||
|
||||
上传期间禁止启动;取消、验证失败、任务/连接改变会保留旧选择,旧请求结果不会污染新连接。相同文件可重选重试,失效来源不会静默退回随机;从头训练须明确选择“不选择(随机初始化)”。支持47维Go2 actor扩展到81/97新输入,不接受未知机器人/任意ONNX/97维源checkpoint。
|
||||
|
||||
**只继承actor,不是完整PPO resume**:两格式critic/optimizer全新、iteration0。`.pt`保留原normalizer count和探索std;ONNX仅继承确定性网络,count合成1,000,000、std/var按受限模板推导、探索std使用目标默认1.0,源迭代未知。详细上限、安全边界及点击步骤见[基础策略上传与迁移](PRETRAINED.md)。
|
||||
|
||||
默认训练工程是仓库内的 `training_server/rl`。如需使用包含其他已注册任务的外部训练工程,仍可通过 `--trainer-root /path/to/trainer` 或 `UNITREE_RL_MJLAB_ROOT` 覆盖。默认端口是 `8765`。如果前端不是从 `localhost` 或 `127.0.0.1` 提供,可显式添加来源:
|
||||
|
||||
@@ -64,6 +75,7 @@ python training_server/server.py \
|
||||
## 接口
|
||||
|
||||
- `GET /api/training/health`:运行环境、允许的任务和活动任务;
|
||||
- `POST /api/training/pretrained-sources/upload?format=pt|onnx&template=go2-legacy47-v1&name=显示名`:认证有界二进制单文件上传,返回持久化内容ID;
|
||||
- `POST /api/training/jobs`:发起训练;
|
||||
- `GET /api/training/jobs/{id}`:状态、迭代进度和最近日志;
|
||||
- `DELETE /api/training/jobs/{id}`:停止训练;
|
||||
@@ -80,7 +92,7 @@ python training_server/server.py \
|
||||
- `GET /api/tuning/sessions/{id}/artifacts/best/policy.onnx`:下载最佳策略;
|
||||
- `GET /api/tuning/presets`:列出可供普通训练复用的最佳奖励 preset。
|
||||
|
||||
普通任务状态在服务重启后丢失,但日志、checkpoint 和 ONNX 保留在 `training_server/rl/logs/rsl_rl/`;调参状态、调度令牌、参数护栏、回滚基准及产物持久化在 `logs/auto_tuning/`。候选配置数可在 1–100 间设置(包含基线配置,仍受连续无提升早停约束)。API 只接收 32 位资源 ID,不接收客户端文件路径;奖励 patch 受到名称、符号、上下界、Session 护栏、每轮最多 4 项及 `0.5×–2×` 变化率校验。
|
||||
普通任务状态在服务重启后丢失,但日志、checkpoint 和 ONNX 保留在 `training_server/rl/logs/rsl_rl/`;调参状态、调度令牌、参数护栏、回滚基准及产物持久化在 `logs/auto_tuning/`。候选配置数可在 1–100 间设置(包含基线配置,仍受连续无提升早停约束)。作业/session API 接收32位资源ID,基础策略选择使用64位内容ID;上传仅接收文件字节,不接收客户端服务器文件路径;奖励 patch 受到名称、符号、上下界、Session 护栏、每轮最多 4 项及 `0.5×–2×` 变化率校验。
|
||||
|
||||
## 测试
|
||||
|
||||
@@ -90,3 +102,7 @@ python -m pip install -r requirements-dev.txt
|
||||
npm run lint:python
|
||||
npm run test:training-server
|
||||
```
|
||||
|
||||
## 使用已训练基础策略
|
||||
|
||||
普通训练与自调参面板均提供“基础策略”选择。管理员通过 `--pretrained-sources` 注册只读本地ONNX及数值匹配的.pt;服务创建SHA绑定快照,浏览器只能选择内容ID,不能提交文件路径。完整可直接运行的注册/启动命令、兼容47/81/97范围、重启恢复规则与验证方法见 [基础策略迁移](PRETRAINED.md)。缺.pt、Rough不匹配或快照失效均明确拒绝,绝不静默随机初始化。
|
||||
|
||||
@@ -0,0 +1,502 @@
|
||||
"""Validated Go2 velocity actor warm-start; never an optimizer/iteration resume.
|
||||
|
||||
Callers supply administrator-owned allowed roots, not roots from an HTTP request.
|
||||
Legacy YAML tags are decoded as inert data: no import, eval or object construction.
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import io
|
||||
import json
|
||||
from dataclasses import asdict, dataclass
|
||||
from enum import Enum
|
||||
from importlib.metadata import version
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
import yaml
|
||||
|
||||
BASE_TERMS = [
|
||||
"base_ang_vel",
|
||||
"projected_gravity",
|
||||
"command",
|
||||
"phase",
|
||||
"joint_pos",
|
||||
"joint_vel",
|
||||
"actions",
|
||||
]
|
||||
JOINTS = [
|
||||
f"{leg}_{joint}_joint" for leg in ("FL", "FR", "RL", "RR") for joint in ("hip", "thigh", "calf")
|
||||
]
|
||||
NORMALIZATION_POLICY = "preserve-source-count/unit-new-features"
|
||||
|
||||
|
||||
class PretrainedError(ValueError):
|
||||
"""The supplied source does not prove the supported transfer contract."""
|
||||
|
||||
|
||||
class _DataLoader(yaml.SafeLoader):
|
||||
def construct_mapping(self, node, deep=False):
|
||||
keys = [self.construct_object(key, deep=deep) for key, _ in node.value]
|
||||
if any(type(key) not in (str, int) for key in keys) or len(set(keys)) != len(keys):
|
||||
raise PretrainedError("Configuration mappings require unique string/integer keys")
|
||||
return super().construct_mapping(node, deep=deep)
|
||||
|
||||
|
||||
def _symbol(loader, suffix, node):
|
||||
if loader.construct_scalar(node) != "":
|
||||
raise PretrainedError("Python name tags must have an empty value")
|
||||
return {"symbol": suffix}
|
||||
|
||||
|
||||
def _enum(loader, node):
|
||||
values = loader.construct_sequence(node)
|
||||
if len(values) != 1 or not isinstance(values[0], (str, int)):
|
||||
raise PretrainedError("Invalid legacy enum data")
|
||||
return values[0]
|
||||
|
||||
|
||||
_DataLoader.add_constructor("tag:yaml.org,2002:python/tuple", _DataLoader.construct_yaml_seq)
|
||||
_DataLoader.add_multi_constructor("tag:yaml.org,2002:python/name:", _symbol)
|
||||
for _name in ("mjlab.actuator.actuator.TransmissionType", "mjlab.viewer.viewer_config.OriginType"):
|
||||
_DataLoader.add_constructor("tag:yaml.org,2002:python/object/apply:" + _name, _enum)
|
||||
_DataLoader.add_constructor(
|
||||
"tag:yaml.org,2002:python/object/apply:builtins.slice",
|
||||
lambda loader, node: {"slice": loader.construct_sequence(node)},
|
||||
)
|
||||
|
||||
|
||||
def _plain(value):
|
||||
if isinstance(value, Enum):
|
||||
return value.value
|
||||
if callable(value):
|
||||
return {"symbol": value.__module__ + "." + value.__qualname__}
|
||||
if isinstance(value, dict):
|
||||
return {k: _plain(v) for k, v in value.items()}
|
||||
if isinstance(value, (list, tuple)):
|
||||
return [_plain(v) for v in value]
|
||||
return value
|
||||
|
||||
|
||||
def _require(ok, message):
|
||||
if not ok:
|
||||
raise PretrainedError(message)
|
||||
|
||||
|
||||
def _read_allowed(path, roots, limit):
|
||||
path = Path(path).expanduser().resolve(strict=True)
|
||||
_require(
|
||||
any(path.is_relative_to(root) for root in roots),
|
||||
"Source is outside configured allowed roots",
|
||||
)
|
||||
_require(path.is_file() and path.stat().st_size <= limit, "Source file missing or oversized")
|
||||
# Hash and deserialize the same bytes, not a second pathname lookup.
|
||||
with path.open("rb") as stream:
|
||||
data = stream.read(limit + 1)
|
||||
_require(len(data) <= limit, "Source file oversized")
|
||||
return path, data
|
||||
|
||||
|
||||
def _yaml_data(data):
|
||||
try:
|
||||
result = yaml.load(data, Loader=_DataLoader)
|
||||
_require(isinstance(result, dict), "Expected YAML mapping")
|
||||
return result
|
||||
except yaml.YAMLError as exc:
|
||||
raise PretrainedError("Unsupported or invalid configuration YAML") from exc
|
||||
|
||||
|
||||
def _actor_config(config):
|
||||
actor = dict(config["actor"])
|
||||
distribution = dict(actor["distribution_cfg"])
|
||||
# Old RSL config omitted this default; the tensor contract is checked too.
|
||||
distribution.setdefault("class_name", "GaussianDistribution")
|
||||
actor["distribution_cfg"] = distribution
|
||||
actor.setdefault("class_name", "MLPModel")
|
||||
actor.setdefault("cnn_cfg", None)
|
||||
return actor
|
||||
|
||||
|
||||
def validate_semantics(source_env, source_agent, target_env, target_agent):
|
||||
"""Compare source and target against the repository's supported legacy47 contract."""
|
||||
from src.tasks.velocity.config.go2.env_cfgs import unitree_go2_flat_env_cfg
|
||||
from src.tasks.velocity.config.go2.rl_cfg import unitree_go2_ppo_runner_cfg
|
||||
|
||||
reference = _plain(asdict(unitree_go2_flat_env_cfg()))
|
||||
agent = _plain(asdict(unitree_go2_ppo_runner_cfg()))
|
||||
target_env, target_agent = _plain(target_env), _plain(target_agent)
|
||||
try:
|
||||
for label, env in (("source", source_env), ("target", target_env)):
|
||||
group = env["observations"]["actor"]
|
||||
names = list(group["terms"])
|
||||
_require(names[:7] == BASE_TERMS, f"{label}: unknown base observation order")
|
||||
_require(
|
||||
names == BASE_TERMS
|
||||
if label == "source"
|
||||
else names in (BASE_TERMS, BASE_TERMS + ["forward_depth", "target_error"]),
|
||||
f"{label}: unsupported observation suffix",
|
||||
)
|
||||
for name in BASE_TERMS:
|
||||
_require(
|
||||
group["terms"][name] == reference["observations"]["actor"]["terms"][name],
|
||||
f"{label}: incompatible observation {name}",
|
||||
)
|
||||
for key in (
|
||||
"concatenate_terms",
|
||||
"concatenate_dim",
|
||||
"history_length",
|
||||
"flatten_history_dim",
|
||||
):
|
||||
_require(
|
||||
group[key] == reference["observations"]["actor"][key],
|
||||
f"{label}: incompatible observation {key}",
|
||||
)
|
||||
_require(
|
||||
env["actions"] == reference["actions"], f"{label}: incompatible action transform"
|
||||
)
|
||||
robot, ref_robot = (
|
||||
env["scene"]["entities"]["robot"],
|
||||
reference["scene"]["entities"]["robot"],
|
||||
)
|
||||
for key in ("spec_fn", "articulation", "sort_actuators", "collisions"):
|
||||
_require(robot[key] == ref_robot[key], f"{label}: incompatible robot {key}")
|
||||
for key in ("joint_pos", "joint_vel"):
|
||||
_require(
|
||||
robot["init_state"][key] == ref_robot["init_state"][key],
|
||||
f"{label}: incompatible default {key}",
|
||||
)
|
||||
_require(
|
||||
env["decimation"] == 4 and env["sim"]["mujoco"]["timestep"] == 0.005,
|
||||
f"{label}: policy must run at 50 Hz",
|
||||
)
|
||||
if len(target_env["observations"]["actor"]["terms"]) > 7:
|
||||
from src.tasks.obstacle_avoidance.env_cfg import unitree_go2_obstacle_env_cfg
|
||||
|
||||
suffix = _plain(asdict(unitree_go2_obstacle_env_cfg()))["observations"]["actor"][
|
||||
"terms"
|
||||
]
|
||||
terms = target_env["observations"]["actor"]["terms"]
|
||||
depth = dict(terms["forward_depth"])
|
||||
_require(
|
||||
set(depth["params"]) == {"max_distance"}
|
||||
and 0 < depth["params"]["max_distance"] < float("inf"),
|
||||
"Invalid ray normalization",
|
||||
)
|
||||
depth["params"] = suffix["forward_depth"]["params"]
|
||||
_require(
|
||||
depth == suffix["forward_depth"]
|
||||
and terms["target_error"] == suffix["target_error"],
|
||||
"Unknown ray/goal observation semantics",
|
||||
)
|
||||
for label, cfg in (("source", source_agent), ("target", target_agent)):
|
||||
_require(
|
||||
_actor_config(cfg) == _actor_config(agent),
|
||||
f"{label}: incompatible actor architecture/activation/std",
|
||||
)
|
||||
_require(
|
||||
cfg["obs_groups"]["actor"] == ["actor"] and cfg["clip_actions"] is None,
|
||||
f"{label}: incompatible actor groups/action clipping",
|
||||
)
|
||||
except (KeyError, TypeError) as exc:
|
||||
raise PretrainedError(f"Missing or invalid semantic configuration: {exc}") from exc
|
||||
|
||||
|
||||
def make_reference_actor(dim=47):
|
||||
from rsl_rl.models import MLPModel
|
||||
from tensordict import TensorDict
|
||||
|
||||
return MLPModel(
|
||||
TensorDict({"actor": torch.zeros(1, dim)}, batch_size=[1]),
|
||||
{"actor": ["actor"]},
|
||||
"actor",
|
||||
12,
|
||||
hidden_dims=(512, 256, 128),
|
||||
activation="elu",
|
||||
obs_normalization=True,
|
||||
distribution_cfg={
|
||||
"class_name": "GaussianDistribution",
|
||||
"init_std": 1.0,
|
||||
"std_type": "scalar",
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def validate_actor_state(state, dim=47):
|
||||
expected = make_reference_actor(dim).state_dict()
|
||||
_require(
|
||||
isinstance(state, dict) and state.keys() == expected.keys(), "Unexpected actor state keys"
|
||||
)
|
||||
for key, tensor in state.items():
|
||||
_require(
|
||||
isinstance(tensor, torch.Tensor)
|
||||
and tensor.shape == expected[key].shape
|
||||
and tensor.dtype == expected[key].dtype,
|
||||
f"Incompatible actor tensor: {key}",
|
||||
)
|
||||
_require(bool(torch.isfinite(tensor).all()), f"Nonfinite actor tensor: {key}")
|
||||
_require(state["obs_normalizer.count"].item() >= 0, "Negative normalizer count")
|
||||
_require(
|
||||
bool((state["obs_normalizer._var"] >= 0).all())
|
||||
and bool((state["obs_normalizer._std"] >= 0).all()),
|
||||
"Negative normalization variance/std",
|
||||
)
|
||||
_require(
|
||||
torch.allclose(
|
||||
state["obs_normalizer._std"].square(),
|
||||
state["obs_normalizer._var"],
|
||||
atol=1e-5,
|
||||
rtol=1e-5,
|
||||
),
|
||||
"Inconsistent normalization variance/std",
|
||||
)
|
||||
_require(bool((state["distribution.std_param"] > 0).all()), "Nonpositive exploration std")
|
||||
|
||||
|
||||
def comparison_observations():
|
||||
"""Deterministic random + physically plausible standing/walking probe vectors."""
|
||||
g = torch.Generator().manual_seed(20260825)
|
||||
random = torch.randn(32, 47, generator=g)
|
||||
physical = torch.zeros(16, 47)
|
||||
physical[:, 5] = -1 # body-frame projected gravity
|
||||
physical[:, 6] = torch.linspace(0, 1, 16)
|
||||
phase = torch.linspace(0, 2 * torch.pi, 16)
|
||||
physical[:, 9], physical[:, 10] = phase.sin(), phase.cos()
|
||||
physical[0, 9:11] = 0 # standing phase is masked
|
||||
return torch.cat((random, physical))
|
||||
|
||||
|
||||
def verify_onnx(onnx_bytes, actor):
|
||||
import numpy as np
|
||||
import onnx
|
||||
import onnxruntime as ort
|
||||
|
||||
graph = onnx.load_model_from_string(onnx_bytes)
|
||||
_require(
|
||||
all(t.data_location != onnx.TensorProto.EXTERNAL for t in graph.graph.initializer),
|
||||
"External ONNX tensors are not allowed",
|
||||
)
|
||||
options = ort.SessionOptions()
|
||||
options.intra_op_num_threads = 1
|
||||
options.inter_op_num_threads = 1
|
||||
session = ort.InferenceSession(
|
||||
onnx_bytes, sess_options=options, providers=["CPUExecutionProvider"]
|
||||
)
|
||||
inputs, outputs = session.get_inputs(), session.get_outputs()
|
||||
_require(
|
||||
len(inputs) == len(outputs) == 1
|
||||
and inputs[0].shape == [1, 47]
|
||||
and outputs[0].shape == [1, 12]
|
||||
and inputs[0].type == outputs[0].type == "tensor(float)",
|
||||
"ONNX must be float32 [1,47] -> [1,12]",
|
||||
)
|
||||
metadata = session.get_modelmeta().custom_metadata_map
|
||||
_require(
|
||||
metadata.get("joint_names", "").split(",") == JOINTS
|
||||
and metadata.get("observation_names", "").split(",") == BASE_TERMS
|
||||
and metadata.get("command_names") == "twist",
|
||||
"ONNX joint/observation/command semantics mismatch",
|
||||
)
|
||||
for key, expected in {
|
||||
"action_scale": [0.25],
|
||||
"joint_stiffness": [20, 20, 40] * 4,
|
||||
"joint_damping": [1, 1, 2] * 4,
|
||||
"default_joint_pos": [-0.1, 0.9, -1.8, 0.1, 0.9, -1.8] * 2,
|
||||
}.items():
|
||||
try:
|
||||
actual = [float(v) for v in metadata[key].split(",")]
|
||||
_require(
|
||||
len(actual) == len(expected) and np.allclose(actual, expected, rtol=0, atol=1e-6),
|
||||
f"ONNX {key} mismatch",
|
||||
)
|
||||
except (KeyError, ValueError) as exc:
|
||||
raise PretrainedError(f"Invalid ONNX {key}") from exc
|
||||
actor.eval()
|
||||
obs = comparison_observations()
|
||||
with torch.inference_mode():
|
||||
expected = actor.mlp(actor.obs_normalizer(obs)).numpy()
|
||||
actual = np.concatenate(
|
||||
[session.run(None, {inputs[0].name: row[None].numpy()})[0] for row in obs]
|
||||
)
|
||||
_require(
|
||||
np.isfinite(actual).all() and np.allclose(actual, expected, atol=2e-5, rtol=2e-5),
|
||||
"ONNX does not match the selected checkpoint actor",
|
||||
)
|
||||
return {
|
||||
"provider": "CPUExecutionProvider",
|
||||
"probe_count": len(obs),
|
||||
"max_abs_error": float(np.max(np.abs(actual - expected))),
|
||||
"atol": 2e-5,
|
||||
"rtol": 2e-5,
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class ValidatedSource:
|
||||
actor_state: dict
|
||||
manifest: dict
|
||||
|
||||
|
||||
def read_pretrained_source(checkpoint, *, allowed_roots, target_env, target_agent, onnx_path=None):
|
||||
"""Read an explicit .pt plus paired ONNX/config. Never guess checkpoint by mtime.
|
||||
|
||||
Future services should resolve opaque source IDs to these server-owned paths.
|
||||
An ONNX selection without an explicit corresponding .pt is not a training source.
|
||||
"""
|
||||
roots = [Path(p).expanduser().resolve(strict=True) for p in allowed_roots]
|
||||
_require(
|
||||
bool(roots) and all(p.is_dir() for p in roots),
|
||||
"Configure at least one local allowed source root",
|
||||
)
|
||||
checkpoint, checkpoint_bytes = _read_allowed(checkpoint, roots, 256 * 1024 * 1024)
|
||||
_require(
|
||||
checkpoint.suffix == ".pt",
|
||||
"Warm-start requires a corresponding .pt checkpoint; ONNX cannot resume PPO",
|
||||
)
|
||||
files = {"checkpoint": (checkpoint, checkpoint_bytes)}
|
||||
for name, path, limit in (
|
||||
("onnx", onnx_path or checkpoint.parent / "policy.onnx", 64 * 1024 * 1024),
|
||||
("env", checkpoint.parent / "params/env.yaml", 2 * 1024 * 1024),
|
||||
("agent", checkpoint.parent / "params/agent.yaml", 128 * 1024),
|
||||
):
|
||||
files[name] = _read_allowed(path, roots, limit)
|
||||
source_env, source_agent = _yaml_data(files["env"][1]), _yaml_data(files["agent"][1])
|
||||
validate_semantics(source_env, source_agent, target_env, target_agent)
|
||||
try:
|
||||
state = torch.load(io.BytesIO(checkpoint_bytes), map_location="cpu", weights_only=True)
|
||||
_require(
|
||||
isinstance(state, dict) and isinstance(state.get("iter"), int),
|
||||
"Invalid checkpoint iteration",
|
||||
)
|
||||
actor_state = state["actor_state_dict"]
|
||||
validate_actor_state(actor_state)
|
||||
except (KeyError, RuntimeError) as exc:
|
||||
raise PretrainedError("Unsupported checkpoint") from exc
|
||||
actor = make_reference_actor()
|
||||
actor.load_state_dict(actor_state, strict=True)
|
||||
identity = verify_onnx(files["onnx"][1], actor)
|
||||
artifacts = {
|
||||
name: {"name": path.name, "sha256": hashlib.sha256(data).hexdigest(), "bytes": len(data)}
|
||||
for name, (path, data) in files.items()
|
||||
}
|
||||
manifest = {
|
||||
"schema_version": 1,
|
||||
"mode": "pretrained-warm-start",
|
||||
"contract": "go2-legacy47-v1",
|
||||
"artifacts": artifacts,
|
||||
"source_iteration": state["iter"],
|
||||
"source_actor_dim": 47,
|
||||
"source_normalizer_count": actor_state["obs_normalizer.count"].item(),
|
||||
"verification_versions": {
|
||||
name: version(name)
|
||||
for name in ("torch", "rsl-rl-lib", "mjlab", "onnxruntime", "PyYAML")
|
||||
},
|
||||
"normalization": NORMALIZATION_POLICY,
|
||||
"onnx_identity": identity,
|
||||
"critic": "fresh-target-initialization",
|
||||
"optimizer": "fresh",
|
||||
"iteration": 0,
|
||||
"base_observation_terms": BASE_TERMS,
|
||||
"joint_names": JOINTS,
|
||||
}
|
||||
manifest["source_id"] = hashlib.sha256(
|
||||
json.dumps(artifacts, sort_keys=True).encode()
|
||||
).hexdigest()
|
||||
return ValidatedSource(actor_state, manifest)
|
||||
|
||||
|
||||
def warm_start_actor(actor, source):
|
||||
"""Strictly transfer a validated legacy actor into a fresh 47/81/97 actor."""
|
||||
dim = actor.obs_dim
|
||||
_require(dim in (47, 81, 97), "Only legacy47 and ray/goal 81/97 actors are supported")
|
||||
reference = make_reference_actor(dim)
|
||||
_require(
|
||||
type(actor) is type(reference)
|
||||
and repr(actor.mlp) == repr(reference.mlp)
|
||||
and type(actor.distribution) is type(reference.distribution)
|
||||
and actor.distribution.std_type == "scalar"
|
||||
and actor.obs_normalization
|
||||
and type(actor.obs_normalizer) is type(reference.obs_normalizer)
|
||||
and actor.obs_normalizer.eps == reference.obs_normalizer.eps
|
||||
and actor.obs_normalizer.until is None
|
||||
and list(actor.obs_groups) == ["actor"],
|
||||
"Target actor runtime architecture mismatch",
|
||||
)
|
||||
validate_actor_state(actor.state_dict(), dim)
|
||||
validate_actor_state(source.actor_state)
|
||||
migrated = {}
|
||||
for key, tensor in source.actor_state.items():
|
||||
if key == "mlp.0.weight":
|
||||
value = tensor.new_zeros((512, dim))
|
||||
value[:, :47] = tensor
|
||||
elif key in ("obs_normalizer._mean", "obs_normalizer._var", "obs_normalizer._std"):
|
||||
value = tensor.new_full((1, dim), 0 if key.endswith("_mean") else 1)
|
||||
value[:, :47] = tensor
|
||||
else:
|
||||
value = tensor.clone()
|
||||
migrated[key] = value
|
||||
actor.load_state_dict(migrated, strict=True)
|
||||
return {
|
||||
**source.manifest,
|
||||
"target_actor_dim": dim,
|
||||
"copied_tensors": list(source.actor_state),
|
||||
"zero_initialized_input_columns": [47, dim],
|
||||
"new_feature_statistics": {"mean": 0, "var": 1, "std": 1},
|
||||
}
|
||||
|
||||
|
||||
def validate_runtime_contract(env):
|
||||
"""Confirm compiled joint order/PD/defaults, not only configuration names."""
|
||||
from mjlab.rl.exporter_utils import get_base_metadata
|
||||
|
||||
metadata = get_base_metadata(env, "pretrained-validation")
|
||||
_require(metadata["joint_names"] == JOINTS, "Compiled robot joint order mismatch")
|
||||
_require(
|
||||
list(metadata["observation_names"])
|
||||
in (BASE_TERMS, BASE_TERMS + ["forward_depth", "target_error"]),
|
||||
"Compiled observation order mismatch",
|
||||
)
|
||||
_require(metadata["command_names"] == ["twist"], "Compiled command mismatch")
|
||||
for key, expected in {
|
||||
"action_scale": [0.25],
|
||||
"joint_stiffness": [20, 20, 40] * 4,
|
||||
"joint_damping": [1, 1, 2] * 4,
|
||||
"default_joint_pos": [-0.1, 0.9, -1.8, 0.1, 0.9, -1.8] * 2,
|
||||
}.items():
|
||||
actual = torch.as_tensor(metadata[key], dtype=torch.float64).reshape(-1)
|
||||
target = torch.tensor(expected, dtype=torch.float64)
|
||||
_require(
|
||||
actual.shape == target.shape and torch.allclose(actual, target, atol=1e-6, rtol=0),
|
||||
f"Compiled {key} mismatch",
|
||||
)
|
||||
return metadata
|
||||
|
||||
|
||||
def public_initialization_metadata(initialization):
|
||||
"""Export provenance without local paths, including original single-file identity."""
|
||||
artifacts = initialization["artifacts"]
|
||||
result = {
|
||||
"sourceId": initialization["source_id"],
|
||||
"checkpointSha256": artifacts["checkpoint"]["sha256"],
|
||||
"normalization": initialization["normalization"],
|
||||
"sourceIteration": initialization["source_iteration"],
|
||||
"initialIteration": 0,
|
||||
}
|
||||
if "sourceFormat" in initialization:
|
||||
result.update(
|
||||
sourceFormat=initialization["sourceFormat"],
|
||||
uploadSha256=artifacts["upload"]["sha256"],
|
||||
templateId=initialization["contract"],
|
||||
derivedFields=initialization["derived_fields"],
|
||||
templateConfirmation=initialization["template_confirmation"],
|
||||
)
|
||||
else:
|
||||
result["onnxSha256"] = artifacts["onnx"]["sha256"]
|
||||
return result
|
||||
|
||||
|
||||
def initialize_runner(runner, source):
|
||||
_require(
|
||||
runner.current_learning_iteration == 0 and not runner.alg.optimizer.state,
|
||||
"Warm-start requires a fresh runner, never a resumed trial",
|
||||
)
|
||||
# Do not load critic/optimizer/iteration from the source, even if shapes match.
|
||||
return warm_start_actor(runner.alg.actor, source)
|
||||
@@ -0,0 +1,447 @@
|
||||
"""Administrator-registered, content-bound local training sources (no client paths)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
import stat
|
||||
import subprocess
|
||||
import tempfile
|
||||
import threading
|
||||
import time
|
||||
from contextlib import contextmanager
|
||||
from copy import deepcopy
|
||||
from pathlib import Path
|
||||
|
||||
SOURCE_ID = re.compile(r"^[A-Za-z0-9_-]{1,64}$")
|
||||
DIGEST = re.compile(r"^[0-9a-f]{64}$")
|
||||
LIMITS = {
|
||||
"checkpoint": 256 * 1024**2,
|
||||
"onnx": 64 * 1024**2,
|
||||
"env": 2 * 1024**2,
|
||||
"agent": 128 * 1024,
|
||||
}
|
||||
TASKS = ["Unitree-Go2-Flat", "Unitree-Go2-ObstacleAvoidance"]
|
||||
|
||||
|
||||
class SourceError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
def regular_bytes(path: Path, root: Path, limit: int) -> bytes:
|
||||
"""Walk with dir_fd/O_NOFOLLOW, so symlink substitution cannot escape the root."""
|
||||
if ".." in path.parts or not path.is_absolute():
|
||||
raise SourceError("基础策略路径必须是允许根内的绝对路径,不能含 ..")
|
||||
try:
|
||||
parts = path.relative_to(root).parts
|
||||
except ValueError as error:
|
||||
raise SourceError("基础策略文件不在管理员配置的允许根内") from error
|
||||
if not parts:
|
||||
raise SourceError("基础策略必须是常规文件")
|
||||
try:
|
||||
fd = os.open(root, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW)
|
||||
except OSError as error:
|
||||
raise SourceError("基础策略快照目录缺失或无效,拒绝训练") from error
|
||||
try:
|
||||
for part in parts[:-1]:
|
||||
next_fd = os.open(part, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW, dir_fd=fd)
|
||||
os.close(fd)
|
||||
fd = next_fd
|
||||
leaf = os.open(parts[-1], os.O_RDONLY | os.O_NOFOLLOW | os.O_NONBLOCK, dir_fd=fd)
|
||||
with os.fdopen(leaf, "rb") as stream:
|
||||
info = os.fstat(stream.fileno())
|
||||
if not stat.S_ISREG(info.st_mode) or info.st_size > limit:
|
||||
raise SourceError("基础策略文件必须是常规文件且不能超过大小上限")
|
||||
data = stream.read(limit + 1)
|
||||
if len(data) > limit:
|
||||
raise SourceError("基础策略文件超过大小上限")
|
||||
return data
|
||||
except OSError as error:
|
||||
raise SourceError(
|
||||
"基础策略文件缺失或含symlink;请提供.pt、policy.onnx和params配置常规文件"
|
||||
) from error
|
||||
finally:
|
||||
os.close(fd)
|
||||
|
||||
|
||||
class PretrainedSources:
|
||||
def __init__(self, config_path: Path | None, store: Path, python: str, trainer_root: Path):
|
||||
self.store = store.expanduser().resolve()
|
||||
self.store.mkdir(parents=True, exist_ok=True, mode=0o700)
|
||||
self.python = python
|
||||
self.trainer_root = trainer_root.resolve()
|
||||
self.entries: dict[str, dict] = {}
|
||||
self._upload_lock = threading.Lock()
|
||||
self._catalog_lock = threading.RLock()
|
||||
self._restore_uploads()
|
||||
if config_path is None:
|
||||
return
|
||||
try:
|
||||
if config_path.stat().st_size > 128 * 1024:
|
||||
raise SourceError("基础策略注册配置不能超过128KiB")
|
||||
config = json.loads(config_path.read_text())
|
||||
if not isinstance(config, dict) or set(config) != {"allowedRoots", "sources"}:
|
||||
raise SourceError("注册配置只接受allowedRoots和sources")
|
||||
if (
|
||||
not isinstance(config["allowedRoots"], list)
|
||||
or not isinstance(config["sources"], list)
|
||||
or len(config["sources"]) > 32
|
||||
):
|
||||
raise SourceError("允许根/sources必须为数组,最多32个注册条目")
|
||||
roots = [Path(p).expanduser().resolve(strict=True) for p in config["allowedRoots"]]
|
||||
if not roots or any(not p.is_dir() for p in roots):
|
||||
raise SourceError("请配置存在的本地允许根目录")
|
||||
for entry in config["sources"]:
|
||||
if not isinstance(entry, dict) or set(entry) != {
|
||||
"id",
|
||||
"label",
|
||||
"checkpoint",
|
||||
"onnx",
|
||||
}:
|
||||
raise SourceError("注册条目必须包含id/label/checkpoint/onnx")
|
||||
key = entry["id"]
|
||||
if not isinstance(key, str) or not SOURCE_ID.fullmatch(key) or key in self.entries:
|
||||
raise SourceError("注册id无效或重复")
|
||||
if not isinstance(entry["label"], str) or not 1 <= len(entry["label"]) <= 100:
|
||||
raise SourceError("基础策略label必须为1–100字符")
|
||||
record = {"id": key, "label": entry["label"], "ready": False, "compatibleTasks": []}
|
||||
self.entries[key] = {"public": record}
|
||||
try:
|
||||
checkpoint, onnx = (
|
||||
Path(entry["checkpoint"]).expanduser(),
|
||||
Path(entry["onnx"]).expanduser(),
|
||||
)
|
||||
if not re.fullmatch(r"[A-Za-z0-9_.-]+\.pt", checkpoint.name):
|
||||
raise SourceError("请选择ONNX对应的.pt训练checkpoint,ONNX不能直接续训")
|
||||
files = {
|
||||
"checkpoint": checkpoint,
|
||||
"onnx": onnx,
|
||||
"env": checkpoint.parent / "params/env.yaml",
|
||||
"agent": checkpoint.parent / "params/agent.yaml",
|
||||
}
|
||||
data = {}
|
||||
for name, path in files.items():
|
||||
root = next((r for r in roots if path.is_relative_to(r)), None)
|
||||
if root is None:
|
||||
raise SourceError("注册文件不在允许根内")
|
||||
data[name] = regular_bytes(path, root, LIMITS[name])
|
||||
names = {
|
||||
"checkpoint": checkpoint.name,
|
||||
"onnx": "policy.onnx",
|
||||
"env": "params/env.yaml",
|
||||
"agent": "params/agent.yaml",
|
||||
}
|
||||
with tempfile.TemporaryDirectory(dir=self.store, prefix="import-") as temporary:
|
||||
directory = Path(temporary)
|
||||
for name, content in data.items():
|
||||
destination = directory / names[name]
|
||||
destination.parent.mkdir(exist_ok=True)
|
||||
destination.write_bytes(content)
|
||||
manifest = self._validate(directory, checkpoint.name, TASKS[0], None)
|
||||
digest = manifest["source_id"]
|
||||
bound = {
|
||||
"sourceId": digest,
|
||||
"registeredId": key,
|
||||
"label": entry["label"],
|
||||
"manifest": manifest,
|
||||
}
|
||||
destination = self.store / digest
|
||||
if not destination.exists():
|
||||
shutil.copytree(directory, destination)
|
||||
for file in destination.rglob("*"):
|
||||
file.chmod(0o555 if file.is_dir() else 0o444)
|
||||
destination.chmod(0o555)
|
||||
self.verify(bound)
|
||||
record.update(
|
||||
id=digest,
|
||||
ready=True,
|
||||
compatibleTasks=TASKS,
|
||||
observationSizes=[47, 81, 97],
|
||||
initialization=bound,
|
||||
)
|
||||
except (SourceError, OSError, ValueError, TypeError) as error:
|
||||
record["error"] = f"{error};请修正管理员注册配置并重启服务,不会退回随机初始化"
|
||||
except (OSError, TypeError, ValueError) as error:
|
||||
raise SourceError(f"基础策略注册配置无效:{error}") from error
|
||||
|
||||
def _restore_uploads(self):
|
||||
# Only committed content directories are catalogued. Partial uploads are never sources.
|
||||
for directory in self.store.iterdir():
|
||||
if not DIGEST.fullmatch(directory.name) or not (directory / "upload.json").is_file():
|
||||
continue
|
||||
try:
|
||||
manifest = json.loads(
|
||||
regular_bytes(directory / "upload.json", self.store, 64 * 1024)
|
||||
)
|
||||
label = json.loads(regular_bytes(directory / "label.json", self.store, 1024))[
|
||||
"label"
|
||||
]
|
||||
bound = {
|
||||
"sourceId": directory.name,
|
||||
"registeredId": directory.name,
|
||||
"label": self._display_name(label),
|
||||
"manifest": manifest,
|
||||
}
|
||||
if manifest["source_id"] != directory.name:
|
||||
raise SourceError("上传来源身份损坏")
|
||||
self.verify(bound)
|
||||
self._publish_upload(bound)
|
||||
except (SourceError, ValueError, KeyError, TypeError, OSError):
|
||||
self.entries[directory.name] = {
|
||||
"public": {
|
||||
"id": directory.name,
|
||||
"label": "失效的已上传策略",
|
||||
"ready": False,
|
||||
"compatibleTasks": [],
|
||||
"error": "上传快照损坏,拒绝随机初始化;请重新上传原文件或恢复快照",
|
||||
}
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _display_name(name):
|
||||
if not isinstance(name, str):
|
||||
return "uploaded-policy"
|
||||
name = name.replace("\\", "/").split("/")[-1]
|
||||
return (
|
||||
re.sub(r"[^\w .()\-]", "_", name, flags=re.UNICODE)[:100].strip(" .")
|
||||
or "uploaded-policy"
|
||||
)
|
||||
|
||||
def _publish_upload(self, bound):
|
||||
record = {
|
||||
"id": bound["sourceId"],
|
||||
"label": bound["label"],
|
||||
"ready": True,
|
||||
"compatibleTasks": TASKS,
|
||||
"observationSizes": [47, 81, 97],
|
||||
"initialization": bound,
|
||||
}
|
||||
with self._catalog_lock:
|
||||
self.entries[bound["sourceId"]] = {"public": record}
|
||||
return deepcopy(record)
|
||||
|
||||
@contextmanager
|
||||
def upload_slot(self, length):
|
||||
# At most one receiving/decoding upload, 32 sources and 2GiB cumulative store.
|
||||
if not self._upload_lock.acquire(blocking=False):
|
||||
raise SourceError("已有文件正在上传/验证,请稍后重试")
|
||||
try:
|
||||
files = list(self.store.rglob("*"))
|
||||
used = sum(p.lstat().st_size for p in files if not p.is_symlink() and p.is_file())
|
||||
count = sum(1 for p in self.store.iterdir() if DIGEST.fullmatch(p.name))
|
||||
if count >= 32 or used + length + 16 * 1024**2 > 2 * 1024**3:
|
||||
raise SourceError("基础策略存储已达32个来源或2GiB上限,请由本机维护者释放空间")
|
||||
with tempfile.TemporaryDirectory(dir=self.store, prefix="upload-") as temporary:
|
||||
yield Path(temporary)
|
||||
finally:
|
||||
self._upload_lock.release()
|
||||
|
||||
def receive_upload(self, stream, length, fmt, template, display_name, *, set_timeout=None):
|
||||
if template != "go2-legacy47-v1":
|
||||
raise SourceError("必须明确确认go2-legacy47-v1观测与关节模板")
|
||||
limit = {"pt": 256 * 1024**2, "onnx": 64 * 1024**2}.get(fmt)
|
||||
if limit is None:
|
||||
raise SourceError("仅支持单个.pt或.onnx,不接受ZIP/路径/配套目录")
|
||||
if type(length) is not int or not 0 < length <= limit:
|
||||
raise SourceError("上传文件为空或超过.pt 256MiB / ONNX 64MiB上限")
|
||||
with self.upload_slot(length) as directory:
|
||||
path = directory / f"upload.{fmt}"
|
||||
deadline = time.monotonic() + 60
|
||||
remaining = length
|
||||
try:
|
||||
with path.open("xb") as output:
|
||||
while remaining:
|
||||
timeout = deadline - time.monotonic()
|
||||
if timeout <= 0:
|
||||
raise SourceError("上传超时,请重试")
|
||||
try:
|
||||
if set_timeout is not None:
|
||||
set_timeout(min(10, timeout))
|
||||
# BufferedReader.read(n) may perform many recv calls whose socket
|
||||
# timeouts reset with each trickled byte. read1 returns after one
|
||||
# raw read, allowing the total deadline to be checked every time.
|
||||
read_chunk = getattr(stream, "read1", stream.read)
|
||||
chunk = read_chunk(min(1024 * 1024, remaining))
|
||||
except OSError as error:
|
||||
raise SourceError("上传超时/连接中断,临时文件已清理") from error
|
||||
if not chunk or len(chunk) > remaining:
|
||||
raise SourceError("上传连接中断或实际长度不符")
|
||||
output.write(chunk)
|
||||
remaining -= len(chunk)
|
||||
finally:
|
||||
if set_timeout is not None:
|
||||
set_timeout(10)
|
||||
request = {
|
||||
"path": str(path),
|
||||
"directory": str(directory),
|
||||
"format": fmt,
|
||||
"template": template,
|
||||
}
|
||||
env = {
|
||||
**os.environ,
|
||||
"CUDA_VISIBLE_DEVICES": "",
|
||||
"OMP_NUM_THREADS": "1",
|
||||
"OPENBLAS_NUM_THREADS": "1",
|
||||
"MKL_NUM_THREADS": "1",
|
||||
}
|
||||
try:
|
||||
result = subprocess.run(
|
||||
[self.python, str(Path(__file__).parent / "rl/scripts/validate_upload.py")],
|
||||
input=json.dumps(request),
|
||||
capture_output=True,
|
||||
text=True,
|
||||
cwd=self.trainer_root,
|
||||
env=env,
|
||||
timeout=60,
|
||||
)
|
||||
manifest = json.loads(result.stdout)
|
||||
except (OSError, subprocess.TimeoutExpired, ValueError) as error:
|
||||
raise SourceError("模型验证超时/资源超限或训练Python依赖不可用") from error
|
||||
if result.returncode:
|
||||
raise SourceError(manifest.get("error", "不支持的模型文件"))
|
||||
digest = manifest["source_id"]
|
||||
if not DIGEST.fullmatch(digest):
|
||||
raise SourceError("模型验证返回无效身份")
|
||||
label = self._display_name(display_name)
|
||||
(directory / "upload.json").write_text(json.dumps(manifest), encoding="utf-8")
|
||||
(directory / "label.json").write_text(json.dumps({"label": label}), encoding="utf-8")
|
||||
destination = self.store / digest
|
||||
if destination.exists():
|
||||
# Dedup never replaces a committed artifact or changes an old job binding.
|
||||
old = next((r for r in self.catalog() if r["id"] == digest and r["ready"]), None)
|
||||
if old is None:
|
||||
raise SourceError("同ID旧快照已损坏,请恢复快照后重试;不会覆盖旧任务来源")
|
||||
self.verify(old["initialization"])
|
||||
return old
|
||||
for file in directory.iterdir():
|
||||
file.chmod(0o444)
|
||||
directory.rename(destination)
|
||||
destination.chmod(0o555)
|
||||
bound = {
|
||||
"sourceId": digest,
|
||||
"registeredId": digest,
|
||||
"label": label,
|
||||
"manifest": manifest,
|
||||
}
|
||||
self.verify(bound)
|
||||
return self._publish_upload(bound)
|
||||
|
||||
def _validate(self, directory: Path, checkpoint: str, task_id: str, task_config: dict | None):
|
||||
command = [self.python, str(Path(__file__).parent / "rl/scripts/validate_pretrained.py")]
|
||||
request = {
|
||||
"directory": str(directory),
|
||||
"checkpoint": checkpoint,
|
||||
"taskId": task_id,
|
||||
"taskConfig": task_config,
|
||||
"uploaded": (directory / "upload.json").is_file(),
|
||||
}
|
||||
try:
|
||||
result = subprocess.run(
|
||||
command,
|
||||
input=json.dumps(request),
|
||||
cwd=self.trainer_root,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=60,
|
||||
)
|
||||
except (subprocess.TimeoutExpired, OSError) as error:
|
||||
raise SourceError("基础策略验证器超时或不可用,请检查训练Python环境") from error
|
||||
if result.returncode:
|
||||
# Validator emits only an actionable error, never a traceback or file contents.
|
||||
raise SourceError(
|
||||
result.stderr.strip()[-1500:] or "基础策略验证器失败,请检查训练Python依赖"
|
||||
)
|
||||
try:
|
||||
return json.loads(result.stdout)
|
||||
except ValueError as error:
|
||||
raise SourceError("基础策略验证器未返回有效身份描述") from error
|
||||
|
||||
def catalog(self):
|
||||
with self._catalog_lock:
|
||||
sources = {}
|
||||
for entry in self.entries.values():
|
||||
sources.setdefault(entry["public"]["id"], deepcopy(entry["public"]))
|
||||
return list(sources.values())
|
||||
|
||||
def bind(self, source_id, task_id: str, task_config=None):
|
||||
# Selection is content-addressed, so a stale UI cannot silently bind a newly
|
||||
# registered model under the same administrator-friendly registration name.
|
||||
entry = next(
|
||||
(entry for entry in self.catalog() if entry["id"] == source_id),
|
||||
None,
|
||||
)
|
||||
if not isinstance(source_id, str) or entry is None:
|
||||
raise SourceError("基础策略内容ID不存在或已变化;请刷新服务连接,不接受文件路径")
|
||||
if not entry["ready"]:
|
||||
raise SourceError(entry["error"])
|
||||
if task_id not in entry["compatibleTasks"]:
|
||||
raise SourceError("该基础策略仅兼容Flat47/Obstacle81或97;Rough高度扫描不支持迁移")
|
||||
bound = deepcopy(entry["initialization"])
|
||||
directory = self.verify(bound)
|
||||
manifest = self._validate(
|
||||
directory, bound["manifest"]["artifacts"]["checkpoint"]["name"], task_id, task_config
|
||||
)
|
||||
if manifest["source_id"] != bound["sourceId"]:
|
||||
raise SourceError("基础策略快照身份变化,拒绝初始化")
|
||||
return bound
|
||||
|
||||
def verify(self, bound: dict) -> Path:
|
||||
try:
|
||||
digest = bound["sourceId"]
|
||||
if not isinstance(digest, str) or not DIGEST.fullmatch(digest):
|
||||
raise SourceError("基础策略快照身份无效")
|
||||
directory = self.store / digest
|
||||
artifacts = bound["manifest"]["artifacts"]
|
||||
if "sourceFormat" in bound["manifest"]:
|
||||
fmt = bound["manifest"]["sourceFormat"]
|
||||
if fmt not in ("pt", "onnx"):
|
||||
raise SourceError("上传格式无效")
|
||||
stored = json.loads(regular_bytes(directory / "upload.json", self.store, 64 * 1024))
|
||||
if stored != bound["manifest"]:
|
||||
raise SourceError("上传manifest SHA绑定变化,拒绝初始化")
|
||||
for name, relative, limit in (
|
||||
(
|
||||
"upload",
|
||||
f"upload.{fmt}",
|
||||
LIMITS["checkpoint"] if fmt == "pt" else LIMITS["onnx"],
|
||||
),
|
||||
("checkpoint", "actor.pt", 16 * 1024**2),
|
||||
):
|
||||
content = regular_bytes(directory / relative, self.store, limit)
|
||||
if (
|
||||
artifacts[name]["name"] != relative
|
||||
or hashlib.sha256(content).hexdigest() != artifacts[name]["sha256"]
|
||||
):
|
||||
raise SourceError("上传快照SHA不匹配,拒绝训练")
|
||||
return directory
|
||||
for name, relative in {
|
||||
"checkpoint": artifacts["checkpoint"]["name"],
|
||||
"onnx": "policy.onnx",
|
||||
"env": "params/env.yaml",
|
||||
"agent": "params/agent.yaml",
|
||||
}.items():
|
||||
content = regular_bytes(directory / relative, self.store, LIMITS[name])
|
||||
if hashlib.sha256(content).hexdigest() != artifacts[name]["sha256"]:
|
||||
raise SourceError("基础策略快照SHA不匹配,拒绝训练;请恢复原快照")
|
||||
return directory
|
||||
except (KeyError, TypeError) as error:
|
||||
raise SourceError("持久化基础策略描述损坏,拒绝训练") from error
|
||||
|
||||
def arguments(self, bound: dict) -> list[str]:
|
||||
directory = self.verify(bound)
|
||||
upload_args = (
|
||||
["--pretrained-upload-manifest", str(directory / "upload.json")]
|
||||
if "sourceFormat" in bound["manifest"]
|
||||
else []
|
||||
)
|
||||
return upload_args + [
|
||||
"--pretrained-checkpoint",
|
||||
str(directory / bound["manifest"]["artifacts"]["checkpoint"]["name"]),
|
||||
"--pretrained-allowed-roots",
|
||||
json.dumps([str(directory)]),
|
||||
"--pretrained-source-id",
|
||||
bound["sourceId"],
|
||||
]
|
||||
@@ -0,0 +1,336 @@
|
||||
"""Single-file Go2 legacy47 import. No adjacent files or arbitrary ONNX conversion."""
|
||||
|
||||
import hashlib
|
||||
import io
|
||||
import json
|
||||
from dataclasses import asdict
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from pretrained import (
|
||||
BASE_TERMS,
|
||||
JOINTS,
|
||||
NORMALIZATION_POLICY,
|
||||
PretrainedError,
|
||||
ValidatedSource,
|
||||
_require,
|
||||
make_reference_actor,
|
||||
validate_actor_state,
|
||||
validate_semantics,
|
||||
verify_onnx,
|
||||
)
|
||||
|
||||
TEMPLATE = "go2-legacy47-v1"
|
||||
SYNTHETIC_COUNT = 1_000_000
|
||||
UPLOAD_LIMITS = {"pt": 256 * 1024**2, "onnx": 64 * 1024**2}
|
||||
SEMANTICS = {
|
||||
"joint_names": JOINTS,
|
||||
"observation_names": BASE_TERMS,
|
||||
"command_names": ["twist"],
|
||||
"action_scale": [0.25],
|
||||
"joint_stiffness": [20, 20, 40] * 4,
|
||||
"joint_damping": [1, 1, 2] * 4,
|
||||
"default_joint_pos": [-0.1, 0.9, -1.8, 0.1, 0.9, -1.8] * 2,
|
||||
}
|
||||
|
||||
|
||||
def validate_metadata(metadata, *, required=False):
|
||||
"""Check recognized embedded semantics; absence remains a user template assumption."""
|
||||
import numpy as np
|
||||
|
||||
_require(isinstance(metadata, dict), "模型metadata必须是对象")
|
||||
checked = []
|
||||
for key, expected in SEMANTICS.items():
|
||||
if key not in metadata:
|
||||
_require(not required, f"ONNX缺少语义metadata: {key}")
|
||||
continue
|
||||
actual = metadata[key]
|
||||
if isinstance(actual, str):
|
||||
actual = actual.split(",")
|
||||
if isinstance(expected[0], str):
|
||||
_require(actual == expected, f"模型metadata冲突: {key}")
|
||||
else:
|
||||
try:
|
||||
actual = np.asarray(actual, dtype=np.float64)
|
||||
_require(
|
||||
actual.shape == (len(expected),)
|
||||
and np.isfinite(actual).all()
|
||||
and np.allclose(actual, expected, atol=1e-6, rtol=0),
|
||||
f"模型metadata冲突: {key}",
|
||||
)
|
||||
except (TypeError, ValueError) as error:
|
||||
raise PretrainedError(f"模型metadata无效: {key}") from error
|
||||
checked.append(key)
|
||||
for key in ("contract", "templateId"):
|
||||
if key in metadata:
|
||||
_require(metadata[key] == TEMPLATE, "模型模板metadata冲突")
|
||||
checked.append(key)
|
||||
return checked
|
||||
|
||||
|
||||
def _checkpoint(data):
|
||||
try:
|
||||
checkpoint = torch.load(io.BytesIO(data), map_location="cpu", weights_only=True)
|
||||
except Exception as error:
|
||||
raise PretrainedError(
|
||||
"不支持或不安全的.pt;仅支持weights_only Go2 legacy47 actor checkpoint"
|
||||
) from error
|
||||
_require(
|
||||
isinstance(checkpoint, dict) and "actor_state_dict" in checkpoint,
|
||||
"请选择含actor_state_dict的Go2 legacy47 .pt,不接受ZIP工程或任意模型",
|
||||
)
|
||||
state = checkpoint["actor_state_dict"]
|
||||
if isinstance(state, dict) and isinstance(state.get("mlp.0.weight"), torch.Tensor):
|
||||
_require(
|
||||
tuple(state["mlp.0.weight"].shape) == (512, 47),
|
||||
"当前仅支持Go2 legacy47输入;81/97 checkpoint请使用原trial续训,而不是基础策略上传",
|
||||
)
|
||||
validate_actor_state(state)
|
||||
iteration = checkpoint.get("iter")
|
||||
_require(
|
||||
iteration is None or (type(iteration) is int and iteration >= 0), "无效checkpoint iteration"
|
||||
)
|
||||
checked = validate_metadata(checkpoint)
|
||||
for key in ("metadata", "infos"):
|
||||
if key in checkpoint:
|
||||
checked.extend(validate_metadata(checkpoint[key]))
|
||||
nested = checkpoint[key].get("metadata")
|
||||
if nested is not None:
|
||||
checked.extend(validate_metadata(nested))
|
||||
return state, iteration, sorted(set(checked))
|
||||
|
||||
|
||||
def _onnx(data):
|
||||
import onnx
|
||||
from onnx import helper, numpy_helper
|
||||
|
||||
model = onnx.load_model_from_string(data)
|
||||
graph = model.graph
|
||||
_require(
|
||||
not model.functions
|
||||
and not model.training_info
|
||||
and len(model.opset_import) == 1
|
||||
and model.opset_import[0].domain == ""
|
||||
and model.opset_import[0].version in (17, 18),
|
||||
"仅支持标准opset17/18受限MLP ONNX",
|
||||
)
|
||||
_require(
|
||||
not graph.sparse_initializer and not graph.quantization_annotation, "不支持稀疏或量化ONNX"
|
||||
)
|
||||
_require(len(graph.input) == len(graph.output) == 1, "ONNX必须单输入单输出")
|
||||
for value, shape in ((graph.input[0], [1, 47]), (graph.output[0], [1, 12])):
|
||||
tensor = value.type.tensor_type
|
||||
_require(
|
||||
tensor.elem_type == onnx.TensorProto.FLOAT
|
||||
and [d.dim_value for d in tensor.shape.dim] == shape
|
||||
and all(not d.dim_param for d in tensor.shape.dim),
|
||||
"ONNX仅支持float32 [1,47] -> [1,12];其他输入请使用对应训练器",
|
||||
)
|
||||
expected_shapes = {"obs_normalizer._mean": (1, 47), "onnx::Div_24": (1, 47)}
|
||||
for i, shape in zip((0, 2, 4, 6), ((512, 47), (256, 512), (128, 256), (12, 128)), strict=True):
|
||||
expected_shapes[f"mlp.{i}.weight"] = shape
|
||||
expected_shapes[f"mlp.{i}.bias"] = (shape[0],)
|
||||
_require(
|
||||
len(graph.initializer) == len(expected_shapes)
|
||||
and {t.name for t in graph.initializer} == set(expected_shapes),
|
||||
"ONNX initializer不符合受支持MLP",
|
||||
)
|
||||
tensors = {}
|
||||
for tensor in graph.initializer:
|
||||
_require(
|
||||
tensor.data_location == onnx.TensorProto.DEFAULT and not tensor.external_data,
|
||||
"拒绝ONNX external data;必须是单个自包含文件",
|
||||
)
|
||||
_require(
|
||||
tensor.data_type == onnx.TensorProto.FLOAT
|
||||
and tuple(tensor.dims) == expected_shapes[tensor.name],
|
||||
"ONNX tensor类型/shape不支持",
|
||||
)
|
||||
tensors[tensor.name] = torch.from_numpy(numpy_helper.to_array(tensor).copy())
|
||||
_require(bool(torch.isfinite(tensors[tensor.name]).all()), "ONNX tensor含非有限值")
|
||||
# Exact dataflow, not merely op/tensor names: no branch, reorder, alias or extra op.
|
||||
expected_ops = ["Sub", "Div", "Gemm", "Elu", "Gemm", "Elu", "Gemm", "Elu", "Gemm"]
|
||||
_require(
|
||||
[n.op_type for n in graph.node] == expected_ops, "ONNX必须是Sub/Div及4层Gemm+3层ELU精确链路"
|
||||
)
|
||||
previous = graph.input[0].name
|
||||
seen = set(tensors) | {previous}
|
||||
_require(len(seen) == len(tensors) + 1, "ONNX输入与initializer重名")
|
||||
layer = 0
|
||||
for node in graph.node:
|
||||
_require(node.domain == "" and not node.overload, "拒绝ONNX custom op")
|
||||
attributes = {a.name: helper.get_attribute_value(a) for a in node.attribute}
|
||||
_require(len(attributes) == len(node.attribute), "重复ONNX属性")
|
||||
if node.op_type in ("Sub", "Div"):
|
||||
inputs = [previous, "obs_normalizer._mean" if node.op_type == "Sub" else "onnx::Div_24"]
|
||||
_require(not attributes, "不支持normalizer算子属性")
|
||||
elif node.op_type == "Gemm":
|
||||
inputs = [previous, f"mlp.{layer}.weight", f"mlp.{layer}.bias"]
|
||||
layer += 2
|
||||
_require(
|
||||
set(attributes) <= {"alpha", "beta", "transA", "transB"}
|
||||
and attributes.get("alpha", 1.0) == 1.0
|
||||
and attributes.get("beta", 1.0) == 1.0
|
||||
and attributes.get("transA", 0) == 0
|
||||
and attributes.get("transB", 0) == 1,
|
||||
"不支持Gemm缩放/转置属性",
|
||||
)
|
||||
else:
|
||||
inputs = [previous]
|
||||
_require(
|
||||
set(attributes) <= {"alpha"} and attributes.get("alpha", 1.0) == 1.0,
|
||||
"仅支持ELU alpha=1",
|
||||
)
|
||||
_require(
|
||||
list(node.input) == inputs
|
||||
and len(node.output) == 1
|
||||
and node.output[0]
|
||||
and node.output[0] not in seen,
|
||||
"ONNX实际连边/输出不符合受支持MLP",
|
||||
)
|
||||
previous = node.output[0]
|
||||
seen.add(previous)
|
||||
_require(previous == graph.output[0].name, "ONNX输出必须是最后Gemm结果")
|
||||
metadata = {p.key: p.value for p in model.metadata_props}
|
||||
_require(len(metadata) == len(model.metadata_props), "重复ONNX metadata")
|
||||
checked = validate_metadata(metadata, required=True)
|
||||
onnx.checker.check_model(model, full_check=True)
|
||||
actor = make_reference_actor()
|
||||
_require(
|
||||
actor.obs_normalizer.eps == 0.01
|
||||
and bool((actor.state_dict()["distribution.std_param"] == 1).all()),
|
||||
"目标训练默认normalizer/exploration已变化,需要新模板",
|
||||
)
|
||||
state = actor.state_dict()
|
||||
for key in state:
|
||||
if key in tensors:
|
||||
state[key] = tensors[key]
|
||||
std = tensors["onnx::Div_24"] - 0.01
|
||||
_require(bool((std > 0).all()), "ONNX denominator必须大于模板epsilon=.01")
|
||||
state["obs_normalizer._std"] = std
|
||||
state["obs_normalizer._var"] = std.square()
|
||||
state["obs_normalizer.count"].fill_(SYNTHETIC_COUNT)
|
||||
validate_actor_state(state)
|
||||
actor.load_state_dict(state)
|
||||
identity = verify_onnx(data, actor)
|
||||
return state, checked, identity
|
||||
|
||||
|
||||
def source_identity(fmt, digest):
|
||||
return hashlib.sha256(f"{TEMPLATE}:{fmt}:{digest}".encode()).hexdigest()
|
||||
|
||||
|
||||
def import_upload(path, fmt, template, directory):
|
||||
"""Run only inside the resource-limited validator process."""
|
||||
_require(template == TEMPLATE, "必须明确确认go2-legacy47-v1模板")
|
||||
_require(fmt in UPLOAD_LIMITS, "仅支持单个.pt或.onnx,不支持ZIP")
|
||||
data = Path(path).read_bytes()
|
||||
_require(0 < len(data) <= UPLOAD_LIMITS[fmt], "上传文件为空或过大")
|
||||
if fmt == "pt":
|
||||
state, iteration, checked = _checkpoint(data)
|
||||
identity = None
|
||||
else:
|
||||
state, checked, identity = _onnx(data)
|
||||
iteration = None
|
||||
from pretrained import comparison_observations
|
||||
|
||||
actor = make_reference_actor().eval()
|
||||
actor.load_state_dict(state)
|
||||
with torch.inference_mode():
|
||||
output = actor.mlp(actor.obs_normalizer(comparison_observations()))
|
||||
_require(bool(torch.isfinite(output).all()), "actor在随机/物理probe上产生非有限动作")
|
||||
directory = Path(directory)
|
||||
actor_path = directory / "actor.pt"
|
||||
torch.save({"actor_state_dict": state}, actor_path)
|
||||
digest = hashlib.sha256(data).hexdigest()
|
||||
manifest = {
|
||||
"schema_version": 1,
|
||||
"mode": "pretrained-warm-start",
|
||||
"contract": TEMPLATE,
|
||||
"sourceFormat": fmt,
|
||||
"source_id": source_identity(fmt, digest),
|
||||
"source_iteration": iteration,
|
||||
"source_actor_dim": 47,
|
||||
"source_normalizer_count": state["obs_normalizer.count"].item(),
|
||||
"normalization": NORMALIZATION_POLICY
|
||||
if fmt == "pt"
|
||||
else "synthetic-count/unit-new-features",
|
||||
"template_confirmation": {
|
||||
"id": TEMPLATE,
|
||||
"confirmed_by": "user",
|
||||
"assumptions": (
|
||||
"Go2 legacy47 observation physics, ordering, 50Hz and action semantics; "
|
||||
"not verified source env.yaml"
|
||||
),
|
||||
},
|
||||
"verified_facts": {
|
||||
"actor_tensor_shapes": "47-512-256-128-12/float32",
|
||||
"metadata_fields": checked,
|
||||
"finite_output_probes": 48,
|
||||
"activation": "graph-verified-ELU" if fmt == "onnx" else "user-template-assumed-ELU",
|
||||
},
|
||||
"derived_fields": {}
|
||||
if fmt == "pt"
|
||||
else {
|
||||
"normalizer_count": {"policy": "synthetic", "value": SYNTHETIC_COUNT},
|
||||
"normalizer_std": "denominator - 0.01",
|
||||
"normalizer_var": "std squared",
|
||||
"epsilon": 0.01,
|
||||
"exploration_std": {"policy": "fresh-target-default", "value": 1.0},
|
||||
},
|
||||
"onnx_identity": identity,
|
||||
"critic": "fresh-target-initialization",
|
||||
"optimizer": "fresh",
|
||||
"iteration": 0,
|
||||
"base_observation_terms": BASE_TERMS,
|
||||
"joint_names": JOINTS,
|
||||
"artifacts": {
|
||||
"upload": {"name": f"upload.{fmt}", "sha256": digest, "bytes": len(data)},
|
||||
"checkpoint": {
|
||||
"name": "actor.pt",
|
||||
"sha256": hashlib.sha256(actor_path.read_bytes()).hexdigest(),
|
||||
"bytes": actor_path.stat().st_size,
|
||||
"origin": "service-derived-actor-only",
|
||||
},
|
||||
},
|
||||
}
|
||||
return manifest
|
||||
|
||||
|
||||
def read_uploaded_source(checkpoint, *, allowed_roots, manifest_path, target_env, target_agent):
|
||||
"""Load only the service-derived actor and bound provenance, never adjacent sidecars."""
|
||||
from pretrained_sources import regular_bytes
|
||||
|
||||
roots = [Path(p).resolve(strict=True) for p in allowed_roots]
|
||||
checkpoint, manifest_path = Path(checkpoint), Path(manifest_path)
|
||||
root = next(
|
||||
(r for r in roots if checkpoint.is_relative_to(r) and manifest_path.is_relative_to(r)), None
|
||||
)
|
||||
_require(root is not None, "上传artifact不在受控根内")
|
||||
manifest = json.loads(regular_bytes(manifest_path, root, 64 * 1024))
|
||||
_require(
|
||||
manifest.get("contract") == TEMPLATE and manifest.get("sourceFormat") in UPLOAD_LIMITS,
|
||||
"无效上传manifest模板/格式",
|
||||
)
|
||||
artifacts = manifest["artifacts"]
|
||||
_require(
|
||||
manifest["source_id"]
|
||||
== source_identity(manifest["sourceFormat"], artifacts["upload"]["sha256"]),
|
||||
"上传原始SHA身份不匹配",
|
||||
)
|
||||
data = regular_bytes(checkpoint, root, UPLOAD_LIMITS["pt"])
|
||||
_require(
|
||||
hashlib.sha256(data).hexdigest() == artifacts["checkpoint"]["sha256"], "上传actor SHA不匹配"
|
||||
)
|
||||
from pretrained import _plain
|
||||
from src.tasks.velocity.config.go2.env_cfgs import unitree_go2_flat_env_cfg
|
||||
from src.tasks.velocity.config.go2.rl_cfg import unitree_go2_ppo_runner_cfg
|
||||
|
||||
validate_semantics(
|
||||
_plain(asdict(unitree_go2_flat_env_cfg())),
|
||||
_plain(asdict(unitree_go2_ppo_runner_cfg())),
|
||||
target_env,
|
||||
target_agent,
|
||||
)
|
||||
state, _, _ = _checkpoint(data)
|
||||
return ValidatedSource(state, manifest)
|
||||
@@ -7,3 +7,7 @@ mujoco-warp==3.5.0
|
||||
warp-lang==1.15.0
|
||||
# RSL-RL 5.0.1 仍向 wandb.Settings 传递 start_method,0.29 已删除该字段。
|
||||
wandb==0.28.2
|
||||
# 基础策略迁移依赖已验证的 actor/normalizer/checkpoint 契约。
|
||||
rsl-rl-lib==5.0.1
|
||||
# 基础策略身份校验只使用 CPU ONNX Runtime。
|
||||
onnxruntime==1.29.0
|
||||
|
||||
@@ -32,7 +32,6 @@ from mjlab.tasks.registry import list_tasks, load_env_cfg, load_rl_cfg, load_run
|
||||
from mjlab.tasks.velocity.mdp import UniformVelocityCommandCfg
|
||||
from mjlab.utils.torch import configure_torch_backends
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
|
||||
from tuning.schema import apply_reward_configuration, validate_configuration
|
||||
|
||||
SCENARIOS = (
|
||||
@@ -61,6 +60,7 @@ class EvaluateConfig:
|
||||
checkpoint: str
|
||||
output: str
|
||||
reward_config: str | None = None
|
||||
task_config: str | None = None
|
||||
num_envs: int = 256
|
||||
steps_per_seed: int = 1000
|
||||
seeds: tuple[int, ...] = field(default_factory=lambda: (101, 202, 303))
|
||||
@@ -163,31 +163,41 @@ def run_evaluation(task_id: str, cfg: EvaluateConfig) -> dict:
|
||||
checkpoint = Path(cfg.checkpoint).expanduser().resolve(strict=True)
|
||||
output = Path(cfg.output).expanduser().resolve()
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
per_seed = [_evaluate_seed(task_id, cfg, seed) for seed in cfg.seeds]
|
||||
metrics = {
|
||||
key: fmean(seed_metrics[key] for seed_metrics in per_seed)
|
||||
for key in (*METRIC_NAMES, "fall_rate")
|
||||
}
|
||||
deviations = {
|
||||
key: pstdev(seed_metrics[key] for seed_metrics in per_seed)
|
||||
for key in (*METRIC_NAMES, "fall_rate")
|
||||
}
|
||||
result = {
|
||||
"protocolVersion": 1,
|
||||
"taskId": task_id,
|
||||
"checkpoint": checkpoint.name,
|
||||
"checkpointSha256": _sha256(checkpoint),
|
||||
"seeds": list(cfg.seeds),
|
||||
"numEnvs": cfg.num_envs,
|
||||
"stepsPerSeed": cfg.steps_per_seed,
|
||||
"scenarios": [list(value) for value in SCENARIOS],
|
||||
"metrics": metrics,
|
||||
"metricStd": deviations,
|
||||
"seedMetrics": [
|
||||
{"seed": seed, "metrics": values}
|
||||
for seed, values in zip(cfg.seeds, per_seed, strict=True)
|
||||
],
|
||||
}
|
||||
if task_id == "Unitree-Go2-ObstacleAvoidance":
|
||||
from scripts.evaluate_obstacle import run_obstacle_evaluation
|
||||
checkpoint_hash = _sha256(checkpoint)
|
||||
result = run_obstacle_evaluation(task_id, cfg)
|
||||
if _sha256(checkpoint) != checkpoint_hash:
|
||||
raise ValueError("Checkpoint changed during evaluation")
|
||||
result.update({"taskId": task_id, "checkpoint": checkpoint.name,
|
||||
"checkpointSha256": checkpoint_hash})
|
||||
metrics = result["metrics"]
|
||||
else:
|
||||
per_seed = [_evaluate_seed(task_id, cfg, seed) for seed in cfg.seeds]
|
||||
metrics = {
|
||||
key: fmean(seed_metrics[key] for seed_metrics in per_seed)
|
||||
for key in (*METRIC_NAMES, "fall_rate")
|
||||
}
|
||||
deviations = {
|
||||
key: pstdev(seed_metrics[key] for seed_metrics in per_seed)
|
||||
for key in (*METRIC_NAMES, "fall_rate")
|
||||
}
|
||||
result = {
|
||||
"protocolVersion": 1,
|
||||
"taskId": task_id,
|
||||
"checkpoint": checkpoint.name,
|
||||
"checkpointSha256": _sha256(checkpoint),
|
||||
"seeds": list(cfg.seeds),
|
||||
"numEnvs": cfg.num_envs,
|
||||
"stepsPerSeed": cfg.steps_per_seed,
|
||||
"scenarios": [list(value) for value in SCENARIOS],
|
||||
"metrics": metrics,
|
||||
"metricStd": deviations,
|
||||
"seedMetrics": [
|
||||
{"seed": seed, "metrics": values}
|
||||
for seed, values in zip(cfg.seeds, per_seed, strict=True)
|
||||
],
|
||||
}
|
||||
output.write_text(json.dumps(result, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
|
||||
writer = SummaryWriter(log_dir=str(output.parent / "evaluation-events"))
|
||||
try:
|
||||
|
||||
@@ -0,0 +1,241 @@
|
||||
"""Actual checkpoint inference, capturing first-terminal physics before mjlab resets."""
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
from dataclasses import asdict
|
||||
from pathlib import Path
|
||||
from statistics import fmean
|
||||
from types import SimpleNamespace
|
||||
|
||||
# Also executable as a fresh-interpreter seed worker (never fork a CUDA context).
|
||||
for source in (Path(__file__).resolve().parents[1], Path(__file__).resolve().parents[2]):
|
||||
if str(source) not in sys.path:
|
||||
sys.path.insert(0, str(source))
|
||||
|
||||
import torch
|
||||
from mjlab.envs import ManagerBasedRlEnv
|
||||
from mjlab.rl import MjlabOnPolicyRunner, RslRlVecEnvWrapper
|
||||
from mjlab.sensor import ContactMatch, ContactSensorCfg
|
||||
from mjlab.tasks.registry import load_env_cfg, load_rl_cfg, load_runner_cls
|
||||
from scripts.train import _load_reward_config, _load_task_config
|
||||
from src.tasks.obstacle_avoidance.env_cfg import apply_obstacle_configuration
|
||||
from task_config import OBSTACLE_TASK
|
||||
from tuning.obstacle_scoring import METRICS, STEPS, protocol, score_trajectory
|
||||
from tuning.schema import apply_reward_configuration
|
||||
|
||||
|
||||
class FirstEpisodeRecorder:
|
||||
"""The termination hook runs post-physics, before _reset_idx (including first terminal)."""
|
||||
|
||||
def __init__(self, env, layout):
|
||||
self.env = env
|
||||
self.layout = layout
|
||||
self.samples = [[] for _ in range(env.num_envs)]
|
||||
self.previous = torch.zeros((env.num_envs, 12), device=env.device)
|
||||
self.calls = 0
|
||||
self.compute = env.termination_manager.compute
|
||||
env.termination_manager.compute = self.capture
|
||||
|
||||
def capture(self):
|
||||
env = self.env
|
||||
terminal = self.compute()
|
||||
robot = env.scene["robot"].data
|
||||
# Like mjlab termination itself, derived pose is one physics substep old.
|
||||
position = robot.root_link_pos_w
|
||||
local_xy = position[:, :2] - env.scene.env_origins[:, :2]
|
||||
xy = local_xy + torch.tensor(self.layout["spawn"][:2], device=env.device)
|
||||
boxes = self.layout["boxes"][1:] # Never include the support floor.
|
||||
clearance = torch.full((env.num_envs,), 0.5, device=env.device)
|
||||
if boxes:
|
||||
centers = torch.tensor([b["pos"][:2] for b in boxes], device=env.device)
|
||||
sizes = torch.tensor([b["size"][:2] for b in boxes], device=env.device)
|
||||
clearance = (
|
||||
((xy[:, None, :] - centers).abs() - sizes).clamp(min=0).norm(dim=2).amin(dim=1)
|
||||
)
|
||||
clearance = (clearance - 0.3).clamp(min=0) # Conservative body footprint radius.
|
||||
forces = env.scene["evaluation_obstacles"].data.force_history
|
||||
collision = forces.norm(dim=-1).flatten(1).amax(dim=1) > 1.0
|
||||
else:
|
||||
collision = torch.zeros(env.num_envs, dtype=torch.bool, device=env.device)
|
||||
nonfoot = env.scene["nonfoot_ground_touch"].data.force_history
|
||||
collision |= nonfoot.norm(dim=-1).flatten(1).amax(dim=1) > 10.0
|
||||
fall = (robot.projected_gravity_b[:, 2] > -0.3420201433) | (position[:, 2] < 0.12)
|
||||
actions = env.action_manager.action
|
||||
delta = (actions - self.previous).square().mean(dim=1)
|
||||
self.previous.copy_(actions)
|
||||
rays = env.scene["forward_scan"].data.distances
|
||||
hits = (
|
||||
((rays >= 0) & (rays <= env.scene["forward_scan"].cfg.max_distance)).float().mean(dim=1)
|
||||
)
|
||||
distance = env.command_manager.get_term("twist").errors()[1]
|
||||
values = torch.stack((distance, clearance, delta, hits, collision, fall, terminal), dim=1)
|
||||
if not torch.isfinite(values).all():
|
||||
raise ValueError("Nonfinite rollout measurements")
|
||||
keys = ("distance", "clearance", "action_delta", "ray_hit", "collision", "fall", "terminal")
|
||||
for samples, row in zip(self.samples, values.cpu().tolist(), strict=True):
|
||||
if not samples or not samples[-1]["terminal"]:
|
||||
samples.append(dict(zip(keys, row, strict=True)))
|
||||
self.calls += 1
|
||||
return terminal
|
||||
|
||||
def metrics(self, horizon=STEPS):
|
||||
if self.calls != horizon:
|
||||
raise ValueError("Incomplete rollout horizon")
|
||||
metrics = [score_trajectory(samples, horizon) for samples in self.samples]
|
||||
return {key: fmean(item[key] for item in metrics) for key in METRICS}
|
||||
|
||||
|
||||
def configure_seed(task_id, cfg, scenario, reward):
|
||||
env_cfg = load_env_cfg(task_id, play=False)
|
||||
env_cfg.seed = scenario["seed"]
|
||||
env_cfg.scene.num_envs = cfg.num_envs
|
||||
# Keep the benchmark's declared pair fixed; training itself randomizes every reset.
|
||||
apply_obstacle_configuration(env_cfg, scenario["taskConfig"], randomize_navigation=False)
|
||||
apply_reward_configuration(env_cfg, reward, task_id)
|
||||
env_cfg.curriculum = {}
|
||||
env_cfg.observations["actor"].enable_corruption = False
|
||||
env_cfg.events.pop("push_robot", None)
|
||||
# Primary geom names are literal compiled static terrain geoms, not a regex.
|
||||
names = tuple(f"terrain_{i}" for i in range(1, len(scenario["terrain"]["boxes"])))
|
||||
if names:
|
||||
env_cfg.scene.sensors += (
|
||||
ContactSensorCfg(
|
||||
name="evaluation_obstacles",
|
||||
primary=ContactMatch(mode="geom", pattern=names),
|
||||
secondary=ContactMatch(mode="subtree", pattern="base_link", entity="robot"),
|
||||
fields=("force",),
|
||||
reduce="maxforce",
|
||||
history_length=env_cfg.decimation,
|
||||
),
|
||||
)
|
||||
return env_cfg
|
||||
|
||||
|
||||
def evaluate_seed(task_id, cfg, scenario, reward):
|
||||
torch.manual_seed(scenario["seed"])
|
||||
env_cfg = configure_seed(task_id, cfg, scenario, reward)
|
||||
agent_cfg = load_rl_cfg(task_id)
|
||||
env = ManagerBasedRlEnv(env_cfg, device=cfg.device)
|
||||
wrapped = RslRlVecEnvWrapper(env, clip_actions=agent_cfg.clip_actions)
|
||||
try:
|
||||
runner_cls = load_runner_cls(task_id) or MjlabOnPolicyRunner
|
||||
runner = runner_cls(wrapped, asdict(agent_cfg), log_dir=None, device=wrapped.device)
|
||||
# rsl_rl actor state_dict includes observation_normalizer buffers. Strict load
|
||||
# rejects incompatible architecture/statistics; no fixture or zero-action fallback.
|
||||
runner.load(
|
||||
str(Path(cfg.checkpoint).resolve(strict=True)),
|
||||
load_cfg={"actor": True},
|
||||
strict=True,
|
||||
map_location=str(wrapped.device),
|
||||
)
|
||||
policy = runner.get_inference_policy(device=str(wrapped.device))
|
||||
obs, _ = env.reset(seed=scenario["seed"])
|
||||
recorder = FirstEpisodeRecorder(env, scenario["terrain"])
|
||||
with torch.inference_mode():
|
||||
for _ in range(STEPS):
|
||||
actions = policy(obs)
|
||||
if not torch.isfinite(actions).all():
|
||||
raise ValueError("Nonfinite policy actions")
|
||||
obs, _, _, _ = wrapped.step(actions)
|
||||
return {
|
||||
"seed": scenario["seed"],
|
||||
"metrics": recorder.metrics(),
|
||||
"episodes": cfg.num_envs,
|
||||
"rolloutSteps": STEPS,
|
||||
}
|
||||
finally:
|
||||
wrapped.close()
|
||||
|
||||
|
||||
def run_obstacle_evaluation(task_id, cfg):
|
||||
from tuning.obstacle_scoring import SEEDS
|
||||
|
||||
if tuple(cfg.seeds) != SEEDS or cfg.steps_per_seed != STEPS:
|
||||
raise ValueError("Obstacle evaluation seeds/horizon are fixed")
|
||||
if not cfg.task_config or not cfg.reward_config:
|
||||
raise ValueError("Obstacle evaluation requires task and parameter configurations")
|
||||
raw = json.loads(Path(cfg.task_config).read_text())
|
||||
custom = _load_task_config(task_id, cfg.task_config, raw["seed"])
|
||||
reward = _load_reward_config(cfg.reward_config, None, OBSTACLE_TASK)
|
||||
fixed = protocol(custom, cfg.num_envs)
|
||||
seeds = evaluate_isolated_seeds(task_id, cfg, fixed, reward)
|
||||
return {
|
||||
"protocol": fixed,
|
||||
"seedMetrics": seeds,
|
||||
"metrics": {key: fmean(item["metrics"][key] for item in seeds) for key in METRICS},
|
||||
}
|
||||
|
||||
|
||||
def file_hash(path):
|
||||
digest = hashlib.sha256()
|
||||
with Path(path).open("rb") as stream:
|
||||
for chunk in iter(lambda: stream.read(1024 * 1024), b""):
|
||||
digest.update(chunk)
|
||||
return digest.hexdigest()
|
||||
|
||||
|
||||
def evaluate_isolated_seeds(task_id, cfg, fixed, reward):
|
||||
from tuning.obstacle_scoring import validate_evaluation
|
||||
|
||||
seeds = []
|
||||
# Inherit the evaluation parent's process group: manager cancellation kills
|
||||
# both parent and the active seed. subprocess.run kills/reaps on timeout too.
|
||||
with tempfile.TemporaryDirectory(prefix="go2-obstacle-eval-") as directory:
|
||||
root = Path(directory)
|
||||
checkpoint = root / "checkpoint.pt"
|
||||
source_hash = file_hash(cfg.checkpoint)
|
||||
shutil.copyfile(cfg.checkpoint, checkpoint)
|
||||
if file_hash(checkpoint) != source_hash:
|
||||
raise ValueError("Checkpoint changed during snapshot")
|
||||
for scenario in fixed["scenarios"]:
|
||||
request = root / f"request-{scenario['seed']}.json"
|
||||
output = root / f"result-{scenario['seed']}.json"
|
||||
request.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"taskId": task_id,
|
||||
"config": {**asdict(cfg), "checkpoint": str(checkpoint)},
|
||||
"scenario": scenario,
|
||||
"reward": reward,
|
||||
"checkpointSha256": source_hash,
|
||||
},
|
||||
allow_nan=False,
|
||||
)
|
||||
)
|
||||
subprocess.run(
|
||||
[sys.executable, str(Path(__file__).resolve()), str(request), str(output)],
|
||||
check=True,
|
||||
timeout=1200,
|
||||
shell=False,
|
||||
)
|
||||
result = json.loads(output.read_text())
|
||||
if (
|
||||
result.pop("checkpointSha256", None) != source_hash
|
||||
or file_hash(checkpoint) != source_hash
|
||||
):
|
||||
raise ValueError("Seed checkpoint identity mismatch")
|
||||
seeds.append(result)
|
||||
candidate = {
|
||||
"protocol": fixed,
|
||||
"seedMetrics": seeds,
|
||||
"metrics": {key: fmean(item["metrics"][key] for item in seeds) for key in METRICS},
|
||||
}
|
||||
validate_evaluation(candidate, fixed)
|
||||
return seeds
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import src.tasks # noqa: F401
|
||||
|
||||
request_path, output_path = map(Path, sys.argv[1:])
|
||||
request = json.loads(request_path.read_text())
|
||||
config = SimpleNamespace(**request["config"])
|
||||
if file_hash(config.checkpoint) != request["checkpointSha256"]:
|
||||
raise ValueError("Worker checkpoint identity mismatch")
|
||||
result = evaluate_seed(request["taskId"], config, request["scenario"], request["reward"])
|
||||
result["checkpointSha256"] = file_hash(config.checkpoint)
|
||||
output_path.write_text(json.dumps(result, allow_nan=False))
|
||||
@@ -7,7 +7,7 @@ import sys
|
||||
from dataclasses import asdict, dataclass, field
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Literal, cast
|
||||
from typing import Literal
|
||||
|
||||
# 训练器作为仓库内置子集直接从 scripts/ 启动,不要求额外执行 pip install -e。
|
||||
TRAINER_ROOT = Path(__file__).resolve().parents[1]
|
||||
@@ -34,7 +34,7 @@ from mjlab.utils.gpu import select_gpus
|
||||
from mjlab.utils.os import dump_yaml, get_checkpoint_path
|
||||
from mjlab.utils.torch import configure_torch_backends
|
||||
from mjlab.utils.wrappers import VideoRecorder
|
||||
|
||||
from task_config import OBSTACLE_TASK, deployment_metadata, validate_task_config
|
||||
from tuning.schema import apply_reward_configuration, validate_configuration
|
||||
|
||||
|
||||
@@ -51,8 +51,14 @@ class TrainConfig:
|
||||
gpu_ids: list[int] | Literal["all"] | None = field(default_factory=lambda: [0])
|
||||
output_dir: str | None = None
|
||||
resume_checkpoint: str | None = None
|
||||
pretrained_checkpoint: str | None = None
|
||||
pretrained_upload_manifest: str | None = None
|
||||
pretrained_onnx: str | None = None
|
||||
pretrained_source_id: str | None = None
|
||||
pretrained_allowed_roots: list[str] = field(default_factory=list)
|
||||
reward_config: str | None = None
|
||||
reward_config_json: str | None = None
|
||||
task_config: str | None = None
|
||||
|
||||
@staticmethod
|
||||
def from_task(task_id: str) -> "TrainConfig":
|
||||
@@ -61,27 +67,96 @@ class TrainConfig:
|
||||
return TrainConfig(env=env_cfg, agent=agent_cfg)
|
||||
|
||||
|
||||
def _load_reward_config(path: str | None, inline: str | None) -> dict | None:
|
||||
def _load_reward_config(path: str | None, inline: str | None, task_id="Unitree-Go2-Flat") -> dict | None:
|
||||
if path is not None and inline is not None:
|
||||
raise ValueError("Use only one of reward_config and reward_config_json")
|
||||
if inline is not None:
|
||||
if len(inline.encode("utf-8")) > 64 * 1024:
|
||||
raise ValueError("Reward configuration is larger than 64 KiB")
|
||||
return validate_configuration(json.loads(inline))
|
||||
return validate_configuration(json.loads(inline), task_id)
|
||||
if path is None:
|
||||
return None
|
||||
source = Path(path).expanduser().resolve(strict=True)
|
||||
if source.stat().st_size > 64 * 1024:
|
||||
raise ValueError("Reward configuration is larger than 64 KiB")
|
||||
with source.open(encoding="utf-8") as stream:
|
||||
return validate_configuration(json.load(stream))
|
||||
return validate_configuration(json.load(stream), task_id)
|
||||
|
||||
|
||||
def _load_task_config(task_id: str, path: str | None, seed: int) -> dict | None:
|
||||
if path is None:
|
||||
return validate_task_config(task_id, {}, seed)
|
||||
source = Path(path).expanduser().resolve(strict=True)
|
||||
if source.stat().st_size > 128 * 1024:
|
||||
raise ValueError("Task configuration is larger than 128 KiB")
|
||||
payload = json.loads(source.read_text(encoding="utf-8"))
|
||||
allowed = {"terrainPreset", "terrainParams", "sensorCfg", "seed", "customTerrainBoxes"}
|
||||
if not isinstance(payload, dict) or payload.keys() - allowed:
|
||||
raise ValueError("Unknown task configuration fields")
|
||||
if (
|
||||
isinstance(payload.get("seed"), bool)
|
||||
or not isinstance(payload.get("seed"), int)
|
||||
or payload.get("seed") != seed
|
||||
):
|
||||
raise ValueError("Task configuration seed must equal the agent seed")
|
||||
if payload.get("sensorCfg") is None:
|
||||
payload.pop("sensorCfg", None)
|
||||
return validate_task_config(task_id, payload, seed)
|
||||
|
||||
|
||||
def _configure_task_and_rewards(task_id: str, cfg: TrainConfig):
|
||||
custom = _load_task_config(task_id, cfg.task_config, cfg.agent.seed)
|
||||
if custom is not None:
|
||||
from src.tasks.obstacle_avoidance.env_cfg import apply_obstacle_configuration
|
||||
from src.tasks.obstacle_avoidance.terrain import apply_terrain_configuration
|
||||
if task_id == OBSTACLE_TASK:
|
||||
apply_obstacle_configuration(cfg.env, custom)
|
||||
else:
|
||||
apply_terrain_configuration(cfg.env, custom)
|
||||
deployment = deployment_metadata(task_id, custom, cfg.agent.seed)
|
||||
reward_config = _load_reward_config(cfg.reward_config, cfg.reward_config_json, task_id)
|
||||
if reward_config is not None:
|
||||
apply_reward_configuration(cfg.env, reward_config, task_id)
|
||||
if task_id == OBSTACLE_TASK:
|
||||
deployment["navigation"]["speed"] = reward_config["params"]["target_velocity"]
|
||||
deployment["sensorCfg"]["avoidanceWeight"] = reward_config["weights"]["avoidance_weight"]
|
||||
|
||||
return deployment, reward_config
|
||||
|
||||
|
||||
def _load_pretrained(cfg: TrainConfig):
|
||||
if cfg.pretrained_checkpoint is None:
|
||||
if (cfg.pretrained_onnx is not None or cfg.pretrained_allowed_roots
|
||||
or cfg.pretrained_source_id or cfg.pretrained_upload_manifest):
|
||||
raise ValueError("Pretrained options require --pretrained-checkpoint (.pt), not ONNX alone")
|
||||
return None
|
||||
if cfg.resume_checkpoint is not None or cfg.agent.resume:
|
||||
raise ValueError("Pretrained warm-start and resume are mutually exclusive")
|
||||
from pretrained import read_pretrained_source
|
||||
|
||||
options = {"onnx_path": cfg.pretrained_onnx}
|
||||
if cfg.pretrained_upload_manifest is not None:
|
||||
if cfg.pretrained_onnx is not None:
|
||||
raise ValueError("Uploaded actor cannot use adjacent ONNX sidecars")
|
||||
from pretrained_upload import read_uploaded_source
|
||||
read_pretrained_source = read_uploaded_source
|
||||
options = {"manifest_path": cfg.pretrained_upload_manifest}
|
||||
source = read_pretrained_source(
|
||||
cfg.pretrained_checkpoint,
|
||||
allowed_roots=cfg.pretrained_allowed_roots,
|
||||
**options,
|
||||
target_env=asdict(cfg.env),
|
||||
target_agent=asdict(cfg.agent),
|
||||
)
|
||||
if cfg.pretrained_source_id is not None and source.manifest["source_id"] != cfg.pretrained_source_id:
|
||||
raise ValueError("基础策略SHA身份变化,拒绝初始化")
|
||||
return source
|
||||
|
||||
|
||||
def run_train(task_id: str, cfg: TrainConfig, log_dir: Path) -> None:
|
||||
reward_config = _load_reward_config(cfg.reward_config, cfg.reward_config_json)
|
||||
if reward_config is not None:
|
||||
apply_reward_configuration(cfg.env, reward_config)
|
||||
|
||||
deployment, reward_config = _configure_task_and_rewards(task_id, cfg)
|
||||
# Validate source identity/semantics before allocating a simulation or optimizer.
|
||||
pretrained = _load_pretrained(cfg)
|
||||
cuda_visible = os.environ.get("CUDA_VISIBLE_DEVICES", "")
|
||||
if cuda_visible == "":
|
||||
device = "cpu"
|
||||
@@ -135,6 +210,17 @@ def run_train(task_id: str, cfg: TrainConfig, log_dir: Path) -> None:
|
||||
cfg=cfg.env, device=device, render_mode="rgb_array" if cfg.video else None
|
||||
)
|
||||
|
||||
if pretrained is not None:
|
||||
from pretrained import validate_runtime_contract
|
||||
|
||||
try:
|
||||
validate_runtime_contract(env)
|
||||
except Exception:
|
||||
env.close()
|
||||
raise
|
||||
|
||||
# The ONNX runner exports this exact snapshot beside and inside policy.onnx.
|
||||
env.platform_deployment = deployment
|
||||
log_root_path = log_dir.parent # Go up from specific run dir to experiment dir.
|
||||
|
||||
resume_path: Path | None = None
|
||||
@@ -171,9 +257,30 @@ def run_train(task_id: str, cfg: TrainConfig, log_dir: Path) -> None:
|
||||
runner = runner_cls(env, agent_cfg, str(log_dir), device, **runner_kwargs)
|
||||
|
||||
runner.add_git_repo_to_log(__file__)
|
||||
if pretrained is not None:
|
||||
from pretrained import initialize_runner
|
||||
|
||||
initialization = initialize_runner(runner, pretrained)
|
||||
env.unwrapped.platform_initialization = initialization
|
||||
if rank == 0:
|
||||
(log_dir / "initialization.json").write_text(
|
||||
json.dumps(initialization, ensure_ascii=False, indent=2) + "\n", encoding="utf-8"
|
||||
)
|
||||
print(f"[INFO] Warm-start actor from source {initialization['source_id']}; fresh critic/optimizer, iteration=0")
|
||||
if resume_path is not None:
|
||||
print(f"[INFO]: Loading model checkpoint from: {resume_path}")
|
||||
runner.load(str(resume_path), map_location=device)
|
||||
# Promotion copies the origin manifest, never the original actor weights.
|
||||
origin_path = resume_path.parent / "initialization.json"
|
||||
if origin_path.is_file():
|
||||
if origin_path.stat().st_size > 64 * 1024:
|
||||
raise ValueError("Initialization provenance exceeds 64 KiB")
|
||||
origin = json.loads(origin_path.read_text(encoding="utf-8"))
|
||||
if not isinstance(origin, dict) or origin.get("mode") != "pretrained-warm-start":
|
||||
raise ValueError("Invalid initialization provenance")
|
||||
env.unwrapped.platform_initialization = origin
|
||||
if rank == 0:
|
||||
(log_dir / "initialization.json").write_text(json.dumps(origin, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
|
||||
if explicit_resume:
|
||||
# RSL-RL stores the last completed zero-based iteration and otherwise
|
||||
# repeats it after load. Explicit tuning promotion uses an absolute target.
|
||||
@@ -260,7 +367,7 @@ def main():
|
||||
# Parse first argument to choose the task.
|
||||
# Import tasks to populate the registry.
|
||||
import mjlab.tasks # noqa: F401
|
||||
import src.tasks
|
||||
import src.tasks # noqa: F401
|
||||
|
||||
all_tasks = list_tasks()
|
||||
chosen_task, remaining_args = tyro.cli(
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
"""Service-owned CPU validator; JSON stdin is never accepted directly from HTTP."""
|
||||
|
||||
import json
|
||||
import sys
|
||||
from dataclasses import asdict
|
||||
from pathlib import Path
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[1]
|
||||
for root in (ROOT, ROOT.parent):
|
||||
sys.path.insert(0, str(root))
|
||||
|
||||
|
||||
def main():
|
||||
from pretrained import read_pretrained_source
|
||||
from task_config import OBSTACLE_TASK
|
||||
from src.tasks.velocity.config.go2.env_cfgs import unitree_go2_flat_env_cfg
|
||||
from src.tasks.velocity.config.go2.rl_cfg import unitree_go2_ppo_runner_cfg
|
||||
|
||||
payload = json.loads(sys.stdin.read(128 * 1024))
|
||||
task_id = payload["taskId"]
|
||||
if task_id == "Unitree-Go2-Flat":
|
||||
cfg = unitree_go2_flat_env_cfg()
|
||||
elif task_id == OBSTACLE_TASK:
|
||||
from src.tasks.obstacle_avoidance.env_cfg import unitree_go2_obstacle_env_cfg, apply_obstacle_configuration
|
||||
cfg = unitree_go2_obstacle_env_cfg()
|
||||
if payload.get("taskConfig") is not None:
|
||||
apply_obstacle_configuration(cfg, payload["taskConfig"])
|
||||
else:
|
||||
raise ValueError("基础策略不兼容该任务;支持Flat47与Obstacle81/97,不支持Rough")
|
||||
directory = Path(payload["directory"])
|
||||
options = {}
|
||||
if payload.get("uploaded"):
|
||||
from pretrained_upload import read_uploaded_source
|
||||
read_pretrained_source = read_uploaded_source
|
||||
options["manifest_path"] = directory / "upload.json"
|
||||
source = read_pretrained_source(
|
||||
directory / payload["checkpoint"], allowed_roots=[directory], target_env=asdict(cfg),
|
||||
target_agent=asdict(unitree_go2_ppo_runner_cfg()), **options,
|
||||
)
|
||||
print(json.dumps(source.manifest))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
main()
|
||||
except Exception as error:
|
||||
print(f"基础策略不兼容或缺少配套文件/依赖:{error}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
@@ -0,0 +1,43 @@
|
||||
"""Resource-limited CPU single-file decoder; no network or neighboring model lookup."""
|
||||
|
||||
import contextlib
|
||||
import json
|
||||
import os
|
||||
import resource
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
# Apply limits before importing tensor/protobuf runtimes. Thread env also set by parent.
|
||||
resource.setrlimit(resource.RLIMIT_CPU, (40, 40))
|
||||
resource.setrlimit(resource.RLIMIT_AS, (8 * 1024**3, 8 * 1024**3))
|
||||
resource.setrlimit(resource.RLIMIT_FSIZE, (16 * 1024**2, 16 * 1024**2))
|
||||
resource.setrlimit(resource.RLIMIT_NOFILE, (64, 64))
|
||||
resource.setrlimit(resource.RLIMIT_CORE, (0, 0))
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = ""
|
||||
ROOT = Path(__file__).resolve().parents[1]
|
||||
for root in (ROOT, ROOT.parent):
|
||||
sys.path.insert(0, str(root))
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
payload = json.loads(sys.stdin.read(4096))
|
||||
with contextlib.redirect_stdout(sys.stderr):
|
||||
import torch
|
||||
from pretrained_upload import import_upload
|
||||
|
||||
torch.set_num_threads(1)
|
||||
manifest = import_upload(
|
||||
payload["path"], payload["format"], payload["template"], payload["directory"]
|
||||
)
|
||||
print(json.dumps(manifest))
|
||||
except Exception as error:
|
||||
# Do not echo untrusted pickle/tensor names or absolute local paths.
|
||||
from pretrained import PretrainedError
|
||||
|
||||
message = (
|
||||
str(error)
|
||||
if isinstance(error, PretrainedError)
|
||||
else "文件解析失败;请提供受支持的自包含Go2 legacy47 actor"
|
||||
)
|
||||
print(json.dumps({"error": message[:500]}))
|
||||
sys.exit(1)
|
||||
@@ -0,0 +1,17 @@
|
||||
from mjlab.tasks.registry import register_mjlab_task
|
||||
from task_config import OBSTACLE_TASK
|
||||
|
||||
from src.tasks.velocity.config.go2.rl_cfg import unitree_go2_ppo_runner_cfg
|
||||
from src.tasks.velocity.rl import VelocityOnPolicyRunner
|
||||
|
||||
from .env_cfg import unitree_go2_obstacle_env_cfg
|
||||
|
||||
_rl_cfg = unitree_go2_ppo_runner_cfg()
|
||||
_rl_cfg.experiment_name = "go2_obstacle_avoidance"
|
||||
register_mjlab_task(
|
||||
task_id=OBSTACLE_TASK,
|
||||
env_cfg=unitree_go2_obstacle_env_cfg(),
|
||||
play_env_cfg=unitree_go2_obstacle_env_cfg(play=True),
|
||||
rl_cfg=_rl_cfg,
|
||||
runner_cls=VelocityOnPolicyRunner,
|
||||
)
|
||||
@@ -0,0 +1,89 @@
|
||||
"""81/97-to-12 forward-ray navigation, retaining the Go2 50 Hz position action contract."""
|
||||
|
||||
from mjlab.managers import ObservationTermCfg, RewardTermCfg, TerminationTermCfg
|
||||
from mjlab.managers.event_manager import EventTermCfg
|
||||
from mjlab.sensor import ObjRef, RayCastSensorCfg
|
||||
from task_config import (
|
||||
OBSTACLE_TASK,
|
||||
build_terrain_layout,
|
||||
navigation_candidates,
|
||||
validate_task_config,
|
||||
)
|
||||
|
||||
from src.tasks.velocity.config.go2.env_cfgs import unitree_go2_flat_env_cfg
|
||||
|
||||
from . import mdp
|
||||
from .terrain import apply_terrain_configuration
|
||||
|
||||
|
||||
def apply_obstacle_configuration(cfg, custom, randomize_navigation=True):
|
||||
apply_terrain_configuration(cfg, custom)
|
||||
sensor = custom["sensorCfg"]
|
||||
cfg.scene.sensors = tuple(s for s in cfg.scene.sensors if s.name != "forward_scan") + (
|
||||
RayCastSensorCfg(
|
||||
name="forward_scan",
|
||||
frame=ObjRef(type="body", name="base_link", entity="robot"),
|
||||
pattern=mdp.ForwardFanPatternCfg(fov=sensor["fov"], sensor_mode=sensor["sensorMode"]),
|
||||
ray_alignment="base",
|
||||
max_distance=sensor["maxDistance"],
|
||||
include_geom_groups=(0,),
|
||||
exclude_parent_body=True,
|
||||
),
|
||||
)
|
||||
for group in ("actor", "critic"):
|
||||
cfg.observations[group].terms["forward_depth"] = ObservationTermCfg(
|
||||
func=mdp.forward_depth, params={"max_distance": sensor["maxDistance"]}
|
||||
)
|
||||
cfg.observations[group].terms["target_error"] = ObservationTermCfg(func=mdp.target_error)
|
||||
layout = build_terrain_layout(custom)
|
||||
size = layout["size"]
|
||||
candidates = navigation_candidates(layout) if randomize_navigation else None
|
||||
cfg.commands["twist"] = mdp.NavigationCommandCfg(
|
||||
resampling_time_range=(1e9, 1e9),
|
||||
goal_offset=tuple(layout["target"][i] - layout["spawn"][i] for i in range(2)),
|
||||
distance_scale=size,
|
||||
**(
|
||||
{
|
||||
"navigation_points": tuple(tuple(point) for point in candidates["points"]),
|
||||
"component_starts": tuple(candidates["componentStarts"]),
|
||||
"component_counts": tuple(candidates["componentCounts"]),
|
||||
"fallback_pairs": tuple(
|
||||
(tuple(pair[0]), tuple(pair[1])) for pair in candidates["fallbackPairs"]
|
||||
),
|
||||
"min_goal_distance": candidates["minDistance"],
|
||||
}
|
||||
if candidates
|
||||
else {}
|
||||
),
|
||||
)
|
||||
if randomize_navigation:
|
||||
cfg.events["reset_base"] = EventTermCfg(func=mdp.reset_navigation_episode, mode="reset")
|
||||
cfg.rewards["goal_progress"] = RewardTermCfg(func=mdp.goal_progress, weight=2.0)
|
||||
cfg.rewards["obstacle_proximity"] = RewardTermCfg(
|
||||
func=mdp.obstacle_proximity,
|
||||
weight=-sensor["avoidanceWeight"],
|
||||
params={
|
||||
"max_distance": sensor["maxDistance"],
|
||||
"safety_distance": sensor["safetyDistance"],
|
||||
**({"floor_size": size} if sensor["sensorMode"] == "multi_ring_raycast" else {}),
|
||||
},
|
||||
)
|
||||
cfg.terminations["outside_map"] = TerminationTermCfg(
|
||||
func=mdp.outside_map, params={"size": size}
|
||||
)
|
||||
|
||||
|
||||
def unitree_go2_obstacle_env_cfg(play=False):
|
||||
cfg = unitree_go2_flat_env_cfg(play=False)
|
||||
cfg.sim.nconmax = 128
|
||||
cfg.sim.njmax = 1500
|
||||
cfg.sim.contact_sensor_maxmatch = 500
|
||||
cfg.curriculum = {}
|
||||
for event in ("push_robot", "encoder_bias", "base_com"):
|
||||
cfg.events.pop(event, None)
|
||||
cfg.observations["actor"].enable_corruption = False
|
||||
custom = validate_task_config(OBSTACLE_TASK, {}, 42)
|
||||
apply_obstacle_configuration(cfg, custom)
|
||||
if play:
|
||||
cfg.episode_length_s = 20.0
|
||||
return cfg
|
||||
@@ -0,0 +1,231 @@
|
||||
"""Go2 navigation observations: body-frame +X forward, +Y left, +Z up."""
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
from mjlab.managers.command_manager import CommandTerm, CommandTermCfg
|
||||
from mjlab.sensor import GridPatternCfg
|
||||
|
||||
|
||||
@dataclass
|
||||
class ForwardFanPatternCfg(GridPatternCfg):
|
||||
"""Inclusive yaw, layer-major pitch; base-aligned with one shared offset."""
|
||||
|
||||
fov: float = 90.0
|
||||
sensor_mode: str = "single_ring_raycast"
|
||||
|
||||
def generate_rays(self, mj_model, device):
|
||||
from task_config import sensor_pattern
|
||||
|
||||
pattern = sensor_pattern(self.sensor_mode, self.fov)
|
||||
yaw = torch.tensor(pattern["yawAngles"], device=device) * torch.pi / 180
|
||||
pitch = torch.tensor(pattern["pitchAngles"], device=device) * torch.pi / 180
|
||||
p, y = torch.meshgrid(pitch, yaw, indexing="ij")
|
||||
directions = torch.stack((p.cos() * y.cos(), p.cos() * y.sin(), p.sin()), dim=-1).reshape(
|
||||
-1, 3
|
||||
)
|
||||
offsets = torch.tensor([0.3, 0, 0.05], device=device).repeat(pattern["rayCount"], 1)
|
||||
return offsets, directions
|
||||
|
||||
|
||||
@dataclass(kw_only=True)
|
||||
class NavigationCommandCfg(CommandTermCfg):
|
||||
goal_offset: tuple[float, float] = (10.0, 0.0)
|
||||
distance_scale: float = 12.0
|
||||
speed: float = 0.6
|
||||
arrival_radius: float = 0.5
|
||||
navigation_points: tuple[tuple[float, float], ...] = ()
|
||||
component_starts: tuple[int, ...] = ()
|
||||
component_counts: tuple[int, ...] = ()
|
||||
fallback_pairs: tuple[tuple[tuple[float, float], tuple[float, float]], ...] = ()
|
||||
min_goal_distance: float = 2.0
|
||||
|
||||
def build(self, env):
|
||||
return NavigationCommand(self, env)
|
||||
|
||||
|
||||
class NavigationCommand(CommandTerm):
|
||||
cfg: NavigationCommandCfg
|
||||
|
||||
def __init__(self, cfg, env):
|
||||
super().__init__(cfg, env)
|
||||
self.metrics["distance_to_goal"] = torch.zeros(self.num_envs, device=self.device)
|
||||
self.goals_w = self._env.scene.env_origins[:, :2] + torch.tensor(
|
||||
self.cfg.goal_offset, device=self.device
|
||||
)
|
||||
self._points = (
|
||||
torch.tensor(self.cfg.navigation_points, dtype=torch.float32, device=self.device)
|
||||
if self.cfg.navigation_points
|
||||
else None
|
||||
)
|
||||
self._component_starts = torch.tensor(
|
||||
self.cfg.component_starts, dtype=torch.long, device=self.device
|
||||
)
|
||||
self._component_counts = torch.tensor(
|
||||
self.cfg.component_counts, dtype=torch.long, device=self.device
|
||||
)
|
||||
self._fallback_pairs = (
|
||||
torch.tensor(self.cfg.fallback_pairs, dtype=torch.float32, device=self.device)
|
||||
if self.cfg.fallback_pairs
|
||||
else None
|
||||
)
|
||||
|
||||
def errors(self):
|
||||
robot = self._env.scene["robot"]
|
||||
delta = self.goals_w - robot.data.root_link_pos_w[:, :2]
|
||||
q = robot.data.root_link_quat_w
|
||||
yaw = torch.atan2(
|
||||
2 * (q[:, 0] * q[:, 3] + q[:, 1] * q[:, 2]),
|
||||
1 - 2 * (q[:, 2].square() + q[:, 3].square()),
|
||||
)
|
||||
heading = torch.atan2(delta[:, 1], delta[:, 0]) - yaw
|
||||
heading = torch.atan2(heading.sin(), heading.cos())
|
||||
distance = delta.norm(dim=1)
|
||||
heading = torch.where(distance < self.cfg.arrival_radius, 0.0, heading)
|
||||
return heading, distance, delta
|
||||
|
||||
@property
|
||||
def command(self):
|
||||
# Compute from current physics, not a one-step-old command manager cache.
|
||||
heading, distance, _ = self.errors()
|
||||
moving = distance >= self.cfg.arrival_radius
|
||||
vx = self.cfg.speed * heading.cos().clamp(min=0)
|
||||
return torch.stack(
|
||||
(
|
||||
torch.where(moving, vx, 0.0),
|
||||
torch.zeros_like(vx),
|
||||
torch.where(moving, heading.clamp(-1, 1), 0.0),
|
||||
),
|
||||
dim=1,
|
||||
)
|
||||
|
||||
def _update_metrics(self):
|
||||
self.metrics["distance_to_goal"][:] = self.errors()[1]
|
||||
|
||||
def sample_episode(self, env_ids):
|
||||
"""Randomize a reachable free-space pair and the initial heading."""
|
||||
if self._points is None or len(env_ids) == 0:
|
||||
return
|
||||
count = len(env_ids)
|
||||
component = torch.randint(len(self._component_starts), (count,), device=self.device)
|
||||
starts = self._component_starts[component]
|
||||
counts = self._component_counts[component]
|
||||
|
||||
def indices():
|
||||
return starts + (torch.rand(count, device=self.device) * counts).long()
|
||||
|
||||
spawn = self._points[indices()]
|
||||
target = self._points[indices()]
|
||||
valid = (target - spawn).norm(dim=1) >= self.cfg.min_goal_distance
|
||||
# Keep this branch-free on CUDA; unresolved samples use a guaranteed component pair.
|
||||
for _ in range(16):
|
||||
candidate = self._points[indices()]
|
||||
target = torch.where(valid[:, None], target, candidate)
|
||||
valid = (target - spawn).norm(dim=1) >= self.cfg.min_goal_distance
|
||||
fallback = self._fallback_pairs[component]
|
||||
spawn = torch.where(valid[:, None], spawn, fallback[:, 0])
|
||||
target = torch.where(valid[:, None], target, fallback[:, 1])
|
||||
|
||||
robot = self._env.scene["robot"]
|
||||
default = robot.data.default_root_state[env_ids].clone()
|
||||
position = torch.cat((spawn, default[:, 2:3]), dim=1)
|
||||
yaw = torch.rand(count, device=self.device) * (2 * torch.pi) - torch.pi
|
||||
half = yaw / 2
|
||||
quaternion = torch.stack(
|
||||
(half.cos(), torch.zeros_like(half), torch.zeros_like(half), half.sin()), dim=1
|
||||
)
|
||||
robot.write_root_link_pose_to_sim(torch.cat((position, quaternion), dim=1), env_ids=env_ids)
|
||||
robot.write_root_link_velocity_to_sim(torch.zeros_like(default[:, 7:13]), env_ids=env_ids)
|
||||
self.goals_w[env_ids] = target
|
||||
|
||||
def _resample_command(self, env_ids):
|
||||
# Episode pairs are sampled by the reset event before manager resets.
|
||||
self.metrics["distance_to_goal"][env_ids] = self.errors()[1][env_ids]
|
||||
|
||||
def _update_command(self):
|
||||
# No cached command: the property is evaluated after every physics update.
|
||||
self._update_metrics()
|
||||
|
||||
|
||||
def reset_navigation_episode(env, env_ids):
|
||||
if env_ids is None:
|
||||
env_ids = torch.arange(env.num_envs, device=env.device, dtype=torch.int)
|
||||
env.command_manager.get_term("twist").sample_episode(env_ids)
|
||||
|
||||
|
||||
def forward_depth(env, sensor_name="forward_scan", max_distance=4.0):
|
||||
distances = env.scene[sensor_name].data.distances
|
||||
distances = torch.where(distances < 0, max_distance, distances)
|
||||
return distances.clamp(0, max_distance) / max_distance
|
||||
|
||||
|
||||
def target_error(env):
|
||||
command = env.command_manager.get_term("twist")
|
||||
heading, distance, _ = command.errors()
|
||||
return torch.stack(
|
||||
(heading / torch.pi, (distance / command.cfg.distance_scale).clamp(0, 1)), dim=1
|
||||
)
|
||||
|
||||
|
||||
def standard_floor_id(model, size):
|
||||
"""Fail closed: match the compiled first layout box, not a name alone."""
|
||||
import mujoco
|
||||
import numpy as np
|
||||
|
||||
data = mujoco.MjData(model)
|
||||
mujoco.mj_forward(model, data)
|
||||
ids = [
|
||||
i
|
||||
for i in range(model.ngeom)
|
||||
if model.geom_type[i] == 6
|
||||
and model.body_weldid[model.geom_bodyid[i]] == 0
|
||||
and model.geom_group[i] == 0
|
||||
and np.allclose(data.geom_xpos[i], [0, 0, -0.1], atol=1e-7, rtol=0)
|
||||
and np.allclose(model.geom_size[i], [size / 2, size / 2, 0.1], atol=1e-7, rtol=0)
|
||||
and np.allclose(data.geom_xmat[i], [1, 0, 0, 0, 1, 0, 0, 0, 1], atol=1e-7, rtol=0)
|
||||
]
|
||||
if len(ids) != 1 or model.geom(ids[0]).name != "terrain_0":
|
||||
raise ValueError("multi-ring reward requires one verified static standard floor")
|
||||
return ids[0]
|
||||
|
||||
|
||||
def floor_top_hits(data, size):
|
||||
# Warp ray hits use the same world-centered coordinates in each independent env.
|
||||
# Float32 tolerance: 10 micrometres; never erase 0.05m low obstacles.
|
||||
hit = data.hit_pos_w
|
||||
normal = data.normals_w
|
||||
return (
|
||||
(data.distances >= 0)
|
||||
& torch.isfinite(data.distances)
|
||||
& (hit[..., :2].abs() <= size / 2 + 1e-5).all(dim=-1)
|
||||
& (hit[..., 2].abs() <= 1e-5)
|
||||
& (normal[..., :2].abs() <= 1e-5).all(dim=-1)
|
||||
& ((normal[..., 2] - 1).abs() <= 1e-5)
|
||||
)
|
||||
|
||||
|
||||
def obstacle_proximity(env, max_distance=4.0, safety_distance=0.5, floor_size=None):
|
||||
depths = forward_depth(env, max_distance=max_distance)
|
||||
if floor_size is not None:
|
||||
# Validate once per environment/model. Observation data remains untouched.
|
||||
if not hasattr(env, "_multi_ring_floor_id"):
|
||||
env._multi_ring_floor_id = standard_floor_id(env.sim.mj_model, floor_size)
|
||||
depths = torch.where(
|
||||
floor_top_hits(env.scene["forward_scan"].data, floor_size), 1.0, depths
|
||||
)
|
||||
closest = depths.amin(dim=1) * max_distance
|
||||
return ((safety_distance - closest).clamp(min=0) / safety_distance).square()
|
||||
|
||||
|
||||
def goal_progress(env):
|
||||
command = env.command_manager.get_term("twist")
|
||||
_, distance, delta = command.errors()
|
||||
velocity = env.scene["robot"].data.root_link_lin_vel_w[:, :2]
|
||||
progress = (velocity * delta / distance.clamp(min=1e-6).unsqueeze(1)).sum(dim=1)
|
||||
return torch.where(distance >= command.cfg.arrival_radius, progress, 0.0)
|
||||
|
||||
|
||||
def outside_map(env, size=12.0):
|
||||
# Every Warp environment is an independent copy of the same world-centered map.
|
||||
position = env.scene["robot"].data.root_link_pos_w[:, :2]
|
||||
return (position.abs() > size / 2 - 0.3).any(dim=1)
|
||||
@@ -0,0 +1,70 @@
|
||||
"""Install the exact bounded boxes-v1 deployment layout in mjlab's generator."""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import mujoco
|
||||
import numpy as np
|
||||
from mjlab.envs import mdp
|
||||
from mjlab.managers.event_manager import EventTermCfg
|
||||
from mjlab.terrains import TerrainEntityCfg, TerrainGeneratorCfg
|
||||
from mjlab.terrains.terrain_generator import SubTerrainCfg, TerrainGeometry, TerrainOutput
|
||||
from task_config import build_terrain_layout
|
||||
|
||||
|
||||
@dataclass
|
||||
class DeploymentTerrainCfg(SubTerrainCfg):
|
||||
layout: dict = field(default_factory=dict)
|
||||
|
||||
def function(self, difficulty, spec, rng):
|
||||
# Generator translates patch-local corner coordinates by (-size/2,-size/2).
|
||||
half = self.layout["size"] / 2
|
||||
translation = np.array([half, half, 0])
|
||||
body = spec.body("terrain")
|
||||
geometries = []
|
||||
for box in self.layout["boxes"]:
|
||||
geom = body.add_geom(
|
||||
type=mujoco.mjtGeom.mjGEOM_BOX,
|
||||
pos=np.array(box["pos"]) + translation,
|
||||
size=box["size"],
|
||||
group=0,
|
||||
contype=1,
|
||||
conaffinity=1,
|
||||
friction=[self.layout["friction"], 0.005, 0.0001],
|
||||
priority=1,
|
||||
condim=3,
|
||||
)
|
||||
geometries.append(TerrainGeometry(geom=geom))
|
||||
spawn = np.array(self.layout["spawn"])
|
||||
spawn[2] = 0 # Robot init_state contributes the 0.32 m body height.
|
||||
return TerrainOutput(origin=spawn + translation, geometries=geometries)
|
||||
|
||||
|
||||
def apply_terrain_configuration(cfg, custom):
|
||||
layout = build_terrain_layout(custom)
|
||||
size = layout["size"]
|
||||
cfg.scene.entities["robot"].init_state.pos = (0.0, 0.0, layout["spawn"][2])
|
||||
cfg.scene.entities["robot"].init_state.rot = tuple(layout["spawnQuaternion"])
|
||||
cfg.scene.terrain = TerrainEntityCfg(
|
||||
terrain_type="generator",
|
||||
max_init_terrain_level=0,
|
||||
terrain_generator=TerrainGeneratorCfg(
|
||||
seed=custom["seed"],
|
||||
size=(size, size),
|
||||
num_rows=1,
|
||||
num_cols=1,
|
||||
border_width=0,
|
||||
curriculum=False,
|
||||
color_scheme="none",
|
||||
sub_terrains={"deployment": DeploymentTerrainCfg(layout=layout)},
|
||||
),
|
||||
)
|
||||
cfg.curriculum.pop("terrain_levels", None)
|
||||
cfg.events.pop("randomize_terrain", None)
|
||||
# Install the deterministic reset first; the navigation task may replace it with
|
||||
# collision-free random point-goal sampling after the command term is configured.
|
||||
cfg.events["reset_base"] = EventTermCfg(
|
||||
func=mdp.reset_root_state_uniform,
|
||||
mode="reset",
|
||||
params={"pose_range": {}, "velocity_range": {}},
|
||||
)
|
||||
cfg.events["foot_friction"].params["ranges"] = (layout["friction"], layout["friction"])
|
||||
@@ -1,4 +1,6 @@
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import wandb
|
||||
|
||||
@@ -23,6 +25,19 @@ class VelocityOnPolicyRunner(MjlabOnPolicyRunner):
|
||||
) # type: ignore[assignment]
|
||||
onnx_path = os.path.join(policy_path, filename)
|
||||
metadata = get_base_metadata(self.env.unwrapped, run_name)
|
||||
deployment = getattr(self.env.unwrapped, "platform_deployment", None)
|
||||
if deployment is not None:
|
||||
metadata["platform_deployment"] = json.dumps(deployment, separators=(",", ":"))
|
||||
Path(policy_path, "deployment.json").write_text(
|
||||
json.dumps(deployment, ensure_ascii=False, indent=2) + "\n", encoding="utf-8"
|
||||
)
|
||||
initialization = getattr(self.env.unwrapped, "platform_initialization", None)
|
||||
if initialization is not None:
|
||||
# Only content identity and protocol, never local source paths or arbitrary infos.
|
||||
from pretrained import public_initialization_metadata
|
||||
metadata["pretrained_initialization"] = json.dumps(
|
||||
public_initialization_metadata(initialization), separators=(",", ":")
|
||||
)
|
||||
attach_metadata_to_onnx(onnx_path, metadata)
|
||||
if self.logger.logger_type in ["wandb"]:
|
||||
wandb.save(policy_path + filename, base_path=os.path.dirname(policy_path))
|
||||
|
||||
+150
-12
@@ -25,14 +25,23 @@ from pathlib import Path
|
||||
from typing import Any
|
||||
from urllib.parse import parse_qs, unquote, urlsplit
|
||||
|
||||
from pretrained_sources import PretrainedSources, SourceError
|
||||
from task_config import (
|
||||
OBSTACLE_TASK,
|
||||
TaskConfigError,
|
||||
deployment_metadata,
|
||||
task_metadata,
|
||||
validate_task_config,
|
||||
)
|
||||
from tuning.manager import TuningError, TuningManager
|
||||
from tuning.process import GpuLease, ResourceBusyError
|
||||
from tuning.schema import RewardConfigError
|
||||
from tuning.schema import RewardConfigError, validate_configuration
|
||||
from tuning.scoring import EvaluationError
|
||||
|
||||
VERSION = "0.4.0"
|
||||
# 浏览器当前 ONNX 运行时只实现 Go2 的 47→12 部署契约;其他任务须由服务启动参数显式放行。
|
||||
DEFAULT_TASKS = ("Unitree-Go2-Flat",)
|
||||
MAX_REQUEST_BYTES = 128 * 1024 # Bounded full boxes-v1 payload (<=257 boxes).
|
||||
# Rough 可以训练,但其高度扫描 actor 不允许冒充浏览器 Flat 部署。
|
||||
DEFAULT_TASKS = ("Unitree-Go2-Flat", "Unitree-Go2-Rough", OBSTACLE_TASK)
|
||||
ACTIVE_STATES = {"queued", "running"}
|
||||
MAX_JOBS = 20
|
||||
ANSI_ESCAPE = re.compile(r"\x1b\[[0-?]*[ -/]*[@-~]")
|
||||
@@ -71,6 +80,12 @@ class TrainingConfig:
|
||||
gpu_ids: list[int]
|
||||
wandb_mode: str
|
||||
reward_config: dict[str, Any] | None = None
|
||||
terrain_preset: str | None = None
|
||||
terrain_params: dict[str, Any] = field(default_factory=dict)
|
||||
sensor_cfg: dict[str, Any] | None = None
|
||||
task_config: dict[str, Any] | None = None
|
||||
deployment: dict[str, Any] = field(default_factory=dict)
|
||||
pretrained: dict[str, Any] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -96,6 +111,8 @@ class TrainingJob:
|
||||
"id": self.id,
|
||||
"state": self.state,
|
||||
"taskId": self.config.task_id,
|
||||
"deployment": self.config.deployment,
|
||||
"pretrained": self.config.pretrained,
|
||||
"createdAt": self.created_at,
|
||||
"startedAt": self.started_at,
|
||||
"endedAt": self.ended_at,
|
||||
@@ -117,6 +134,7 @@ class TrainingManager:
|
||||
tasks: tuple[str, ...],
|
||||
check_environment: bool = True,
|
||||
lease: GpuLease | None = None,
|
||||
sources: PretrainedSources | None = None,
|
||||
):
|
||||
self.trainer_root = trainer_root.expanduser().resolve()
|
||||
self.python = str(Path(python).expanduser()) if os.sep in python else python
|
||||
@@ -127,6 +145,7 @@ class TrainingManager:
|
||||
self._environment_error: str | None | bool = False
|
||||
self.lease = lease or GpuLease()
|
||||
self.preset_resolver: Any = None
|
||||
self.sources = sources
|
||||
|
||||
def readiness_error(self) -> str | None:
|
||||
if not self.trainer_root.is_dir():
|
||||
@@ -172,6 +191,13 @@ class TrainingManager:
|
||||
"trainerRoot": str(self.trainer_root),
|
||||
"python": self.python,
|
||||
"tasks": list(self.tasks),
|
||||
"pretrainedSources": self.sources.catalog() if self.sources else [],
|
||||
"pretrainedUpload": {
|
||||
"enabled": self.sources is not None, "templateId": "go2-legacy47-v1",
|
||||
"formats": {"pt": 256 * 1024**2, "onnx": 64 * 1024**2},
|
||||
"endpoint": "/api/training/pretrained-sources/upload",
|
||||
},
|
||||
"taskMetadata": task_metadata(self.tasks),
|
||||
"activeJobId": self.active_job_id(),
|
||||
"error": error,
|
||||
}
|
||||
@@ -179,6 +205,25 @@ class TrainingManager:
|
||||
def parse_config(self, payload: Any) -> TrainingConfig:
|
||||
if not isinstance(payload, dict):
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, "请求体必须是 JSON 对象")
|
||||
allowed = {
|
||||
"taskId",
|
||||
"numEnvs",
|
||||
"maxIterations",
|
||||
"seed",
|
||||
"runName",
|
||||
"device",
|
||||
"gpuIds",
|
||||
"wandbMode",
|
||||
"rewardPresetId",
|
||||
"pretrainedSourceId",
|
||||
"terrainPreset",
|
||||
"terrainParams",
|
||||
"sensorCfg",
|
||||
"sensorType",
|
||||
"customTerrainBoxes",
|
||||
}
|
||||
if payload.keys() - allowed:
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, "请求包含未知字段(不接受配置路径/MJCF)")
|
||||
task_id = payload.get("taskId")
|
||||
if task_id not in self.tasks:
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, f"不允许的训练任务:{task_id}")
|
||||
@@ -216,19 +261,46 @@ class TrainingManager:
|
||||
preset_id = payload.get("rewardPresetId")
|
||||
reward_config = None
|
||||
if preset_id is not None:
|
||||
if task_id != "Unitree-Go2-Flat":
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, "自动调参奖励 preset 仅支持 Flat 任务")
|
||||
if not isinstance(preset_id, str) or not re.fullmatch(r"[0-9a-f]{32}", preset_id):
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, "rewardPresetId 格式无效")
|
||||
if self.preset_resolver is None:
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, "奖励 preset 服务未就绪")
|
||||
try:
|
||||
reward_config = self.preset_resolver(preset_id)
|
||||
reward_config = validate_configuration(
|
||||
self.preset_resolver(preset_id, task_id), task_id
|
||||
)
|
||||
except KeyError as error:
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, "奖励 preset 不存在") from error
|
||||
except RewardConfigError as error:
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, str(error)) from error
|
||||
seed = integer("seed", 0, 2_147_483_647)
|
||||
try:
|
||||
custom = validate_task_config(task_id, payload, seed)
|
||||
except TaskConfigError as error:
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, str(error)) from error
|
||||
pretrained = None
|
||||
if "pretrainedSourceId" in payload:
|
||||
if self.sources is None:
|
||||
raise ApiError(
|
||||
HTTPStatus.BAD_REQUEST, "服务尚未注册基础策略,请配置--pretrained-sources"
|
||||
)
|
||||
try:
|
||||
pretrained = self.sources.bind(payload["pretrainedSourceId"], task_id, custom)
|
||||
except SourceError as error:
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, str(error)) from error
|
||||
return TrainingConfig(
|
||||
pretrained=pretrained,
|
||||
task_id=task_id,
|
||||
terrain_preset=custom["terrainPreset"] if custom else None,
|
||||
terrain_params=custom["terrainParams"] if custom else {},
|
||||
sensor_cfg=custom["sensorCfg"] if custom else None,
|
||||
task_config=custom,
|
||||
deployment=deployment_metadata(task_id, custom, seed),
|
||||
num_envs=integer("numEnvs", 1, 16384),
|
||||
max_iterations=integer("maxIterations", 1, 1_000_000),
|
||||
seed=integer("seed", 0, 2_147_483_647),
|
||||
seed=seed,
|
||||
run_name=run_name,
|
||||
device=device,
|
||||
gpu_ids=raw_gpu_ids,
|
||||
@@ -325,7 +397,9 @@ class TrainingManager:
|
||||
os.killpg(process.pid, signal.SIGKILL)
|
||||
process.wait(timeout=2)
|
||||
|
||||
def command_for(self, config: TrainingConfig) -> list[str]:
|
||||
def command_for(
|
||||
self, config: TrainingConfig, task_config_path: Path | None = None
|
||||
) -> list[str]:
|
||||
command = [
|
||||
self.python,
|
||||
"-u",
|
||||
@@ -336,6 +410,12 @@ class TrainingManager:
|
||||
f"--agent.seed={config.seed}",
|
||||
f"--agent.run-name={config.run_name}",
|
||||
]
|
||||
if task_config_path is not None:
|
||||
command.extend(("--task-config", str(task_config_path)))
|
||||
if config.pretrained is not None:
|
||||
if self.sources is None:
|
||||
raise SourceError("基础策略快照服务未配置,拒绝随机初始化")
|
||||
command.extend(self.sources.arguments(config.pretrained))
|
||||
if config.reward_config is not None:
|
||||
command.extend(
|
||||
(
|
||||
@@ -387,13 +467,21 @@ class TrainingManager:
|
||||
|
||||
def _run(self, job: TrainingJob) -> None:
|
||||
before = self._artifact_snapshot()
|
||||
command = self.command_for(job.config)
|
||||
environment = os.environ.copy()
|
||||
# 默认离线记录,保留本地 W&B 指标但不要求 API Key;只有前端明确选择
|
||||
# online 时才允许 wandb 发起登录/联网。
|
||||
environment["WANDB_MODE"] = job.config.wandb_mode
|
||||
environment["WANDB_SILENT"] = "true"
|
||||
try:
|
||||
config_path = None
|
||||
if job.config.task_config is not None:
|
||||
job_dir = self.trainer_root / "logs" / "rsl_rl" / "web_jobs" / job.id
|
||||
job_dir.mkdir(parents=True, exist_ok=True)
|
||||
config_path = job_dir / "training_config.json"
|
||||
config_path.write_text(json.dumps(job.config.task_config), encoding="utf-8")
|
||||
command = self.command_for(job.config, config_path)
|
||||
if config_path is not None:
|
||||
command.extend(("--output-dir", str(config_path.parent)))
|
||||
# Popen 与 process 登记必须和取消检查处于同一个临界区:cancel() 要么在
|
||||
# 创建前标记取消,要么在创建后取得进程并终止,不能落入二者之间。
|
||||
with self.lock:
|
||||
@@ -500,7 +588,7 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
|
||||
self._json(HTTPStatus.NOT_FOUND, {"error": "调参 session、trial 或 proposal 不存在"})
|
||||
elif isinstance(error, ResourceBusyError):
|
||||
self._json(HTTPStatus.CONFLICT, {"error": str(error)})
|
||||
elif isinstance(error, (TuningError, RewardConfigError, EvaluationError)):
|
||||
elif isinstance(error, (TuningError, RewardConfigError, EvaluationError, SourceError)):
|
||||
self._json(HTTPStatus.BAD_REQUEST, {"error": str(error)})
|
||||
else:
|
||||
self._json(
|
||||
@@ -523,15 +611,47 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
|
||||
length = int(self.headers.get("Content-Length", "0"))
|
||||
except ValueError as error:
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, "Content-Length 无效") from error
|
||||
if length <= 0 or length > 32 * 1024:
|
||||
if length <= 0 or length > MAX_REQUEST_BYTES:
|
||||
raise ApiError(
|
||||
HTTPStatus.REQUEST_ENTITY_TOO_LARGE, "训练请求体不能为空且不能超过 32 KiB"
|
||||
HTTPStatus.REQUEST_ENTITY_TOO_LARGE, "训练请求体不能为空且不能超过 128 KiB"
|
||||
)
|
||||
try:
|
||||
return json.loads(self.rfile.read(length))
|
||||
except (UnicodeDecodeError, json.JSONDecodeError) as error:
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, "训练请求不是有效 JSON") from error
|
||||
|
||||
def _upload(self):
|
||||
# This dedicated binary route must never pass through the 128KiB JSON reader.
|
||||
self.close_connection = True # no unread extra bytes may become a second request
|
||||
if self.manager.sources is None:
|
||||
raise ApiError(HTTPStatus.SERVICE_UNAVAILABLE, "基础策略上传服务不可用")
|
||||
if self.headers.get("Transfer-Encoding") or self.headers.get("Content-Encoding"):
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, "上传不接受chunked/压缩编码")
|
||||
lengths = self.headers.get_all("Content-Length", [])
|
||||
if len(lengths) != 1 or not re.fullmatch(r"[0-9]{1,10}", lengths[0]):
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, "上传必须提供唯一Content-Length")
|
||||
length = int(lengths[0])
|
||||
if self.headers.get("Content-Type") != "application/octet-stream":
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, "上传Content-Type必须是application/octet-stream")
|
||||
query = parse_qs(urlsplit(self.path).query, keep_blank_values=True)
|
||||
if set(query) != {"format", "template", "name"} or any(len(v) != 1 for v in query.values()):
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, "上传必须显式提供format/template/name")
|
||||
fmt, template, name = (query[key][0] for key in ("format", "template", "name"))
|
||||
limit = {"pt": 256 * 1024**2, "onnx": 64 * 1024**2}.get(fmt)
|
||||
if limit is not None and (length == 0 or length > limit):
|
||||
raise ApiError(HTTPStatus.REQUEST_ENTITY_TOO_LARGE, "上传为空或超过格式大小上限")
|
||||
if len(name) > 255:
|
||||
raise ApiError(HTTPStatus.BAD_REQUEST, "上传显示名称过长")
|
||||
try:
|
||||
result = self.manager.sources.receive_upload(
|
||||
self.rfile, length, fmt, template, name, set_timeout=self.connection.settimeout,
|
||||
)
|
||||
except OSError as error:
|
||||
raise ApiError(
|
||||
HTTPStatus.SERVICE_UNAVAILABLE, "上传存储写入失败,请检查本机磁盘空间/权限"
|
||||
) from error
|
||||
self._json(HTTPStatus.CREATED, result)
|
||||
|
||||
@staticmethod
|
||||
def _route(path: str) -> tuple[str | None, bool]:
|
||||
match = re.fullmatch(r"/api/training/jobs/([0-9a-f]{32})(/artifacts/policy\.onnx)?", path)
|
||||
@@ -628,6 +748,9 @@ class TrainingRequestHandler(BaseHTTPRequestHandler):
|
||||
try:
|
||||
self._ensure_request()
|
||||
path = urlsplit(self.path).path
|
||||
if path == "/api/training/pretrained-sources/upload":
|
||||
self._upload()
|
||||
return
|
||||
if path == "/api/training/jobs":
|
||||
self._json(HTTPStatus.ACCEPTED, self.manager.start(self._payload()))
|
||||
return
|
||||
@@ -742,6 +865,12 @@ def parse_args() -> argparse.Namespace:
|
||||
default=None,
|
||||
help="调参 SQLite 与 trial 产物目录;默认位于训练工程 logs/auto_tuning",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--pretrained-sources",
|
||||
type=Path,
|
||||
default=None,
|
||||
help="可选旧管理员基础策略注册JSON;不配置也可直接上传单个.pt/.onnx",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--task", action="append", dest="tasks", help="允许前端启动的任务 ID;可重复"
|
||||
)
|
||||
@@ -762,14 +891,23 @@ def main() -> None:
|
||||
if len(token) < 16:
|
||||
raise SystemExit("训练服务访问令牌至少需要 16 个字符")
|
||||
lease = GpuLease()
|
||||
tuning_root = args.tuning_data_root or (Path(args.trainer_root) / "logs" / "auto_tuning")
|
||||
sources = PretrainedSources(
|
||||
args.pretrained_sources,
|
||||
tuning_root / "pretrained_sources",
|
||||
args.trainer_python,
|
||||
args.trainer_root,
|
||||
)
|
||||
manager = TrainingManager(
|
||||
args.trainer_root,
|
||||
args.trainer_python,
|
||||
tuple(args.tasks or DEFAULT_TASKS),
|
||||
lease=lease,
|
||||
sources=sources,
|
||||
)
|
||||
tuning_manager = TuningManager(
|
||||
args.trainer_root, args.trainer_python, tuning_root, lease, sources=sources
|
||||
)
|
||||
tuning_root = args.tuning_data_root or (Path(args.trainer_root) / "logs" / "auto_tuning")
|
||||
tuning_manager = TuningManager(args.trainer_root, args.trainer_python, tuning_root, lease)
|
||||
manager.preset_resolver = tuning_manager.preset_config
|
||||
TrainingRequestHandler.manager = manager
|
||||
TrainingRequestHandler.tuning_manager = tuning_manager
|
||||
|
||||
@@ -0,0 +1,479 @@
|
||||
"""Validated, dependency-free browser training/deployment contract (version 1)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import math
|
||||
import random
|
||||
from typing import Any
|
||||
|
||||
FLAT_TASK = "Unitree-Go2-Flat"
|
||||
ROUGH_TASK = "Unitree-Go2-Rough"
|
||||
OBSTACLE_TASK = "Unitree-Go2-ObstacleAvoidance"
|
||||
TERRAIN_PRESETS = ("plane", "discrete_obstacles", "rough", "pyramid_stairs", "wave", "custom_boxes")
|
||||
# Range metadata is also the sole validation source for browser configuration.
|
||||
TERRAIN_PARAMETERS = {
|
||||
"size": {"min": 8, "max": 24, "default": 12},
|
||||
"obstacle_count": {"min": 1, "max": 100, "default": 24, "integer": True},
|
||||
"obstacle_height_min": {"min": 0.05, "max": 1.5, "default": 0.2},
|
||||
"obstacle_height_max": {"min": 0.05, "max": 1.5, "default": 0.6},
|
||||
"spacing": {"min": 0.6, "max": 3, "default": 1.2},
|
||||
"friction": {"min": 0.2, "max": 2, "default": 0.8},
|
||||
"roughness": {"min": 0.01, "max": 0.2, "default": 0.06},
|
||||
"step_height": {"min": 0.03, "max": 0.2, "default": 0.08},
|
||||
"wave_amplitude": {"min": 0.01, "max": 0.2, "default": 0.08},
|
||||
}
|
||||
SENSOR_PARAMETERS = {
|
||||
"fov": {"min": 30, "max": 120, "default": 90},
|
||||
"maxDistance": {"min": 1, "max": 5, "default": 4},
|
||||
"safetyDistance": {"min": 0.1, "max": 1, "default": 0.5},
|
||||
"avoidanceWeight": {"min": 0, "max": 10, "default": 2},
|
||||
}
|
||||
JOINT_NAMES = [
|
||||
f"{leg}_{joint}_joint" for leg in ("FL", "FR", "RL", "RR") for joint in ("hip", "thigh", "calf")
|
||||
]
|
||||
DEFAULT_JOINT_POSITION = [-0.1, 0.9, -1.8, 0.1, 0.9, -1.8] * 2
|
||||
|
||||
|
||||
class TaskConfigError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
def _parameters(value: Any, schema: dict, label: str) -> dict:
|
||||
if not isinstance(value, dict) or value.keys() - schema.keys():
|
||||
raise TaskConfigError(f"{label} 包含未知参数或不是对象")
|
||||
result = {}
|
||||
for name, bounds in schema.items():
|
||||
number = value.get(name, bounds["default"])
|
||||
if (
|
||||
isinstance(number, bool)
|
||||
or not isinstance(number, (int, float))
|
||||
or not bounds["min"] <= number <= bounds["max"]
|
||||
or not math.isfinite(number)
|
||||
or (bounds.get("integer") and not isinstance(number, int))
|
||||
):
|
||||
raise TaskConfigError(f"{label}.{name} 超出允许的有限数值范围")
|
||||
result[name] = number
|
||||
return result
|
||||
|
||||
|
||||
def sensor_pattern(mode: str, fov: float) -> dict:
|
||||
if mode not in ("single_ring_raycast", "multi_ring_raycast"):
|
||||
raise TaskConfigError("不支持的 sensorMode")
|
||||
multi = mode == "multi_ring_raycast"
|
||||
count = 16 if multi else 32
|
||||
return {
|
||||
"sensorMode": mode,
|
||||
"rayCount": 48 if multi else 32,
|
||||
"pitchAngles": [0, -20, -45] if multi else [0],
|
||||
"yawCount": count,
|
||||
"yawAngles": [-fov / 2 + i * fov / (count - 1) for i in range(count)],
|
||||
"angleUnit": "deg",
|
||||
"rayOrder": "layer-major",
|
||||
}
|
||||
|
||||
|
||||
def validate_sensor_config(raw: dict) -> dict:
|
||||
pattern_keys = set(sensor_pattern("single_ring_raycast", 90))
|
||||
sensor = _parameters(
|
||||
{k: v for k, v in raw.items() if k not in pattern_keys | {"type"}},
|
||||
SENSOR_PARAMETERS,
|
||||
"sensorCfg",
|
||||
)
|
||||
pattern = sensor_pattern(raw.get("sensorMode", "single_ring_raycast"), sensor["fov"])
|
||||
for key, expected in pattern.items():
|
||||
if key not in raw:
|
||||
continue
|
||||
actual = raw[key]
|
||||
if isinstance(expected, list):
|
||||
valid = (
|
||||
isinstance(actual, list)
|
||||
and len(actual) == len(expected)
|
||||
and all(
|
||||
not isinstance(a, bool)
|
||||
and isinstance(a, (int, float))
|
||||
and math.isfinite(a)
|
||||
and abs(a - b) <= 1e-10
|
||||
for a, b in zip(actual, expected, strict=True)
|
||||
)
|
||||
)
|
||||
else:
|
||||
valid = type(actual) is type(expected) and actual == expected
|
||||
if not valid:
|
||||
raise TaskConfigError(f"sensorCfg.{key} 与mode/FOV矛盾")
|
||||
return {**sensor, "type": "raycast", **pattern}
|
||||
|
||||
|
||||
def validate_custom_terrain(value: Any) -> dict:
|
||||
"""Untrusted full layout: exact schema, no RNG, no geometry repair or clipping."""
|
||||
fields = {
|
||||
"representation",
|
||||
"approximation",
|
||||
"size",
|
||||
"friction",
|
||||
"boxes",
|
||||
"spawn",
|
||||
"spawnQuaternion",
|
||||
"target",
|
||||
"actualObstacleCount",
|
||||
}
|
||||
if not isinstance(value, dict) or set(value) != fields:
|
||||
raise TaskConfigError("customTerrainBoxes 字段缺失或未知(不接受路径/MJCF)")
|
||||
|
||||
def number(v, lo, hi):
|
||||
if (
|
||||
isinstance(v, bool)
|
||||
or not isinstance(v, (int, float))
|
||||
or not lo <= v <= hi
|
||||
or not math.isfinite(v)
|
||||
):
|
||||
raise TaskConfigError("customTerrainBoxes 必须使用范围内有限数值")
|
||||
return v
|
||||
|
||||
def vector(v, n, lo=-12, hi=12):
|
||||
if not isinstance(v, list) or len(v) != n:
|
||||
raise TaskConfigError("customTerrainBoxes 向量长度无效")
|
||||
return [number(x, lo, hi) for x in v]
|
||||
|
||||
size = number(value["size"], 8, 24)
|
||||
number(value["friction"], 0.2, 2)
|
||||
if value["representation"] != "boxes-v1" or value["approximation"] is not True:
|
||||
raise TaskConfigError(
|
||||
"custom_boxes 必须声明 boxes-v1 和 approximation=true(AABB/底板标准化)"
|
||||
)
|
||||
boxes = value["boxes"]
|
||||
if not isinstance(boxes, list) or not 1 <= len(boxes) <= 257:
|
||||
raise TaskConfigError("customTerrainBoxes 只能包含底板及最多256障碍")
|
||||
for box in boxes:
|
||||
if not isinstance(box, dict) or set(box) != {"pos", "size", "yaw"}:
|
||||
raise TaskConfigError("box 字段缺失或未知")
|
||||
center = vector(box["pos"], 3)
|
||||
half = vector(box["size"], 3, 0, 12)
|
||||
number(box["yaw"], 0, 0)
|
||||
if any(x <= 0 for x in half):
|
||||
raise TaskConfigError("box 半尺寸必须严格大于零")
|
||||
if any(abs(center[i]) + half[i] > size / 2 + 1e-6 for i in range(2)):
|
||||
raise TaskConfigError("box 超出世界地图边界")
|
||||
if center[2] - half[2] < -0.2 - 1e-6 or center[2] + half[2] > 12:
|
||||
raise TaskConfigError("box 高度超出边界")
|
||||
if boxes[0] != {"pos": [0, 0, -0.1], "size": [size / 2, size / 2, 0.1], "yaw": 0}:
|
||||
raise TaskConfigError("必须使用标准 floor z=[-0.2,0]")
|
||||
count = value["actualObstacleCount"]
|
||||
if isinstance(count, bool) or not isinstance(count, int) or count != len(boxes) - 1:
|
||||
raise TaskConfigError("actualObstacleCount 与布局不一致")
|
||||
spawn = vector(value["spawn"], 3)
|
||||
target = vector(value["target"], 2)
|
||||
if spawn[2] != 0.32:
|
||||
raise TaskConfigError("出生高度必须为0.32")
|
||||
quaternion = vector(value["spawnQuaternion"], 4, -1, 1)
|
||||
if abs(sum(x * x for x in quaternion) - 1) > 1e-6:
|
||||
raise TaskConfigError("出生四元数必须归一化")
|
||||
for point in (spawn, target):
|
||||
if any(abs(point[i]) > size / 2 - 0.5 for i in range(2)):
|
||||
raise TaskConfigError("起终点0.5m安全区超出地图")
|
||||
for box in boxes[1:]:
|
||||
distance_sq = sum(
|
||||
max(abs(point[i] - box["pos"][i]) - box["size"][i], 0) ** 2 for i in range(2)
|
||||
)
|
||||
if distance_sq <= 0.5**2:
|
||||
raise TaskConfigError("障碍物侵占起终点0.5m圆形安全区;请修改坐标,不会清除障碍")
|
||||
return copy.deepcopy(value)
|
||||
|
||||
|
||||
def validate_task_config(task_id: str, payload: dict, seed: int) -> dict | None:
|
||||
"""Validate preset parameters or an authoritative boxes-v1 layout, never paths/XML."""
|
||||
custom_keys = {
|
||||
"terrainPreset",
|
||||
"terrainParams",
|
||||
"sensorCfg",
|
||||
"sensorType",
|
||||
"customTerrainBoxes",
|
||||
}
|
||||
if task_id != OBSTACLE_TASK and not custom_keys.intersection(payload):
|
||||
return None
|
||||
if task_id not in (FLAT_TASK, ROUGH_TASK, OBSTACLE_TASK):
|
||||
raise TaskConfigError("该任务不支持自定义地形")
|
||||
preset = payload.get(
|
||||
"terrainPreset", "discrete_obstacles" if task_id == OBSTACLE_TASK else "plane"
|
||||
)
|
||||
if not isinstance(preset, str) or preset not in TERRAIN_PRESETS:
|
||||
raise TaskConfigError("不支持的 terrainPreset")
|
||||
layout = None
|
||||
if preset == "custom_boxes":
|
||||
layout = validate_custom_terrain(payload.get("customTerrainBoxes"))
|
||||
expected = {"size": layout["size"], "friction": layout["friction"]}
|
||||
if "terrainParams" in payload and (
|
||||
not isinstance(payload["terrainParams"], dict)
|
||||
or payload["terrainParams"] != expected
|
||||
or any(
|
||||
isinstance(v, bool) or not isinstance(v, (int, float))
|
||||
for v in payload["terrainParams"].values()
|
||||
)
|
||||
):
|
||||
raise TaskConfigError("custom_boxes terrainParams 必须与布局size/friction完全一致")
|
||||
terrain = expected
|
||||
else:
|
||||
if "customTerrainBoxes" in payload:
|
||||
raise TaskConfigError("customTerrainBoxes 仅允许 custom_boxes,不能降级预设")
|
||||
terrain = _parameters(payload.get("terrainParams", {}), TERRAIN_PARAMETERS, "terrainParams")
|
||||
if not layout and terrain["obstacle_height_min"] > terrain["obstacle_height_max"]:
|
||||
raise TaskConfigError("障碍物最小高度不能超过最大高度")
|
||||
raw_sensor = payload.get("sensorCfg", {})
|
||||
if not isinstance(raw_sensor, dict):
|
||||
raise TaskConfigError("sensorCfg 必须是对象")
|
||||
sensor_type = payload.get("sensorType", raw_sensor.get("type", "raycast"))
|
||||
if sensor_type != "raycast" or raw_sensor.get("type", "raycast") != "raycast":
|
||||
raise TaskConfigError("首版只支持 raycast,未实现 camera_depth")
|
||||
if task_id != OBSTACLE_TASK and (raw_sensor or "sensorType" in payload):
|
||||
raise TaskConfigError("只有避障任务支持 sensorCfg")
|
||||
sensor = None
|
||||
if task_id == OBSTACLE_TASK:
|
||||
sensor = validate_sensor_config(raw_sensor)
|
||||
sensor["type"] = "raycast"
|
||||
if sensor["safetyDistance"] >= sensor["maxDistance"]:
|
||||
raise TaskConfigError("安全距离必须小于探测距离")
|
||||
return {
|
||||
"terrainPreset": preset,
|
||||
"terrainParams": terrain,
|
||||
"sensorCfg": sensor,
|
||||
"seed": seed,
|
||||
**({"customTerrainBoxes": layout} if layout else {}),
|
||||
}
|
||||
|
||||
|
||||
def navigation_candidates(
|
||||
layout: dict, spacing: float = 0.5, clearance: float = 0.55, min_distance: float = 2.0
|
||||
) -> dict:
|
||||
"""Build connected, collision-free point-goal samples for episode resets.
|
||||
|
||||
Obstacles are inflated by ``clearance`` and the map is sampled on a regular grid.
|
||||
Only components containing a pair at least ``min_distance`` apart are retained, so
|
||||
runtime sampling cannot place a goal across an impassable wall.
|
||||
"""
|
||||
size = layout["size"]
|
||||
lo, hi = -size / 2 + clearance, size / 2 - clearance
|
||||
count = max(1, int(math.floor((hi - lo) / spacing)) + 1)
|
||||
axis = [lo + i * (hi - lo) / max(count - 1, 1) for i in range(count)]
|
||||
obstacles = layout["boxes"][1:]
|
||||
|
||||
def free(x: float, y: float) -> bool:
|
||||
return all(
|
||||
math.hypot(
|
||||
max(abs(x - box["pos"][0]) - box["size"][0], 0),
|
||||
max(abs(y - box["pos"][1]) - box["size"][1], 0),
|
||||
)
|
||||
> clearance
|
||||
for box in obstacles
|
||||
)
|
||||
|
||||
cells = {(i, j) for i, x in enumerate(axis) for j, y in enumerate(axis) if free(x, y)}
|
||||
components = []
|
||||
while cells:
|
||||
pending = [cells.pop()]
|
||||
component = []
|
||||
while pending:
|
||||
cell = pending.pop()
|
||||
component.append(cell)
|
||||
i, j = cell
|
||||
for neighbour in ((i - 1, j), (i + 1, j), (i, j - 1), (i, j + 1)):
|
||||
if neighbour in cells:
|
||||
cells.remove(neighbour)
|
||||
pending.append(neighbour)
|
||||
points = [(axis[i], axis[j]) for i, j in sorted(component)]
|
||||
extrema = [
|
||||
(
|
||||
min(points, key=lambda point: point[axis_index]),
|
||||
max(points, key=lambda point: point[axis_index]),
|
||||
)
|
||||
for axis_index in (0, 1)
|
||||
]
|
||||
first, second = max(extrema, key=lambda pair: math.dist(*pair))
|
||||
if math.dist(first, second) >= min_distance:
|
||||
components.append((points, first, second))
|
||||
if not components:
|
||||
raise TaskConfigError(
|
||||
f"地图没有可用于随机起终点的连通自由区域(至少需要{min_distance:g}m间距)"
|
||||
)
|
||||
flattened = []
|
||||
starts = []
|
||||
counts = []
|
||||
fallbacks = []
|
||||
for points, first, second in components:
|
||||
starts.append(len(flattened))
|
||||
counts.append(len(points))
|
||||
flattened.extend(points)
|
||||
fallbacks.append((first, second))
|
||||
return {
|
||||
"points": flattened,
|
||||
"componentStarts": starts,
|
||||
"componentCounts": counts,
|
||||
"fallbackPairs": fallbacks,
|
||||
"minDistance": min_distance,
|
||||
"clearance": clearance,
|
||||
"spacing": spacing,
|
||||
}
|
||||
|
||||
|
||||
def build_terrain_layout(config: dict) -> dict:
|
||||
"""Emit exact world-centered boxes; consumers do not reimplement RNG/terrain presets.
|
||||
|
||||
size values are MuJoCo half-extents, pos is the center, yaw is always zero.
|
||||
TerrainGenerator uses one patch, shifted back to these coordinates. All parallel
|
||||
environments are independent worlds with the same map and spawn, not a grid.
|
||||
"""
|
||||
if config["terrainPreset"] == "custom_boxes":
|
||||
return validate_custom_terrain(config.get("customTerrainBoxes"))
|
||||
p = config["terrainParams"]
|
||||
size = p["size"]
|
||||
rng = random.Random(config["seed"])
|
||||
boxes = [{"pos": [0, 0, -0.1], "size": [size / 2, size / 2, 0.1], "yaw": 0}]
|
||||
preset = config["terrainPreset"]
|
||||
|
||||
def box(x: float, y: float, sx: float, sy: float, height: float) -> None:
|
||||
boxes.append({"pos": [x, y, height / 2], "size": [sx, sy, height / 2], "yaw": 0})
|
||||
|
||||
# Leave two flat end strips for the spawn and goal; no rejection sampling.
|
||||
if preset == "discrete_obstacles":
|
||||
spacing = p["spacing"]
|
||||
nx = max(1, int((size - 4) / spacing))
|
||||
ny = max(1, int((size - 2) / spacing))
|
||||
cells = [
|
||||
(
|
||||
-size / 2 + 2 + (i + 0.5) * (size - 4) / nx,
|
||||
-size / 2 + 1 + (j + 0.5) * (size - 2) / ny,
|
||||
)
|
||||
for i in range(nx)
|
||||
for j in range(ny)
|
||||
]
|
||||
rng.shuffle(cells)
|
||||
for x, y in cells[: p["obstacle_count"]]:
|
||||
box(x, y, 0.2, 0.2, rng.uniform(p["obstacle_height_min"], p["obstacle_height_max"]))
|
||||
elif preset in ("rough", "wave"):
|
||||
# 16x16 bounded box approximation, deliberately not the editor heightfield.
|
||||
nx = ny = 16
|
||||
sx, sy = (size - 4) / nx, size / ny
|
||||
for i in range(nx):
|
||||
for j in range(ny):
|
||||
height = (
|
||||
rng.uniform(0.005, p["roughness"])
|
||||
if preset == "rough"
|
||||
else 0.005 + p["wave_amplitude"] * (1 + math.sin(i * math.pi / 4)) / 2
|
||||
)
|
||||
box(
|
||||
-size / 2 + 2 + (i + 0.5) * sx,
|
||||
-size / 2 + (j + 0.5) * sy,
|
||||
sx / 2,
|
||||
sy / 2,
|
||||
height,
|
||||
)
|
||||
elif preset == "pyramid_stairs":
|
||||
for i in range(4):
|
||||
half = (size - 4) / 2 - i * (size - 4) / 10
|
||||
box(0, 0, half, half, p["step_height"] * (i + 1))
|
||||
return {
|
||||
"representation": "boxes-v1",
|
||||
"approximation": preset in ("rough", "wave", "pyramid_stairs"),
|
||||
"size": size,
|
||||
"friction": p["friction"],
|
||||
"boxes": boxes,
|
||||
"spawn": [-size / 2 + 1, 0, 0.32],
|
||||
"spawnQuaternion": [1, 0, 0, 0],
|
||||
"target": [size / 2 - 1, 0],
|
||||
"actualObstacleCount": len(boxes) - 1,
|
||||
}
|
||||
|
||||
|
||||
def deployment_metadata(task_id: str, config: dict | None, seed: int) -> dict:
|
||||
obstacle = task_id == OBSTACLE_TASK
|
||||
metadata = {
|
||||
"version": 1,
|
||||
"taskId": task_id,
|
||||
"browserCompatible": task_id in (FLAT_TASK, OBSTACLE_TASK),
|
||||
"observationSize": (49 + config["sensorCfg"]["rayCount"])
|
||||
if obstacle
|
||||
else (234 if task_id == ROUGH_TASK else 47),
|
||||
"actionSize": 12,
|
||||
"controlHz": 50,
|
||||
"gaitPeriod": 0.6,
|
||||
"jointNames": JOINT_NAMES,
|
||||
"defaultJointPosition": DEFAULT_JOINT_POSITION,
|
||||
"actionScale": [0.25] * 12,
|
||||
"stiffness": [20, 20, 40] * 4,
|
||||
"damping": [1, 1, 2] * 4,
|
||||
"effortLimits": [23.5, 23.5, 45] * 4,
|
||||
"observationTerms": [
|
||||
"base_ang_vel",
|
||||
"projected_gravity",
|
||||
"command",
|
||||
"phase",
|
||||
"joint_pos",
|
||||
"joint_vel",
|
||||
"actions",
|
||||
]
|
||||
+ (
|
||||
["forward_depth", "target_error"]
|
||||
if obstacle
|
||||
else (["height_scan"] if task_id == ROUGH_TASK else [])
|
||||
),
|
||||
"seed": seed,
|
||||
}
|
||||
if task_id == ROUGH_TASK:
|
||||
metadata["incompatibilityReason"] = "旧 Rough actor 含向下高度扫描;浏览器未实现此部署契约"
|
||||
if config is not None:
|
||||
metadata.update(
|
||||
{
|
||||
"terrainPreset": config["terrainPreset"],
|
||||
"terrainParams": config["terrainParams"],
|
||||
"terrain": build_terrain_layout(config),
|
||||
}
|
||||
)
|
||||
if obstacle:
|
||||
assert config is not None
|
||||
metadata["sensorCfg"] = {
|
||||
**config["sensorCfg"],
|
||||
"offset": [0.3, 0, 0.05],
|
||||
"alignment": "base",
|
||||
"terrainOnly": True,
|
||||
"includeGround": True,
|
||||
}
|
||||
metadata["navigation"] = {
|
||||
"speed": 0.6,
|
||||
"arrivalRadius": 0.5,
|
||||
"distanceScale": config["terrainParams"]["size"],
|
||||
"headingScale": math.pi,
|
||||
"yawGain": 1.0,
|
||||
"maxYawRate": 1.0,
|
||||
"episodeSeconds": 20,
|
||||
"onArrival": "stop",
|
||||
"onReset": "respawn",
|
||||
"trainingReset": "random-connected-free-pair",
|
||||
"trainingMinGoalDistance": 2.0,
|
||||
"trainingClearance": 0.55,
|
||||
}
|
||||
return metadata
|
||||
|
||||
|
||||
def task_metadata(tasks: tuple[str, ...]) -> list[dict]:
|
||||
names = {
|
||||
FLAT_TASK: "平地速度控制",
|
||||
ROUGH_TASK: "地形自适应(仅后端)",
|
||||
OBSTACLE_TASK: "前视射线避障导航",
|
||||
}
|
||||
return [
|
||||
{
|
||||
"id": task,
|
||||
"name": names.get(task, task),
|
||||
"browserCompatible": task in (FLAT_TASK, OBSTACLE_TASK),
|
||||
"terrainPresets": list(TERRAIN_PRESETS) if task in names else [],
|
||||
"terrainParameters": TERRAIN_PARAMETERS if task in names else {},
|
||||
"sensorTypes": ["raycast"] if task == OBSTACLE_TASK else [],
|
||||
"sensorModes": ["single_ring_raycast", "multi_ring_raycast"]
|
||||
if task == OBSTACLE_TASK
|
||||
else [],
|
||||
"sensorParameters": SENSOR_PARAMETERS if task == OBSTACLE_TASK else {},
|
||||
"mapSyncScope": (
|
||||
"已应用静态碰撞场景→custom_boxes(boxes-v1);AABB与标准底板近似;不支持mesh/hfield"
|
||||
),
|
||||
}
|
||||
for task in tasks
|
||||
]
|
||||
@@ -0,0 +1,14 @@
|
||||
{
|
||||
"representation": "boxes-v1",
|
||||
"approximation": true,
|
||||
"size": 12,
|
||||
"friction": 0.8,
|
||||
"boxes": [
|
||||
{ "pos": [0, 0, -0.1], "size": [6, 6, 0.1], "yaw": 0 },
|
||||
{ "pos": [1, 2, 0.5], "size": [0.4, 0.3, 0.5], "yaw": 0 }
|
||||
],
|
||||
"spawn": [-2, -1, 0.32],
|
||||
"spawnQuaternion": [0.7071067811865476, 0, 0, 0.7071067811865476],
|
||||
"target": [2, -1],
|
||||
"actualObstacleCount": 1
|
||||
}
|
||||
@@ -0,0 +1,90 @@
|
||||
"""Regenerate CPU mj_ray references: run from repository root in .venv."""
|
||||
|
||||
import json
|
||||
import math
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import mujoco
|
||||
import numpy as np
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
from task_config import OBSTACLE_TASK, deployment_metadata, validate_task_config # noqa: E402
|
||||
|
||||
|
||||
def generate():
|
||||
config = validate_task_config(
|
||||
OBSTACLE_TASK, {"sensorCfg": {"sensorMode": "multi_ring_raycast"}}, 42
|
||||
)
|
||||
deployment = deployment_metadata(OBSTACLE_TASK, config, 42)
|
||||
boxes = deployment["terrain"]["boxes"]
|
||||
# Low obstacle directly ahead; platform edge is a bounded-floor drop (no invented height map).
|
||||
layouts = {
|
||||
"default": boxes,
|
||||
"low": [boxes[0], {"pos": [0.9, 0, 0.025], "size": [0.2, 0.5, 0.025], "yaw": 0}],
|
||||
"edge": [boxes[0]],
|
||||
}
|
||||
cases = []
|
||||
for layout, position, euler in [
|
||||
("default", [-5, 0, 0.32], [0, 0, 0]),
|
||||
("default", [-3, 1, 0.42], [0.23, -0.31, 0.51]),
|
||||
("default", [-2, -2, 0.5], [-0.35, 0.27, -0.63]),
|
||||
("low", [0, 0, 0.32], [0, 0, 0]),
|
||||
("low", [0, 0, 0.32], [0.2, 0.1, -0.1]),
|
||||
("edge", [5.45, 0, 0.32], [0, 0, 0]),
|
||||
]:
|
||||
xml = (
|
||||
"<mujoco><worldbody>"
|
||||
+ "".join(
|
||||
'<geom type="box" pos="{}" size="{}"/>'.format(
|
||||
" ".join(map(str, b["pos"])), " ".join(map(str, b["size"]))
|
||||
)
|
||||
for b in layouts[layout]
|
||||
)
|
||||
+ "</worldbody></mujoco>"
|
||||
)
|
||||
model = mujoco.MjModel.from_xml_string(xml)
|
||||
data = mujoco.MjData(model)
|
||||
mujoco.mj_forward(model, data)
|
||||
q = np.zeros(4)
|
||||
mujoco.mju_euler2Quat(q, np.array(euler), "xyz")
|
||||
matrix = np.zeros(9)
|
||||
mujoco.mju_quat2Mat(matrix, q)
|
||||
matrix = matrix.reshape(3, 3)
|
||||
origin = np.array(position) + matrix @ np.array([0.3, 0, 0.05])
|
||||
distances, ids = [], []
|
||||
for pitch in [0, -20, -45]:
|
||||
for i in range(16):
|
||||
yaw = math.radians(-45 + i * 90 / 15)
|
||||
p = math.radians(pitch)
|
||||
direction = matrix @ np.array(
|
||||
[math.cos(p) * math.cos(yaw), math.cos(p) * math.sin(yaw), math.sin(p)]
|
||||
)
|
||||
geom = np.array([-1], dtype=np.int32)
|
||||
distance = mujoco.mj_ray(model, data, origin, direction, None, 1, -1, geom)
|
||||
distances.append(distance)
|
||||
ids.append(int(geom[0]))
|
||||
cases.append(
|
||||
dict(
|
||||
layout=layout,
|
||||
position=position,
|
||||
quaternion=q.tolist(),
|
||||
distances=distances,
|
||||
hitIds=ids,
|
||||
depth=[1 if d < 0 else min(1, d / 4) for d in distances],
|
||||
)
|
||||
)
|
||||
root = Path("web_platform/src/rl/fixtures")
|
||||
(root / "multiRingGolden.json").write_text(
|
||||
json.dumps(
|
||||
dict(source=f"CPU MuJoCo {mujoco.__version__} mj_ray", layouts=layouts, cases=cases),
|
||||
indent=2,
|
||||
)
|
||||
+ "\n"
|
||||
)
|
||||
(root / "multiRingDeployment.json").write_text(json.dumps(deployment, indent=2) + "\n")
|
||||
Path("/tmp/go2-multi-ring-stage5/task.json").write_text(json.dumps(config))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
generate()
|
||||
@@ -0,0 +1,254 @@
|
||||
"""Authoritative custom boxes validation and CPU compilation, no training/RNG."""
|
||||
|
||||
import copy
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[1]
|
||||
for path in (ROOT, ROOT / "rl"):
|
||||
sys.path.insert(0, str(path))
|
||||
from task_config import ( # noqa: E402
|
||||
OBSTACLE_TASK,
|
||||
TaskConfigError,
|
||||
build_terrain_layout,
|
||||
deployment_metadata,
|
||||
validate_custom_terrain,
|
||||
validate_task_config,
|
||||
)
|
||||
|
||||
|
||||
def layout():
|
||||
return json.loads((Path(__file__).parent / "fixtures/custom-boxes.json").read_text())
|
||||
|
||||
|
||||
class CustomBoxesTest(unittest.TestCase):
|
||||
def test_payload_roundtrip_without_rng_and_aliasing(self):
|
||||
source = layout()
|
||||
with patch("task_config.random.Random", side_effect=AssertionError("must not run RNG")):
|
||||
config = validate_task_config(
|
||||
OBSTACLE_TASK,
|
||||
{
|
||||
"terrainPreset": "custom_boxes",
|
||||
"customTerrainBoxes": source,
|
||||
},
|
||||
42,
|
||||
)
|
||||
self.assertEqual(build_terrain_layout(config), source)
|
||||
self.assertEqual(deployment_metadata(OBSTACLE_TASK, config, 42)["terrain"], source)
|
||||
source["boxes"][1]["pos"][0] = 9
|
||||
self.assertEqual(config["customTerrainBoxes"], layout())
|
||||
|
||||
def test_strict_fields_numbers_floor_safety_and_count(self):
|
||||
mutations = [
|
||||
lambda t: t.update(path="../../etc/passwd"),
|
||||
lambda t: t.pop("target"),
|
||||
lambda t: t.update(approximation=False),
|
||||
lambda t: t.update(actualObstacleCount=True),
|
||||
lambda t: t.update(actualObstacleCount=9),
|
||||
lambda t: t.update(size=25),
|
||||
lambda t: t.update(friction=True),
|
||||
lambda t: t.update(spawnQuaternion=[2, 0, 0, 0]),
|
||||
lambda t: t.update(spawn=[-2, -1, 0.4]),
|
||||
lambda t: t.update(target=[6, 0]),
|
||||
lambda t: t["boxes"][0]["pos"].__setitem__(2, 0),
|
||||
lambda t: t["boxes"][1].update(mesh="../../model.stl"),
|
||||
lambda t: t["boxes"][1].update(yaw=True),
|
||||
lambda t: t["boxes"][1].update(pos=[6, 2, 0.5]),
|
||||
lambda t: t["boxes"][1].update(pos=[1, 2, -0.3]),
|
||||
lambda t: t["boxes"][1].update(pos=[1, 2, 12]),
|
||||
lambda t: t.update(
|
||||
boxes=t["boxes"] + [copy.deepcopy(t["boxes"][1])] * 256, actualObstacleCount=257
|
||||
),
|
||||
]
|
||||
for bad in (float("nan"), float("inf"), -float("inf"), 10**400, 0, -0.0, -1, True):
|
||||
mutations.append(lambda t, bad=bad: t["boxes"][1]["size"].__setitem__(0, bad))
|
||||
for mutate in mutations:
|
||||
value = layout()
|
||||
mutate(value)
|
||||
with self.subTest(value=str(value)[:200]), self.assertRaises(TaskConfigError):
|
||||
validate_custom_terrain(value)
|
||||
for key in ("spawn", "target"):
|
||||
value = layout()
|
||||
# Circle tangent to the right face: reject; floor alone is exempt.
|
||||
value[key][:2] = [1.9, 2]
|
||||
with self.assertRaisesRegex(TaskConfigError, "安全区"):
|
||||
validate_custom_terrain(value)
|
||||
value[key][:2] = [1.91, 2]
|
||||
validate_custom_terrain(value)
|
||||
# Precise circular corner test, not an expanded square approximation.
|
||||
value[key][:2] = [1.8, 2.7]
|
||||
validate_custom_terrain(value)
|
||||
|
||||
def test_conflict_missing_path_and_training_entry(self):
|
||||
from scripts.train import _load_task_config
|
||||
|
||||
for payload in (
|
||||
{"terrainPreset": "custom_boxes"},
|
||||
{"terrainPreset": "custom_boxes", "customTerrainBoxes": "../../file.json"},
|
||||
{"terrainPreset": "plane", "customTerrainBoxes": layout()},
|
||||
{
|
||||
"terrainPreset": "custom_boxes",
|
||||
"customTerrainBoxes": layout(),
|
||||
"terrainParams": {"size": 8, "friction": 0.8},
|
||||
},
|
||||
):
|
||||
with self.assertRaises(TaskConfigError):
|
||||
validate_task_config(OBSTACLE_TASK, payload, 42)
|
||||
config = validate_task_config(
|
||||
OBSTACLE_TASK, {"terrainPreset": "custom_boxes", "customTerrainBoxes": layout()}, 42
|
||||
)
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
path = Path(directory) / "task.json"
|
||||
path.write_text(json.dumps(config))
|
||||
self.assertEqual(_load_task_config(OBSTACLE_TASK, str(path), 42), config)
|
||||
config["customTerrainBoxes"]["boxes"][1]["size"][0] = -0.0
|
||||
path.write_text(json.dumps(config))
|
||||
with self.assertRaises(TaskConfigError):
|
||||
_load_task_config(OBSTACLE_TASK, str(path), 42)
|
||||
config["customTerrainBoxesPath"] = "/etc/passwd"
|
||||
path.write_text(json.dumps(config))
|
||||
with self.assertRaises(ValueError):
|
||||
_load_task_config(OBSTACLE_TASK, str(path), 42)
|
||||
|
||||
def test_server_payload_and_limit(self):
|
||||
from server import DEFAULT_TASKS, MAX_REQUEST_BYTES, ApiError, TrainingManager
|
||||
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
manager = TrainingManager(
|
||||
Path(directory), sys.executable, DEFAULT_TASKS, check_environment=False
|
||||
)
|
||||
payload = dict(
|
||||
taskId=OBSTACLE_TASK,
|
||||
terrainPreset="custom_boxes",
|
||||
customTerrainBoxes=layout(),
|
||||
numEnvs=2,
|
||||
maxIterations=1,
|
||||
seed=42,
|
||||
device="cpu",
|
||||
gpuIds=[],
|
||||
)
|
||||
config = manager.parse_config(payload)
|
||||
self.assertEqual(config.deployment["terrain"], layout())
|
||||
self.assertEqual(config.task_config["customTerrainBoxes"], layout())
|
||||
self.assertIn("custom_boxes", manager.health()["taskMetadata"][2]["terrainPresets"])
|
||||
self.assertIn(
|
||||
"--task-config", manager.command_for(config, Path(directory) / "server-owned.json")
|
||||
)
|
||||
with self.assertRaises(ApiError):
|
||||
manager.parse_config({**payload, "customTerrainBoxesPath": "/etc/passwd"})
|
||||
full = layout()
|
||||
full["boxes"] += [copy.deepcopy(full["boxes"][1])] * 255
|
||||
full["actualObstacleCount"] = 256
|
||||
# Worst typical double precision expansion still fits the bounded request envelope.
|
||||
full["boxes"][1:] = [
|
||||
dict(
|
||||
pos=[1.123456789012345, 2.123456789012345, 0.5123456789012345],
|
||||
size=[0.4123456789012345, 0.3123456789012345, 0.5123456789012345],
|
||||
yaw=0,
|
||||
)
|
||||
for _ in range(256)
|
||||
]
|
||||
self.assertLess(
|
||||
len(json.dumps({**payload, "customTerrainBoxes": full}).encode()), MAX_REQUEST_BYTES
|
||||
)
|
||||
manager.parse_config({**payload, "customTerrainBoxes": full})
|
||||
|
||||
def test_http_body_length_is_bounded_and_allows_full_layout(self):
|
||||
import io
|
||||
|
||||
from server import MAX_REQUEST_BYTES, ApiError, TrainingRequestHandler
|
||||
|
||||
value = layout()
|
||||
value["boxes"] += [copy.deepcopy(value["boxes"][1])] * 255
|
||||
value["actualObstacleCount"] = 256
|
||||
body = json.dumps({"terrainPreset": "custom_boxes", "customTerrainBoxes": value}).encode()
|
||||
handler = object.__new__(TrainingRequestHandler)
|
||||
handler.headers = {"Content-Length": str(len(body))}
|
||||
handler.rfile = io.BytesIO(body)
|
||||
self.assertEqual(handler._payload()["customTerrainBoxes"], value)
|
||||
for size in (0, MAX_REQUEST_BYTES + 1):
|
||||
handler.headers = {"Content-Length": str(size)}
|
||||
with self.assertRaises(ApiError):
|
||||
handler._payload()
|
||||
|
||||
def test_real_cpu_compilation_and_custom_origin_quaternion_goal(self):
|
||||
import mujoco
|
||||
import numpy as np
|
||||
from mjlab.terrains import TerrainGenerator
|
||||
from src.tasks.obstacle_avoidance.env_cfg import (
|
||||
apply_obstacle_configuration,
|
||||
unitree_go2_obstacle_env_cfg,
|
||||
)
|
||||
|
||||
value = layout()
|
||||
config = validate_task_config(
|
||||
OBSTACLE_TASK, {"terrainPreset": "custom_boxes", "customTerrainBoxes": value}, 42
|
||||
)
|
||||
cfg = unitree_go2_obstacle_env_cfg()
|
||||
apply_obstacle_configuration(cfg, config)
|
||||
generator = TerrainGenerator(cfg.scene.terrain.terrain_generator)
|
||||
spec = mujoco.MjSpec()
|
||||
generator.compile(spec)
|
||||
model = spec.compile()
|
||||
np.testing.assert_allclose(model.geom_pos, [b["pos"] for b in value["boxes"]], atol=1e-12)
|
||||
np.testing.assert_allclose(model.geom_size, [b["size"] for b in value["boxes"]], atol=1e-12)
|
||||
np.testing.assert_allclose(generator.terrain_origins[0, 0], [-2, -1, 0])
|
||||
self.assertEqual(
|
||||
cfg.scene.entities["robot"].init_state.rot, tuple(value["spawnQuaternion"])
|
||||
)
|
||||
self.assertEqual(cfg.commands["twist"].goal_offset, (4, 0))
|
||||
self.assertEqual(cfg.terminations["outside_map"].params, {"size": 12})
|
||||
self.assertTrue(math.isfinite(model.geom_size.sum()))
|
||||
|
||||
@unittest.skipUnless(os.environ.get("GO2_CUSTOM_BOXES_SMOKE") == "1", "opt-in 2env/1step smoke")
|
||||
def test_two_env_one_step_custom_layout(self):
|
||||
import numpy as np
|
||||
import torch
|
||||
import warp as wp
|
||||
from mjlab.envs import ManagerBasedRlEnv
|
||||
from src.tasks.obstacle_avoidance.env_cfg import (
|
||||
apply_obstacle_configuration,
|
||||
unitree_go2_obstacle_env_cfg,
|
||||
)
|
||||
from warp._src import context
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
self.skipTest("CUDA unavailable")
|
||||
if not hasattr(wp, "context"):
|
||||
wp.context = context
|
||||
cfg = unitree_go2_obstacle_env_cfg()
|
||||
custom = validate_task_config(
|
||||
OBSTACLE_TASK, {"terrainPreset": "custom_boxes", "customTerrainBoxes": layout()}, 42
|
||||
)
|
||||
apply_obstacle_configuration(cfg, custom, randomize_navigation=False)
|
||||
cfg.scene.num_envs = 2
|
||||
env = ManagerBasedRlEnv(cfg, device="cuda:0")
|
||||
try:
|
||||
obs, _ = env.reset()
|
||||
self.assertEqual(tuple(obs["actor"].shape), (2, 81))
|
||||
np.testing.assert_allclose(env.scene.env_origins.cpu(), [[-2, -1, 0]] * 2, atol=1e-6)
|
||||
np.testing.assert_allclose(
|
||||
env.scene["robot"].data.root_link_pos_w.cpu(), [layout()["spawn"]] * 2, atol=1e-6
|
||||
)
|
||||
np.testing.assert_allclose(
|
||||
env.scene["robot"].data.root_link_quat_w.cpu(),
|
||||
[layout()["spawnQuaternion"]] * 2,
|
||||
atol=1e-6,
|
||||
)
|
||||
np.testing.assert_allclose(
|
||||
env.command_manager.get_term("twist").errors()[1].cpu(), [4, 4], atol=1e-6
|
||||
)
|
||||
obs, reward, terminated, truncated, _ = env.step(
|
||||
torch.zeros((2, 12), device=env.device)
|
||||
)
|
||||
self.assertTrue(torch.isfinite(obs["actor"]).all())
|
||||
self.assertTrue(torch.isfinite(reward).all())
|
||||
self.assertFalse(terminated.any() or truncated.any())
|
||||
finally:
|
||||
env.close()
|
||||
@@ -0,0 +1,214 @@
|
||||
"""Multi-ring contract, CPU reference generation and opt-in real 97-D environment."""
|
||||
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[1]
|
||||
for p in (ROOT, ROOT / "rl"):
|
||||
sys.path.insert(0, str(p))
|
||||
|
||||
from task_config import ( # noqa: E402
|
||||
OBSTACLE_TASK,
|
||||
TaskConfigError,
|
||||
deployment_metadata,
|
||||
validate_task_config,
|
||||
)
|
||||
|
||||
|
||||
class MultiRingTest(unittest.TestCase):
|
||||
def test_fresh_cli_registers_flat_rough_obstacle(self):
|
||||
import subprocess
|
||||
|
||||
for task in ("Unitree-Go2-Flat", "Unitree-Go2-Rough", OBSTACLE_TASK):
|
||||
result = subprocess.run(
|
||||
[sys.executable, "-u", "scripts/train.py", task, "--help"],
|
||||
cwd=ROOT / "rl",
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=60,
|
||||
)
|
||||
self.assertEqual(result.returncode, 0, result.stdout + result.stderr)
|
||||
self.assertIn("--env.scene.num-envs", result.stdout)
|
||||
|
||||
def test_whitelist_and_legacy_defaults(self):
|
||||
for mode, count in (("single_ring_raycast", 32), ("multi_ring_raycast", 48)):
|
||||
c = validate_task_config(OBSTACLE_TASK, {"sensorCfg": {"sensorMode": mode}}, 42)
|
||||
s = c["sensorCfg"]
|
||||
self.assertEqual(s["rayCount"], count)
|
||||
self.assertEqual(
|
||||
deployment_metadata(OBSTACLE_TASK, c, 42)["observationSize"], 49 + count
|
||||
)
|
||||
self.assertEqual(validate_task_config(OBSTACLE_TASK, c, 42), c)
|
||||
for patch in (
|
||||
{"rayCount": 64},
|
||||
{"pitchAngles": [0, -45, -20]},
|
||||
{"yawCount": True},
|
||||
{"angleUnit": "rad"},
|
||||
{"rayOrder": "yaw-major"},
|
||||
{"sensorMode": "camera_depth"},
|
||||
{"yawAngles": [0] * s["yawCount"]},
|
||||
{"garbage": 1},
|
||||
):
|
||||
with self.subTest(patch=patch), self.assertRaises(TaskConfigError):
|
||||
validate_task_config(OBSTACLE_TASK, {"sensorCfg": {**s, **patch}}, 42)
|
||||
self.assertEqual(validate_task_config(OBSTACLE_TASK, {}, 42)["sensorCfg"]["rayCount"], 32)
|
||||
|
||||
def test_floor_identity_fail_closed_and_fixed_body_world_transform(self):
|
||||
import mujoco
|
||||
from src.tasks.obstacle_avoidance.mdp import standard_floor_id
|
||||
|
||||
floor = '<geom name="terrain_0" type="box" pos="-.5 0 -.1" size="6 6 .1"/>'
|
||||
model = mujoco.MjModel.from_xml_string(
|
||||
'<mujoco><worldbody><body pos=".5 0 0">' + floor + "</body></worldbody></mujoco>"
|
||||
)
|
||||
self.assertEqual(standard_floor_id(model, 12), 0)
|
||||
for geoms in (
|
||||
floor.replace("terrain_0", "other"),
|
||||
floor.replace("6 6 .1", "5 5 .1"),
|
||||
floor + floor.replace("terrain_0", "duplicate"),
|
||||
):
|
||||
m = mujoco.MjModel.from_xml_string(
|
||||
'<mujoco><worldbody><body pos=".5 0 0">' + geoms + "</body></worldbody></mujoco>"
|
||||
)
|
||||
with self.assertRaises(ValueError):
|
||||
standard_floor_id(m, 12)
|
||||
|
||||
def test_pattern_floor_classification_and_reward(self):
|
||||
import torch
|
||||
from src.tasks.obstacle_avoidance.mdp import (
|
||||
ForwardFanPatternCfg,
|
||||
forward_depth,
|
||||
obstacle_proximity,
|
||||
)
|
||||
|
||||
offsets, rays = ForwardFanPatternCfg(sensor_mode="multi_ring_raycast").generate_rays(
|
||||
None, "cpu"
|
||||
)
|
||||
self.assertEqual(tuple(rays.shape), (48, 3))
|
||||
torch.testing.assert_close(offsets, torch.tensor([0.3, 0, 0.05]).repeat(48, 1))
|
||||
torch.testing.assert_close(rays.norm(dim=1), torch.ones(48))
|
||||
for i, pitch in enumerate([0, -20, -45]):
|
||||
self.assertAlmostEqual(
|
||||
rays[i * 16, 2].item(),
|
||||
__import__("math").sin(pitch * __import__("math").pi / 180),
|
||||
places=6,
|
||||
)
|
||||
self.assertLess(rays[i * 16, 1], 0)
|
||||
self.assertGreater(rays[i * 16 + 15, 1], 0)
|
||||
# floor, 5cm obstacle, side, outside floor, miss; obs is never overwritten.
|
||||
data = SimpleNamespace(
|
||||
distances=torch.tensor([[0.2], [0.2], [0.2], [0.2], [-1.0]]),
|
||||
hit_pos_w=torch.tensor(
|
||||
[[[0.0, 0, 0]], [[0, 0, 0.05]], [[0, 0, 0]], [[7, 0, 0]], [[0, 0, 0]]]
|
||||
),
|
||||
normals_w=torch.tensor(
|
||||
[[[0.0, 0, 1]], [[0, 0, 1]], [[1, 0, 0]], [[0, 0, 1]], [[0, 0, 0]]]
|
||||
),
|
||||
)
|
||||
env = SimpleNamespace(
|
||||
scene={"forward_scan": SimpleNamespace(data=data)}, _multi_ring_floor_id=0
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
forward_depth(env)[:, 0], torch.tensor([0.05, 0.05, 0.05, 0.05, 1])
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
obstacle_proximity(env, floor_size=12), torch.tensor([0.0, 0.36, 0.36, 0.36, 0.0])
|
||||
)
|
||||
self.assertGreater(obstacle_proximity(env)[0], 0) # legacy unchanged
|
||||
|
||||
@unittest.skipUnless(os.environ.get("GO2_RUN_MULTI_SMOKE") == "1", "opt-in GPU smoke")
|
||||
def test_real_2env_1step_97_shape_floor_reward_and_cpu_parity(self):
|
||||
import mujoco
|
||||
import numpy as np
|
||||
import torch
|
||||
import warp as wp
|
||||
from mjlab.envs import ManagerBasedRlEnv
|
||||
from src.tasks.obstacle_avoidance.env_cfg import (
|
||||
apply_obstacle_configuration,
|
||||
unitree_go2_obstacle_env_cfg,
|
||||
)
|
||||
from src.tasks.obstacle_avoidance.mdp import (
|
||||
ForwardFanPatternCfg,
|
||||
floor_top_hits,
|
||||
obstacle_proximity,
|
||||
)
|
||||
from warp._src import context
|
||||
|
||||
if not hasattr(wp, "context"):
|
||||
wp.context = context
|
||||
custom = validate_task_config(
|
||||
OBSTACLE_TASK,
|
||||
{
|
||||
"terrainPreset": "plane",
|
||||
"sensorCfg": {"sensorMode": "multi_ring_raycast", "safetyDistance": 1},
|
||||
},
|
||||
42,
|
||||
)
|
||||
cfg = unitree_go2_obstacle_env_cfg()
|
||||
apply_obstacle_configuration(cfg, custom)
|
||||
cfg.scene.num_envs = 2
|
||||
env = ManagerBasedRlEnv(cfg, device="cuda:0")
|
||||
try:
|
||||
obs, _ = env.reset()
|
||||
m = env.sim.mj_model
|
||||
print(
|
||||
"FLOOR_DIAGNOSTIC",
|
||||
[
|
||||
(
|
||||
m.geom(i).name,
|
||||
int(m.geom_bodyid[i]),
|
||||
m.geom_pos[i].tolist(),
|
||||
m.geom_size[i].tolist(),
|
||||
m.geom_quat[i].tolist(),
|
||||
)
|
||||
for i in range(m.ngeom)
|
||||
if m.geom_group[i] == 0
|
||||
],
|
||||
)
|
||||
obs, reward, _, _, _ = env.step(torch.zeros((2, 12), device=env.device))
|
||||
self.assertEqual(tuple(obs["actor"].shape), (2, 97))
|
||||
self.assertTrue(torch.isfinite(obs["actor"]).all() and torch.isfinite(reward).all())
|
||||
scan = env.scene["forward_scan"].data
|
||||
self.assertTrue(floor_top_hits(scan, 12).any())
|
||||
self.assertTrue((obs["actor"][:, 63:95] < 1).any())
|
||||
torch.testing.assert_close(
|
||||
obstacle_proximity(env, safety_distance=1, floor_size=12),
|
||||
torch.zeros(2, device=env.device),
|
||||
)
|
||||
model = env.sim.mj_model
|
||||
data = mujoco.MjData(model)
|
||||
offsets, rays = ForwardFanPatternCfg(sensor_mode="multi_ring_raycast").generate_rays(
|
||||
None, "cpu"
|
||||
)
|
||||
for e in range(2):
|
||||
data.qpos[:] = env.sim.data.qpos[e].cpu().numpy()
|
||||
mujoco.mj_forward(model, data)
|
||||
body = model.body("robot/base_link").id
|
||||
rotation = data.xmat[body].reshape(3, 3)
|
||||
expected = []
|
||||
for o, d in zip(offsets.numpy(), rays.numpy(), strict=True):
|
||||
t = mujoco.mj_ray(
|
||||
model,
|
||||
data,
|
||||
data.xpos[body] + rotation @ o,
|
||||
rotation @ d,
|
||||
np.array([1, 0, 0, 0, 0, 0], dtype=np.uint8),
|
||||
1,
|
||||
-1,
|
||||
np.array([-1], dtype=np.int32),
|
||||
)
|
||||
expected.append(1 if t < 0 else min(1, t / 4))
|
||||
np.testing.assert_allclose(obs["actor"][e, 47:95].cpu(), expected, atol=2e-5)
|
||||
print(
|
||||
"MULTI_SMOKE: 2env x 1step actor=(2,97); CPU mj_ray all48 parity; "
|
||||
"floor obs retained, proximity=0; reward finite"
|
||||
)
|
||||
finally:
|
||||
env.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,223 @@
|
||||
"""Optional installed-mjlab checks; GPU smoke is opt-in, never a full training run."""
|
||||
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
SERVICE_ROOT = Path(__file__).resolve().parents[1]
|
||||
for root in (SERVICE_ROOT, SERVICE_ROOT / "rl"):
|
||||
sys.path.insert(0, str(root))
|
||||
|
||||
HAS_MJLAB = importlib.util.find_spec("mjlab") is not None
|
||||
|
||||
|
||||
@unittest.skipUnless(HAS_MJLAB, "mjlab is not installed in this Python")
|
||||
class ObstacleContractTest(unittest.TestCase):
|
||||
def test_training_config_file_is_revalidated(self):
|
||||
import json
|
||||
import tempfile
|
||||
|
||||
from scripts.train import _load_task_config
|
||||
from task_config import OBSTACLE_TASK, TaskConfigError, validate_task_config
|
||||
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
source = Path(directory) / "training_config.json"
|
||||
custom = validate_task_config(OBSTACLE_TASK, {"sensorCfg": {"fov": 60}}, 42)
|
||||
source.write_text(json.dumps(custom))
|
||||
self.assertEqual(_load_task_config(OBSTACLE_TASK, str(source), 42), custom)
|
||||
with self.assertRaises(ValueError):
|
||||
_load_task_config(OBSTACLE_TASK, str(source), 43)
|
||||
custom["sensorCfg"]["maxDistance"] = -1
|
||||
source.write_text(json.dumps(custom))
|
||||
with self.assertRaises(TaskConfigError):
|
||||
_load_task_config(OBSTACLE_TASK, str(source), 42)
|
||||
source.write_text("x" * (128 * 1024 + 1))
|
||||
with self.assertRaises(ValueError):
|
||||
_load_task_config(OBSTACLE_TASK, str(source), 42)
|
||||
|
||||
def test_pattern_order_normalization_and_body_offsets(self):
|
||||
import torch
|
||||
from src.tasks.obstacle_avoidance.mdp import ForwardFanPatternCfg, forward_depth
|
||||
|
||||
offsets, directions = ForwardFanPatternCfg(fov=90).generate_rays(None, "cpu")
|
||||
self.assertEqual(tuple(directions.shape), (32, 3))
|
||||
torch.testing.assert_close(offsets, torch.tensor([0.3, 0, 0.05]).repeat(32, 1))
|
||||
torch.testing.assert_close(directions.norm(dim=1), torch.ones(32))
|
||||
self.assertAlmostEqual(directions[0, 1].item(), -(0.5**0.5), places=6)
|
||||
self.assertAlmostEqual(directions[-1, 1].item(), 0.5**0.5, places=6)
|
||||
scene = {
|
||||
"forward_scan": SimpleNamespace(
|
||||
data=SimpleNamespace(
|
||||
distances=torch.tensor([[-1, 0, 1, 4, 8]], dtype=torch.float32)
|
||||
)
|
||||
)
|
||||
}
|
||||
result = forward_depth(SimpleNamespace(scene=scene), max_distance=4)
|
||||
torch.testing.assert_close(result, torch.tensor([[1, 0, 0.25, 1, 1]]))
|
||||
|
||||
def test_navigation_matches_body_heading_and_arrival_stop(self):
|
||||
import torch
|
||||
from src.tasks.obstacle_avoidance.mdp import NavigationCommandCfg
|
||||
|
||||
data = SimpleNamespace(
|
||||
root_link_pos_w=torch.tensor([[-5.0, 0, 0.32]]),
|
||||
root_link_quat_w=torch.tensor([[1.0, 0, 0, 0]]),
|
||||
)
|
||||
|
||||
class Scene(dict):
|
||||
env_origins = torch.tensor([[-5.0, 0, 0]])
|
||||
|
||||
env = SimpleNamespace(
|
||||
num_envs=1, device="cpu", scene=Scene(robot=SimpleNamespace(data=data))
|
||||
)
|
||||
command = NavigationCommandCfg(resampling_time_range=(1e9, 1e9)).build(env)
|
||||
torch.testing.assert_close(command.command, torch.tensor([[0.6, 0, 0]]))
|
||||
self.assertAlmostEqual(command.errors()[1].item(), 10)
|
||||
data.root_link_quat_w[:] = torch.tensor([[0.5**0.5, 0, 0, 0.5**0.5]])
|
||||
self.assertAlmostEqual(command.command[0, 2].item(), -1)
|
||||
data.root_link_pos_w[0, 0] = 4.8
|
||||
torch.testing.assert_close(command.command, torch.zeros((1, 3)))
|
||||
|
||||
def test_navigation_reset_randomizes_safe_distant_pairs_and_heading(self):
|
||||
import torch
|
||||
from src.tasks.obstacle_avoidance.mdp import NavigationCommandCfg
|
||||
|
||||
class Robot:
|
||||
def __init__(self, count):
|
||||
self.data = SimpleNamespace(
|
||||
default_root_state=torch.tensor([[0, 0, 0.32] + [0] * 10] * count),
|
||||
root_link_pos_w=torch.zeros((count, 3)),
|
||||
root_link_quat_w=torch.zeros((count, 4)),
|
||||
)
|
||||
|
||||
def write_root_link_pose_to_sim(self, pose, env_ids):
|
||||
self.data.root_link_pos_w[env_ids] = pose[:, :3]
|
||||
self.data.root_link_quat_w[env_ids] = pose[:, 3:]
|
||||
|
||||
def write_root_link_velocity_to_sim(self, velocity, env_ids):
|
||||
self.velocity = velocity
|
||||
|
||||
class Scene(dict):
|
||||
env_origins = torch.zeros((64, 3))
|
||||
|
||||
robot = Robot(64)
|
||||
env = SimpleNamespace(num_envs=64, device="cpu", scene=Scene(robot=robot))
|
||||
command = NavigationCommandCfg(
|
||||
resampling_time_range=(1e9, 1e9),
|
||||
navigation_points=((-2, 0), (-1, 0), (0, 0), (1, 0), (2, 0)),
|
||||
component_starts=(0,),
|
||||
component_counts=(5,),
|
||||
fallback_pairs=(((-2, 0), (2, 0)),),
|
||||
min_goal_distance=2,
|
||||
).build(env)
|
||||
torch.manual_seed(7)
|
||||
command.sample_episode(torch.arange(64))
|
||||
self.assertTrue(
|
||||
((command.goals_w - robot.data.root_link_pos_w[:, :2]).norm(dim=1) >= 2).all()
|
||||
)
|
||||
self.assertGreater(torch.unique(robot.data.root_link_pos_w[:, :2], dim=0).shape[0], 1)
|
||||
self.assertGreater(torch.unique(command.goals_w, dim=0).shape[0], 1)
|
||||
torch.testing.assert_close(robot.data.root_link_quat_w.norm(dim=1), torch.ones(64))
|
||||
self.assertTrue((robot.velocity == 0).all())
|
||||
|
||||
def test_all_terrain_presets_compile_at_exported_world_coordinates(self):
|
||||
import mujoco
|
||||
import numpy as np
|
||||
from mjlab.terrains import TerrainGenerator
|
||||
from src.tasks.obstacle_avoidance.env_cfg import unitree_go2_obstacle_env_cfg
|
||||
from src.tasks.obstacle_avoidance.terrain import apply_terrain_configuration
|
||||
from task_config import (
|
||||
OBSTACLE_TASK,
|
||||
TERRAIN_PRESETS,
|
||||
build_terrain_layout,
|
||||
validate_task_config,
|
||||
)
|
||||
|
||||
for preset in (p for p in TERRAIN_PRESETS if p != "custom_boxes"):
|
||||
with self.subTest(preset=preset):
|
||||
custom = validate_task_config(OBSTACLE_TASK, {"terrainPreset": preset}, 7)
|
||||
cfg = unitree_go2_obstacle_env_cfg()
|
||||
apply_terrain_configuration(cfg, custom)
|
||||
generator = TerrainGenerator(cfg.scene.terrain.terrain_generator)
|
||||
spec = mujoco.MjSpec()
|
||||
generator.compile(spec)
|
||||
model = spec.compile()
|
||||
layout = build_terrain_layout(custom)
|
||||
np.testing.assert_allclose(model.geom_pos, [box["pos"] for box in layout["boxes"]])
|
||||
np.testing.assert_allclose(
|
||||
model.geom_size, [box["size"] for box in layout["boxes"]]
|
||||
)
|
||||
np.testing.assert_allclose(generator.terrain_origins[0, 0], [-5, 0, 0])
|
||||
np.testing.assert_allclose(model.geom_friction[:, 0], layout["friction"])
|
||||
self.assertTrue((model.geom_group == 0).all())
|
||||
|
||||
@unittest.skipUnless(os.environ.get("GO2_RUN_MJLAB_SMOKE") == "1", "opt-in GPU smoke")
|
||||
def test_real_environment_81_observation_and_mujoco_raycast_parity(self):
|
||||
import mujoco
|
||||
import numpy as np
|
||||
import torch
|
||||
import warp as wp
|
||||
from mjlab.envs import ManagerBasedRlEnv
|
||||
from src.tasks.obstacle_avoidance.env_cfg import (
|
||||
apply_obstacle_configuration,
|
||||
unitree_go2_obstacle_env_cfg,
|
||||
)
|
||||
from src.tasks.obstacle_avoidance.mdp import ForwardFanPatternCfg
|
||||
from task_config import JOINT_NAMES, OBSTACLE_TASK, validate_task_config
|
||||
from warp._src import context
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
self.skipTest("CUDA is unavailable")
|
||||
if not hasattr(wp, "context"):
|
||||
wp.context = context
|
||||
cfg = unitree_go2_obstacle_env_cfg()
|
||||
apply_obstacle_configuration(
|
||||
cfg, validate_task_config(OBSTACLE_TASK, {}, 42), randomize_navigation=False
|
||||
)
|
||||
cfg.scene.num_envs = 2
|
||||
env = ManagerBasedRlEnv(cfg, device="cuda:0")
|
||||
try:
|
||||
obs, _ = env.reset()
|
||||
self.assertEqual(tuple(obs["actor"].shape), (2, 81))
|
||||
self.assertEqual(list(env.scene["robot"].joint_names), JOINT_NAMES)
|
||||
np.testing.assert_allclose(env.scene.env_origins.cpu(), [[-5, 0, 0]] * 2)
|
||||
np.testing.assert_allclose(
|
||||
env.scene["robot"].data.root_link_pos_w.cpu(), [[-5, 0, 0.32]] * 2, atol=1e-6
|
||||
)
|
||||
for _ in range(3):
|
||||
obs, reward, terminated, truncated, _ = env.step(
|
||||
torch.zeros((2, 12), device=env.device)
|
||||
)
|
||||
self.assertTrue(torch.isfinite(obs["actor"]).all())
|
||||
self.assertTrue(torch.isfinite(reward).all())
|
||||
self.assertFalse(terminated.any() or truncated.any())
|
||||
model = env.sim.mj_model
|
||||
data = mujoco.MjData(model)
|
||||
data.qpos[:] = env.sim.data.qpos[0].cpu().numpy()
|
||||
mujoco.mj_forward(model, data)
|
||||
body = model.body("robot/base_link").id
|
||||
rotation = data.xmat[body].reshape(3, 3)
|
||||
offsets, directions = ForwardFanPatternCfg().generate_rays(None, "cpu")
|
||||
expected = []
|
||||
for offset, direction in zip(offsets.numpy(), directions.numpy(), strict=True):
|
||||
distance = mujoco.mj_ray(
|
||||
model,
|
||||
data,
|
||||
data.xpos[body] + rotation @ offset,
|
||||
rotation @ direction,
|
||||
np.array([1, 0, 0, 0, 0, 0], dtype=np.uint8),
|
||||
1,
|
||||
-1,
|
||||
np.array([-1], dtype=np.int32),
|
||||
)
|
||||
expected.append(1 if distance < 0 else min(1, distance / 4))
|
||||
np.testing.assert_allclose(obs["actor"][0, 47:79].cpu(), expected, atol=2e-5)
|
||||
finally:
|
||||
env.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,512 @@
|
||||
"""Obstacle-specific schema, orchestration, objective math and opt-in real rollout."""
|
||||
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
import threading
|
||||
import unittest
|
||||
from copy import deepcopy
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
SERVICE = Path(__file__).resolve().parents[1]
|
||||
for root in (SERVICE, SERVICE / "rl"):
|
||||
sys.path.insert(0, str(root))
|
||||
|
||||
from task_config import build_terrain_layout, validate_task_config # noqa: E402
|
||||
from tuning import obstacle_scoring as scoring # noqa: E402
|
||||
from tuning.advisor import DeepSeekAdvisor # noqa: E402
|
||||
from tuning.manager import TuningError, TuningManager # noqa: E402
|
||||
from tuning.process import GpuLease, ResourceBusyError # noqa: E402
|
||||
from tuning.schema import ( # noqa: E402
|
||||
OBSTACLE_TASK,
|
||||
RewardConfigError,
|
||||
apply_reward_configuration,
|
||||
base_configuration,
|
||||
merge_proposal,
|
||||
validate_configuration,
|
||||
validate_constraints,
|
||||
validate_proposal,
|
||||
)
|
||||
from tuning.scoring import EvaluationError # noqa: E402
|
||||
|
||||
|
||||
def sample(**kw):
|
||||
return (
|
||||
dict(
|
||||
distance=1.0,
|
||||
clearance=0.5,
|
||||
action_delta=0.0,
|
||||
ray_hit=0.0,
|
||||
collision=0.0,
|
||||
fall=0.0,
|
||||
terminal=0.0,
|
||||
**{},
|
||||
)
|
||||
| kw
|
||||
)
|
||||
|
||||
|
||||
def evaluation(custom, n=2):
|
||||
metrics = scoring.score_trajectory([sample()] * scoring.STEPS)
|
||||
return {
|
||||
"protocol": scoring.protocol(custom, n),
|
||||
"metrics": metrics,
|
||||
"seedMetrics": [
|
||||
{"seed": seed, "episodes": n, "rolloutSteps": scoring.STEPS, "metrics": metrics}
|
||||
for seed in scoring.SEEDS
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
class ObstacleSchemaTest(unittest.TestCase):
|
||||
def test_ranges_task_isolation_nonfinite_and_patch_no_mutation(self):
|
||||
base = base_configuration(OBSTACLE_TASK)
|
||||
original = deepcopy(base)
|
||||
for section, key, low, high in (
|
||||
("weights", "avoidance_weight", 0.5, 5),
|
||||
("weights", "collision_penalty", -10, -0.5),
|
||||
("weights", "action_smoothness", -0.05, -0.001),
|
||||
("params", "target_velocity", 0.3, 1.2),
|
||||
):
|
||||
for value in (low, high):
|
||||
changed = deepcopy(base)
|
||||
changed[section][key] = value
|
||||
self.assertEqual(validate_configuration(changed, OBSTACLE_TASK), changed)
|
||||
for value in (low - 0.00001, high + 0.00001, float("nan"), float("inf"), True):
|
||||
with self.subTest(key=key, value=value), self.assertRaises(RewardConfigError):
|
||||
validate_proposal({section: {key: value}}, base, task_id=OBSTACLE_TASK)
|
||||
for proposal in (
|
||||
{"weights": {"pose": 1.1}},
|
||||
{"weights": {"target_velocity": 0.8}},
|
||||
{"params": {"seed": 101}},
|
||||
{"unknown": {}},
|
||||
{"weights": {"avoidance_weight": 4.01}},
|
||||
):
|
||||
with self.assertRaises(RewardConfigError):
|
||||
validate_proposal(proposal, base, task_id=OBSTACLE_TASK)
|
||||
with self.assertRaises(RewardConfigError):
|
||||
validate_configuration(base)
|
||||
with self.assertRaises(RewardConfigError):
|
||||
validate_configuration(base_configuration(), OBSTACLE_TASK)
|
||||
with self.assertRaises(RewardConfigError):
|
||||
validate_constraints({"weights.pose": {"kind": "fixed", "value": 1}}, OBSTACLE_TASK)
|
||||
candidate = merge_proposal(
|
||||
base, {"params": {"target_velocity": 1.2}}, task_id=OBSTACLE_TASK
|
||||
)
|
||||
self.assertEqual(candidate["params"]["target_velocity"], 1.2)
|
||||
self.assertEqual(base, original)
|
||||
|
||||
def test_real_training_config_mapping_and_command(self):
|
||||
import torch
|
||||
from scripts.train import (
|
||||
_configure_task_and_rewards,
|
||||
_load_reward_config,
|
||||
_load_task_config,
|
||||
)
|
||||
from src.tasks.obstacle_avoidance.env_cfg import (
|
||||
apply_obstacle_configuration,
|
||||
unitree_go2_obstacle_env_cfg,
|
||||
)
|
||||
|
||||
base = base_configuration(OBSTACLE_TASK)
|
||||
base["weights"].update(avoidance_weight=3, collision_penalty=-7, action_smoothness=-0.02)
|
||||
base["params"]["target_velocity"] = 0.9
|
||||
custom = validate_task_config(OBSTACLE_TASK, {"sensorCfg": {"fov": 60}}, 42)
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
reward, task = Path(directory) / "reward.json", Path(directory) / "task.json"
|
||||
reward.write_text(json.dumps(base))
|
||||
task.write_text(json.dumps(custom))
|
||||
cfg = unitree_go2_obstacle_env_cfg()
|
||||
apply_obstacle_configuration(cfg, _load_task_config(OBSTACLE_TASK, str(task), 42))
|
||||
apply_reward_configuration(
|
||||
cfg, _load_reward_config(str(reward), None, OBSTACLE_TASK), OBSTACLE_TASK
|
||||
)
|
||||
deployment, checked = _configure_task_and_rewards(
|
||||
OBSTACLE_TASK,
|
||||
SimpleNamespace(
|
||||
env=unitree_go2_obstacle_env_cfg(),
|
||||
task_config=str(task),
|
||||
reward_config=str(reward),
|
||||
reward_config_json=None,
|
||||
agent=SimpleNamespace(seed=42),
|
||||
),
|
||||
)
|
||||
self.assertEqual(checked, base)
|
||||
self.assertEqual(deployment["navigation"]["speed"], 0.9)
|
||||
self.assertEqual(deployment["sensorCfg"]["avoidanceWeight"], 3)
|
||||
self.assertEqual(cfg.rewards["obstacle_proximity"].weight, -3)
|
||||
self.assertEqual(cfg.rewards["obstacle_collision"].weight, -7)
|
||||
self.assertIs(
|
||||
cfg.rewards["obstacle_collision"].func, cfg.terminations["illegal_contact"].func
|
||||
)
|
||||
self.assertEqual(cfg.rewards["action_rate_l2"].weight, -0.02)
|
||||
self.assertNotIn("target_velocity", cfg.rewards)
|
||||
self.assertEqual(cfg.commands["twist"].speed, 0.9)
|
||||
|
||||
class Scene(dict):
|
||||
env_origins = torch.tensor([[-5.0, 0, 0]])
|
||||
|
||||
scene = Scene(
|
||||
robot=SimpleNamespace(
|
||||
data=SimpleNamespace(
|
||||
root_link_pos_w=torch.tensor([[-5.0, 0, 0.32]]),
|
||||
root_link_quat_w=torch.tensor([[1.0, 0, 0, 0]]),
|
||||
)
|
||||
)
|
||||
)
|
||||
command = cfg.commands["twist"].build(
|
||||
SimpleNamespace(num_envs=1, device="cpu", scene=scene)
|
||||
)
|
||||
self.assertAlmostEqual(command.command[0, 0].item(), 0.9, places=6)
|
||||
|
||||
def test_mock_deepseek_task_context_no_network(self):
|
||||
advisor = DeepSeekAdvisor()
|
||||
output = SimpleNamespace(
|
||||
weights={"avoidance_weight": 2.2},
|
||||
params={"target_velocity": 0.7},
|
||||
rationale="绕行与目标导航",
|
||||
expected_impact={},
|
||||
confidence=0.8,
|
||||
)
|
||||
fake = SimpleNamespace(run_sync=lambda prompt: SimpleNamespace(output=output))
|
||||
with patch.object(advisor, "_agent", return_value=fake):
|
||||
result = advisor.propose({"task": OBSTACLE_TASK}, base_configuration(OBSTACLE_TASK))
|
||||
self.assertEqual(result["patch"]["params"], {"target_velocity": 0.7})
|
||||
|
||||
|
||||
class ObstacleScoringTest(unittest.TestCase):
|
||||
def test_hand_calculated_scores_and_first_terminal(self):
|
||||
samples = [
|
||||
sample(clearance=0.25, action_delta=0.5),
|
||||
sample(distance=0.4, clearance=0.25, action_delta=0.5),
|
||||
]
|
||||
metrics = scoring.score_trajectory(samples, 2)
|
||||
self.assertEqual(metrics["success"], 1)
|
||||
self.assertEqual(metrics["time"], 0)
|
||||
self.assertEqual(metrics["clearance"], 0.5)
|
||||
self.assertEqual(metrics["smooth"], 0.5)
|
||||
self.assertAlmostEqual(scoring.score_evaluation(metrics)["score"], 0.65)
|
||||
failed = scoring.score_trajectory([sample(distance=0.1, terminal=1, fall=1)], 1000)
|
||||
self.assertEqual((failed["success"], failed["time"], failed["clearance"]), (0, 0, 0))
|
||||
contact = scoring.score_trajectory([sample(distance=0.1, terminal=1, collision=1)], 1000)
|
||||
self.assertEqual((contact["arrival_rate"], contact["success"], contact["time"]), (1, 0, 0))
|
||||
with self.assertRaises(EvaluationError):
|
||||
scoring.score_trajectory([sample()], 1000)
|
||||
with self.assertRaises(EvaluationError):
|
||||
scoring.score_trajectory([sample(terminal=1), sample()], 2)
|
||||
with self.assertRaises(EvaluationError):
|
||||
scoring.score_trajectory([sample(clearance=float("nan"))], 1)
|
||||
|
||||
def test_recorder_excludes_floor_and_keeps_terminal_before_reset(self):
|
||||
import torch
|
||||
from scripts.evaluate_obstacle import FirstEpisodeRecorder
|
||||
|
||||
class Scene(dict):
|
||||
env_origins = torch.zeros((2, 3))
|
||||
|
||||
robot = SimpleNamespace(
|
||||
root_link_pos_w=torch.tensor([[0.0, 0, 0.32], [0.0, 0, 0.1]]),
|
||||
projected_gravity_b=torch.tensor([[0.0, 0, -1.0], [0.0, 1.0, 0.0]]),
|
||||
)
|
||||
scene = Scene(
|
||||
robot=SimpleNamespace(data=robot),
|
||||
nonfoot_ground_touch=SimpleNamespace(
|
||||
data=SimpleNamespace(force_history=torch.zeros((2, 1, 4, 3)))
|
||||
),
|
||||
forward_scan=SimpleNamespace(
|
||||
data=SimpleNamespace(distances=torch.ones((2, 32))),
|
||||
cfg=SimpleNamespace(max_distance=4),
|
||||
),
|
||||
)
|
||||
command = SimpleNamespace(errors=lambda: (None, torch.tensor([0.2, 0.2]), None))
|
||||
env = SimpleNamespace(
|
||||
scene=scene,
|
||||
device="cpu",
|
||||
num_envs=2,
|
||||
action_manager=SimpleNamespace(action=torch.ones((2, 12)) * 0.5),
|
||||
command_manager=SimpleNamespace(get_term=lambda name: command),
|
||||
termination_manager=SimpleNamespace(compute=lambda: torch.ones(2, dtype=torch.bool)),
|
||||
)
|
||||
floor = {"pos": [0, 0, -0.1], "size": [6, 6, 0.1]}
|
||||
recorder = FirstEpisodeRecorder(env, {"spawn": [0, 0, 0.32], "boxes": [floor]})
|
||||
env.termination_manager.compute()
|
||||
robot.root_link_pos_w[:] = torch.tensor([0.0, 0, 0.32]) # Simulated auto-reset overwrite.
|
||||
robot.projected_gravity_b[:] = torch.tensor([0.0, 0, -1.0])
|
||||
env.termination_manager.compute()
|
||||
self.assertEqual([len(s) for s in recorder.samples], [1, 1])
|
||||
upright, lying = [scoring.score_trajectory(s, 2) for s in recorder.samples]
|
||||
self.assertEqual(upright["clearance"], 0.5) # 1 valid safe step / 2, not floor distance 0.
|
||||
self.assertEqual(upright["success"], 1)
|
||||
self.assertEqual(lying["clearance"], 0)
|
||||
self.assertEqual(lying["success"], 0)
|
||||
self.assertEqual(lying["time"], 0)
|
||||
self.assertEqual(lying["fall_rate"], 1)
|
||||
|
||||
def test_threshold_boundary_and_fail_closed(self):
|
||||
base = scoring.score_trajectory([sample(distance=0.1)], 1)
|
||||
base.update(success=0.5, fall_rate=0.1, no_fall=0.9)
|
||||
current = dict(base, success=0.48, fall_rate=0.12, no_fall=0.88)
|
||||
self.assertTrue(scoring.score_evaluation(current, base)["eligible"])
|
||||
for changed in (
|
||||
dict(current, success=0.479999),
|
||||
dict(current, fall_rate=0.120001, no_fall=0.879999),
|
||||
):
|
||||
self.assertFalse(scoring.score_evaluation(changed, base)["eligible"])
|
||||
self.assertEqual(scoring.score_evaluation(changed, base)["score"], -1)
|
||||
custom = validate_task_config(OBSTACLE_TASK, {}, 42)
|
||||
expected = scoring.protocol(custom, 2)
|
||||
valid = evaluation(custom)
|
||||
scoring.validate_evaluation(valid, expected)
|
||||
for mutate in (
|
||||
lambda v: v["seedMetrics"].pop(),
|
||||
lambda v: v["metrics"].pop("success"),
|
||||
lambda v: v["protocol"].update(stepsPerSeed=999),
|
||||
lambda v: v["seedMetrics"][0].update(episodes=1),
|
||||
lambda v: v["metrics"].update(smooth=float("inf")),
|
||||
):
|
||||
bad = deepcopy(valid)
|
||||
mutate(bad)
|
||||
with self.assertRaises(EvaluationError):
|
||||
scoring.validate_evaluation(bad, expected)
|
||||
|
||||
def test_fixed_three_seed_and_custom_authoritative_map(self):
|
||||
custom = validate_task_config(OBSTACLE_TASK, {}, 42)
|
||||
scenes = scoring.evaluation_scenarios(custom)
|
||||
self.assertEqual([s["seed"] for s in scenes], [101, 202, 303])
|
||||
self.assertNotEqual(scenes[0]["terrain"], scenes[1]["terrain"])
|
||||
layout = build_terrain_layout(custom)
|
||||
layout["approximation"] = True
|
||||
custom = validate_task_config(
|
||||
OBSTACLE_TASK, {"terrainPreset": "custom_boxes", "customTerrainBoxes": layout}, 42
|
||||
)
|
||||
fixed = scoring.protocol(custom, 2)
|
||||
self.assertEqual(fixed["sceneMode"], "fixed-custom-map")
|
||||
self.assertEqual([s["terrain"] for s in fixed["scenarios"]], [layout] * 3)
|
||||
self.assertEqual(fixed, scoring.protocol(deepcopy(custom), 2))
|
||||
|
||||
|
||||
class ObstacleManagerTest(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.temp = tempfile.TemporaryDirectory()
|
||||
root = Path(self.temp.name)
|
||||
self.manager = TuningManager(SERVICE / "rl", sys.executable, root, GpuLease())
|
||||
payload = {
|
||||
"taskId": OBSTACLE_TASK,
|
||||
"mode": "approval",
|
||||
"evalNumEnvs": 2,
|
||||
"numEnvs": 2,
|
||||
"initialIterations": 1,
|
||||
"middleIterations": 1,
|
||||
"finalIterations": 1,
|
||||
"sensorCfg": {"fov": 60, "avoidanceWeight": 3},
|
||||
}
|
||||
mode, config, objective, fallback = self.manager.parse_create(payload)
|
||||
self.session = self.manager.storage.create_session(mode, config, objective, fallback)
|
||||
self.trial = self.manager.storage.create_trial(
|
||||
self.session["id"],
|
||||
0,
|
||||
0,
|
||||
1,
|
||||
self.manager._base_configuration(self.session),
|
||||
None,
|
||||
"trial-000-rung-0",
|
||||
)
|
||||
|
||||
def tearDown(self):
|
||||
self.manager.shutdown()
|
||||
self.temp.cleanup()
|
||||
|
||||
def test_argv_json_real_train_and_eval_ingress(self):
|
||||
commands = []
|
||||
custom = self.session["config"]["taskConfig"]
|
||||
|
||||
def run(session_id, command, cwd, environment, log_path):
|
||||
commands.append(command)
|
||||
task_path = Path(command[command.index("--task-config") + 1])
|
||||
reward_path = Path(command[command.index("--reward-config") + 1])
|
||||
from scripts.train import _load_reward_config, _load_task_config
|
||||
|
||||
self.assertEqual(_load_task_config(OBSTACLE_TASK, str(task_path), 42), custom)
|
||||
self.assertEqual(
|
||||
_load_reward_config(str(reward_path), None, OBSTACLE_TASK),
|
||||
self.trial["rewardConfig"],
|
||||
)
|
||||
if "scripts/train.py" in command:
|
||||
(log_path.parent / "model_0.pt").write_bytes(b"mock-checkpoint")
|
||||
(log_path.parent / "policy.onnx").write_bytes(b"mock-onnx")
|
||||
else:
|
||||
Path(command[command.index("--output") + 1]).write_text(
|
||||
json.dumps(evaluation(custom))
|
||||
)
|
||||
return 0
|
||||
|
||||
self.manager.cancel_events[self.session["id"]] = threading.Event()
|
||||
with (
|
||||
patch.object(self.manager, "_run_command", side_effect=run),
|
||||
patch.object(self.manager.studies, "record", return_value=0),
|
||||
):
|
||||
result = self.manager._execute_trial(self.session, self.trial)
|
||||
self.assertEqual(result["state"], "completed")
|
||||
self.assertIn("--steps-per-seed=1000", commands[1])
|
||||
self.assertIn(OBSTACLE_TASK, commands[0])
|
||||
context = self.manager._proposal_context(self.session)
|
||||
self.assertIn("32", context["taskContext"])
|
||||
self.assertEqual(set(context["allowlist"]["params"]), {"target_velocity"})
|
||||
|
||||
def test_cas_stale_guardrails_and_mode_roundtrip(self):
|
||||
sid = self.session["id"]
|
||||
constraint = {"params.target_velocity": {"kind": "range", "min": 0.4, "max": 0.8}}
|
||||
self.manager.set_constraints(sid, {"revision": 0, "constraints": constraint})
|
||||
before = self.manager.detail(sid)
|
||||
with self.assertRaises(ResourceBusyError):
|
||||
self.manager.set_constraints(sid, {"revision": 0, "constraints": {}})
|
||||
self.assertEqual(self.manager.detail(sid), before)
|
||||
for bad in (
|
||||
{"weights.pose": {"kind": "fixed", "value": 1}},
|
||||
{"params.target_velocity": {"kind": "fixed", "value": float("nan")}},
|
||||
):
|
||||
with self.assertRaises(RewardConfigError):
|
||||
self.manager.set_constraints(sid, {"revision": 1, "constraints": bad})
|
||||
self.assertEqual(self.manager.detail(sid), before)
|
||||
self.manager.storage.update_session(sid, state="awaiting_approval")
|
||||
proposal = self.manager.storage.create_proposal(
|
||||
sid, self.trial["id"], {"params": {"target_velocity": 0.7}}, "test", {}, 0.8, "agent"
|
||||
)
|
||||
self.assertEqual(self.manager.set_mode(sid, {"mode": "automatic"})["mode"], "automatic")
|
||||
self.assertEqual(self.manager.storage.get_proposal(proposal["id"])["state"], "approved")
|
||||
self.assertEqual(self.manager.set_mode(sid, {"mode": "approval"})["mode"], "approval")
|
||||
with self.assertRaises(RewardConfigError):
|
||||
validate_proposal(
|
||||
{"params": {"target_velocity": 0.9}},
|
||||
self.trial["rewardConfig"],
|
||||
constraint,
|
||||
OBSTACLE_TASK,
|
||||
)
|
||||
|
||||
def test_protocol_fields_cannot_be_changed(self):
|
||||
for values in (
|
||||
{"evalSteps": 999},
|
||||
{"seeds": [1, 2, 3]},
|
||||
{"objectiveWeights": {}},
|
||||
{"taskConfig": {"unknown": 1}},
|
||||
{"taskConfig": {"seed": True}},
|
||||
{"seed": 1, "taskConfig": {"seed": True}},
|
||||
):
|
||||
with self.assertRaises(TuningError):
|
||||
self.manager.parse_create({"taskId": OBSTACLE_TASK, **values})
|
||||
|
||||
|
||||
class IsolatedEvaluationTest(unittest.TestCase):
|
||||
def test_spawn_results_and_failure_cancel_missing_seed_fail_closed(self):
|
||||
import subprocess
|
||||
|
||||
from scripts.evaluate import EvaluateConfig
|
||||
from scripts.evaluate_obstacle import evaluate_isolated_seeds
|
||||
|
||||
custom = validate_task_config(OBSTACLE_TASK, {}, 42)
|
||||
fixed = scoring.protocol(custom, 2)
|
||||
reward = base_configuration(OBSTACLE_TASK)
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
checkpoint = Path(directory) / "model.pt"
|
||||
checkpoint.write_bytes(b"test checkpoint snapshot")
|
||||
cfg = EvaluateConfig(checkpoint=str(checkpoint), output="unused", num_envs=2)
|
||||
|
||||
def run(command, **kwargs):
|
||||
self.assertFalse(kwargs["shell"])
|
||||
self.assertEqual(command[0], sys.executable)
|
||||
request = json.loads(Path(command[2]).read_text())
|
||||
result = evaluation(custom)["seedMetrics"][0]
|
||||
result["seed"] = request["scenario"]["seed"]
|
||||
result["checkpointSha256"] = request["checkpointSha256"]
|
||||
Path(command[3]).write_text(json.dumps(result))
|
||||
|
||||
with patch("scripts.evaluate_obstacle.subprocess.run", side_effect=run) as mocked:
|
||||
results = evaluate_isolated_seeds(OBSTACLE_TASK, cfg, fixed, reward)
|
||||
self.assertEqual([r["seed"] for r in results], [101, 202, 303])
|
||||
self.assertEqual(mocked.call_count, 3)
|
||||
for failure in (
|
||||
subprocess.CalledProcessError(1, "worker"),
|
||||
subprocess.TimeoutExpired("worker", 1200),
|
||||
KeyboardInterrupt(),
|
||||
):
|
||||
with patch(
|
||||
"scripts.evaluate_obstacle.subprocess.run", side_effect=failure
|
||||
) as mocked:
|
||||
with self.assertRaises(type(failure)):
|
||||
evaluate_isolated_seeds(OBSTACLE_TASK, cfg, fixed, reward)
|
||||
self.assertEqual(mocked.call_count, 1)
|
||||
with (
|
||||
patch("scripts.evaluate_obstacle.subprocess.run"),
|
||||
self.assertRaises(FileNotFoundError),
|
||||
):
|
||||
evaluate_isolated_seeds(OBSTACLE_TASK, cfg, fixed, reward)
|
||||
|
||||
def wrong_seed(command, **kwargs):
|
||||
run(command, **kwargs)
|
||||
path = Path(command[3])
|
||||
result = json.loads(path.read_text())
|
||||
result["seed"] = 0
|
||||
path.write_text(json.dumps(result))
|
||||
|
||||
with (
|
||||
patch("scripts.evaluate_obstacle.subprocess.run", side_effect=wrong_seed),
|
||||
self.assertRaises(EvaluationError),
|
||||
):
|
||||
evaluate_isolated_seeds(OBSTACLE_TASK, cfg, fixed, reward)
|
||||
|
||||
|
||||
@unittest.skipUnless(
|
||||
os.environ.get("GO2_RUN_TUNING_SMOKE") == "1", "opt-in real GPU checkpoint rollout"
|
||||
)
|
||||
class ObstacleRolloutSmoke(unittest.TestCase):
|
||||
def test_actual_checkpoint_policy_statistics_and_first_terminal(self):
|
||||
from dataclasses import asdict
|
||||
|
||||
import src.tasks # noqa: F401
|
||||
import torch
|
||||
from mjlab.envs import ManagerBasedRlEnv
|
||||
from mjlab.rl import RslRlVecEnvWrapper
|
||||
from mjlab.tasks.registry import load_rl_cfg, load_runner_cls
|
||||
from scripts.evaluate import EvaluateConfig
|
||||
from scripts.evaluate_obstacle import FirstEpisodeRecorder, configure_seed
|
||||
|
||||
cfg = EvaluateConfig(
|
||||
checkpoint=os.environ["GO2_TUNING_CHECKPOINT"],
|
||||
output="/tmp/unused.json",
|
||||
num_envs=2,
|
||||
device="cuda:0",
|
||||
)
|
||||
custom = validate_task_config(OBSTACLE_TASK, {}, 42)
|
||||
scenario = scoring.evaluation_scenarios(custom)[0]
|
||||
env_cfg = configure_seed(OBSTACLE_TASK, cfg, scenario, base_configuration(OBSTACLE_TASK))
|
||||
env_cfg.episode_length_s = 0.02 # Test-only one-step timeout exercises auto-reset capture.
|
||||
env = ManagerBasedRlEnv(env_cfg, device="cuda:0")
|
||||
agent_cfg = load_rl_cfg(OBSTACLE_TASK)
|
||||
wrapped = RslRlVecEnvWrapper(env, clip_actions=agent_cfg.clip_actions)
|
||||
try:
|
||||
runner = load_runner_cls(OBSTACLE_TASK)(
|
||||
wrapped, asdict(agent_cfg), log_dir=None, device="cuda:0"
|
||||
)
|
||||
runner.load(
|
||||
cfg.checkpoint, load_cfg={"actor": True}, strict=True, map_location="cuda:0"
|
||||
)
|
||||
saved = torch.load(cfg.checkpoint, weights_only=False)["actor_state_dict"]
|
||||
actual = runner.alg.actor.state_dict()
|
||||
normalizers = [k for k in saved if "normaliz" in k]
|
||||
self.assertTrue(normalizers)
|
||||
for key in normalizers:
|
||||
torch.testing.assert_close(actual[key], saved[key].to(actual[key].device))
|
||||
obs, _ = env.reset(seed=101)
|
||||
recorder = FirstEpisodeRecorder(env, scenario["terrain"])
|
||||
with torch.inference_mode():
|
||||
policy = runner.get_inference_policy(device="cuda:0")
|
||||
for _ in range(2):
|
||||
obs, _, _, _ = wrapped.step(policy(obs))
|
||||
self.assertEqual([len(s) for s in recorder.samples], [1, 1])
|
||||
self.assertTrue(all(s[0]["terminal"] for s in recorder.samples))
|
||||
scoring.validate_metrics(recorder.metrics(horizon=2))
|
||||
finally:
|
||||
wrapped.close()
|
||||
@@ -0,0 +1,475 @@
|
||||
"""CPU transfer regression; real source and <=4env/1iteration GPU checks are opt-in."""
|
||||
|
||||
import copy
|
||||
import importlib.util
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
from dataclasses import asdict, replace
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[1]
|
||||
for root in (ROOT, ROOT / "rl"):
|
||||
sys.path.insert(0, str(root))
|
||||
HAS_STACK = all(importlib.util.find_spec(m) is not None for m in ("mjlab", "torch", "onnxruntime"))
|
||||
|
||||
|
||||
@unittest.skipUnless(HAS_STACK, "installed RL + CPU ORT stack required")
|
||||
class PretrainedTest(unittest.TestCase):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
import torch
|
||||
|
||||
torch.set_num_threads(1)
|
||||
from pretrained import ValidatedSource, make_reference_actor
|
||||
from src.tasks.velocity.config.go2.env_cfgs import unitree_go2_flat_env_cfg
|
||||
from src.tasks.velocity.config.go2.rl_cfg import unitree_go2_ppo_runner_cfg
|
||||
|
||||
torch.manual_seed(42)
|
||||
actor = make_reference_actor()
|
||||
with torch.no_grad():
|
||||
actor.obs_normalizer._mean.uniform_(-0.2, 0.2)
|
||||
actor.obs_normalizer._var.uniform_(0.1, 1.2)
|
||||
actor.obs_normalizer._std.copy_(actor.obs_normalizer._var.sqrt())
|
||||
actor.obs_normalizer.count.fill_(983138304)
|
||||
cls.source = ValidatedSource(actor.state_dict(), {"source_id": "synthetic"})
|
||||
cls.env_cfg = asdict(unitree_go2_flat_env_cfg())
|
||||
cls.agent_cfg = asdict(unitree_go2_ppo_runner_cfg())
|
||||
|
||||
def test_extension_preserves_actions_and_learns_new_columns(self):
|
||||
import torch
|
||||
from pretrained import comparison_observations, make_reference_actor, warm_start_actor
|
||||
from tensordict import TensorDict
|
||||
|
||||
source_actor = make_reference_actor().eval()
|
||||
source_actor.load_state_dict(self.source.actor_state)
|
||||
base = comparison_observations()
|
||||
expected = source_actor.mlp(source_actor.obs_normalizer(base)).detach()
|
||||
for dim in (47, 81, 97):
|
||||
with self.subTest(dim=dim):
|
||||
actor = make_reference_actor(dim)
|
||||
warm_start_actor(actor, self.source)
|
||||
for key in ("_mean", "_var", "_std"):
|
||||
value = getattr(actor.obs_normalizer, key)
|
||||
self.assertTrue(
|
||||
torch.equal(value[:, :47], self.source.actor_state[f"obs_normalizer.{key}"])
|
||||
)
|
||||
self.assertTrue(
|
||||
torch.equal(
|
||||
value[:, 47:],
|
||||
torch.full_like(value[:, 47:], 0 if key == "_mean" else 1),
|
||||
)
|
||||
)
|
||||
self.assertTrue(
|
||||
torch.equal(
|
||||
actor.obs_normalizer.count, self.source.actor_state["obs_normalizer.count"]
|
||||
)
|
||||
)
|
||||
self.assertEqual(actor.mlp[0].weight[:, 47:].count_nonzero().item(), 0)
|
||||
self.assertTrue(actor.mlp[0].weight.requires_grad)
|
||||
extra = torch.rand(len(base), dim - 47) * 2 - 1
|
||||
x = torch.cat((base, extra), dim=1)
|
||||
torch.testing.assert_close(
|
||||
actor.mlp(actor.obs_normalizer(x)), expected, atol=3e-6, rtol=3e-6
|
||||
)
|
||||
initial_error = (actor.mlp(actor.obs_normalizer(x)) - expected).abs().max().item()
|
||||
x[:, 47:] *= -1
|
||||
torch.testing.assert_close(
|
||||
actor.mlp(actor.obs_normalizer(x)), expected, atol=3e-6, rtol=3e-6
|
||||
)
|
||||
before = actor.obs_normalizer._mean.clone()
|
||||
obs = TensorDict({"actor": x}, batch_size=[len(base)])
|
||||
actor.update_normalization(obs)
|
||||
self.assertEqual(actor.obs_normalizer.count.item(), 983138304 + len(base))
|
||||
self.assertLess(
|
||||
(actor.obs_normalizer._mean[:, :47] - before[:, :47]).abs().max().item(), 1e-6
|
||||
)
|
||||
# Updating statistics is not eval identity. Random gravity probes include
|
||||
# out-of-distribution values in nearly constant source channels.
|
||||
updated_source = copy.deepcopy(source_actor).train()
|
||||
updated_source.update_normalization(
|
||||
TensorDict({"actor": base}, batch_size=[len(base)])
|
||||
)
|
||||
updated_expected = updated_source.mlp(updated_source.obs_normalizer(base)).detach()
|
||||
torch.testing.assert_close(actor(obs), updated_expected, atol=3e-6, rtol=3e-6)
|
||||
self.assertLess((actor(obs) - expected).abs().max().item(), 2e-4)
|
||||
if hasattr(self, "evidence"):
|
||||
self.evidence.setdefault("cpu_transfer", {})[str(dim)] = {
|
||||
"initial_max_abs_error": initial_error,
|
||||
"updated_source_max_abs_error": (actor(obs) - updated_expected)
|
||||
.abs()
|
||||
.max()
|
||||
.item(),
|
||||
"random_batch_action_max_abs_change": (actor(obs) - expected)
|
||||
.abs()
|
||||
.max()
|
||||
.item(),
|
||||
}
|
||||
loss = actor(obs).square().mean()
|
||||
loss.backward()
|
||||
grad = actor.mlp[0].weight.grad
|
||||
self.assertTrue(torch.isfinite(grad).all())
|
||||
if dim > 47:
|
||||
self.assertGreater(grad[:, 47:].abs().max().item(), 0)
|
||||
torch.optim.Adam(actor.parameters(), lr=1e-4).step()
|
||||
if dim > 47:
|
||||
self.assertGreater(actor.mlp[0].weight[:, 47:].abs().max().item(), 0)
|
||||
stream = io.BytesIO()
|
||||
torch.save(actor.state_dict(), stream)
|
||||
stream.seek(0)
|
||||
restored = make_reference_actor(dim)
|
||||
restored.load_state_dict(torch.load(stream, weights_only=True), strict=True)
|
||||
for k, v in actor.state_dict().items():
|
||||
self.assertTrue(torch.equal(v, restored.state_dict()[k]), k)
|
||||
|
||||
def test_rejects_shapes_nonfinite_and_runtime_activation(self):
|
||||
import torch
|
||||
from pretrained import (
|
||||
PretrainedError,
|
||||
ValidatedSource,
|
||||
make_reference_actor,
|
||||
warm_start_actor,
|
||||
)
|
||||
|
||||
for key, value in (
|
||||
("mlp.0.weight", torch.zeros(512, 46)),
|
||||
("obs_normalizer._std", torch.full((1, 47), float("nan"))),
|
||||
("distribution.std_param", torch.zeros(12)),
|
||||
):
|
||||
state = dict(self.source.actor_state)
|
||||
state[key] = value
|
||||
with self.assertRaises(PretrainedError):
|
||||
warm_start_actor(make_reference_actor(81), ValidatedSource(state, {}))
|
||||
actor = make_reference_actor(81)
|
||||
actor.mlp[1] = torch.nn.ReLU()
|
||||
with self.assertRaisesRegex(PretrainedError, "architecture"):
|
||||
warm_start_actor(actor, self.source)
|
||||
with self.assertRaises(PretrainedError):
|
||||
warm_start_actor(make_reference_actor(82), self.source)
|
||||
|
||||
def test_semantics_fail_closed(self):
|
||||
from pretrained import PretrainedError, _plain, validate_semantics
|
||||
|
||||
env, agent = _plain(self.env_cfg), _plain(self.agent_cfg)
|
||||
validate_semantics(env, agent, self.env_cfg, self.agent_cfg)
|
||||
for edit in (
|
||||
lambda e: e["observations"]["actor"]["terms"]["phase"]["params"].update(period=0.7),
|
||||
lambda e: e["observations"]["actor"]["terms"]["joint_vel"].update(scale=0.1),
|
||||
lambda e: e["actions"]["joint_pos"].update(scale=0.5),
|
||||
lambda e: e["scene"]["entities"]["robot"]["articulation"]["actuators"][0].update(
|
||||
armature=0.1
|
||||
),
|
||||
):
|
||||
bad = copy.deepcopy(env)
|
||||
edit(bad)
|
||||
with self.assertRaises(PretrainedError):
|
||||
validate_semantics(bad, agent, self.env_cfg, self.agent_cfg)
|
||||
with self.assertRaises(PretrainedError):
|
||||
validate_semantics(env, agent, bad, self.agent_cfg)
|
||||
bad_agent = copy.deepcopy(agent)
|
||||
bad_agent["actor"]["activation"] = "relu"
|
||||
with self.assertRaises(PretrainedError):
|
||||
validate_semantics(env, bad_agent, self.env_cfg, self.agent_cfg)
|
||||
|
||||
def test_compiled_runtime_contract(self):
|
||||
from unittest.mock import patch
|
||||
|
||||
from pretrained import BASE_TERMS, JOINTS, PretrainedError, validate_runtime_contract
|
||||
|
||||
metadata = {
|
||||
"joint_names": JOINTS,
|
||||
"observation_names": BASE_TERMS + ["forward_depth", "target_error"],
|
||||
"command_names": ["twist"],
|
||||
"action_scale": 0.25,
|
||||
"joint_stiffness": [20, 20, 40] * 4,
|
||||
"joint_damping": [1, 1, 2] * 4,
|
||||
"default_joint_pos": [-0.1, 0.9, -1.8, 0.1, 0.9, -1.8] * 2,
|
||||
}
|
||||
with patch("mjlab.rl.exporter_utils.get_base_metadata", return_value=metadata):
|
||||
validate_runtime_contract(None)
|
||||
metadata["joint_names"] = list(reversed(JOINTS))
|
||||
with self.assertRaisesRegex(PretrainedError, "joint order"):
|
||||
validate_runtime_contract(None)
|
||||
|
||||
def test_safe_yaml_and_allowed_roots(self):
|
||||
from pretrained import PretrainedError, _read_allowed, _yaml_data
|
||||
|
||||
with self.assertRaises(PretrainedError):
|
||||
_yaml_data(b"x: !!python/object/apply:os.system ['touch /tmp/never-pretrained']")
|
||||
with self.assertRaises(PretrainedError):
|
||||
_yaml_data(b"x: 1\nx: 2")
|
||||
self.assertEqual(_yaml_data(b"axis: {0: 1, 1: 2}"), {"axis": {0: 1, 1: 2}})
|
||||
self.assertEqual(
|
||||
_yaml_data(b"x: !!python/name:os.system ''"), {"x": {"symbol": "os.system"}}
|
||||
)
|
||||
with tempfile.TemporaryDirectory() as directory, tempfile.TemporaryDirectory() as outside:
|
||||
root = Path(directory)
|
||||
secret = Path(outside) / "secret.pt"
|
||||
secret.write_bytes(b"x")
|
||||
(root / "link.pt").symlink_to(secret)
|
||||
for path in (secret, root / "link.pt", root / ".." / Path(outside).name / "secret.pt"):
|
||||
with self.assertRaises(PretrainedError):
|
||||
_read_allowed(path, [root], 128)
|
||||
with self.assertRaises(PretrainedError):
|
||||
_read_allowed(secret, [Path(outside)], 0)
|
||||
|
||||
def test_fresh_runner_only_and_trial_consistency(self):
|
||||
import torch
|
||||
from pretrained import PretrainedError, initialize_runner, make_reference_actor
|
||||
|
||||
actors = []
|
||||
for _ in range(2):
|
||||
actor = make_reference_actor(81)
|
||||
critic = torch.nn.Linear(108, 1)
|
||||
original = copy.deepcopy(critic.state_dict())
|
||||
optimizer = torch.optim.Adam(list(actor.parameters()) + list(critic.parameters()))
|
||||
runner = SimpleNamespace(
|
||||
current_learning_iteration=0,
|
||||
alg=SimpleNamespace(actor=actor, critic=critic, optimizer=optimizer),
|
||||
)
|
||||
initialize_runner(runner, self.source)
|
||||
self.assertFalse(optimizer.state)
|
||||
self.assertEqual(runner.current_learning_iteration, 0)
|
||||
for k in original:
|
||||
self.assertTrue(torch.equal(original[k], critic.state_dict()[k]))
|
||||
actors.append(actor.state_dict())
|
||||
runner.current_learning_iteration = 1
|
||||
with self.assertRaises(PretrainedError):
|
||||
initialize_runner(runner, self.source)
|
||||
runner.current_learning_iteration = 0
|
||||
optimizer.state[actor.mlp[0].weight] = {"step": torch.tensor(1)}
|
||||
with self.assertRaises(PretrainedError):
|
||||
initialize_runner(runner, self.source)
|
||||
for k in actors[0]:
|
||||
self.assertTrue(torch.equal(actors[0][k], actors[1][k]))
|
||||
|
||||
def test_cli_modes_do_not_reinitialize_resume(self):
|
||||
from scripts.train import TrainConfig, _load_pretrained
|
||||
from src.tasks.velocity.config.go2.env_cfgs import unitree_go2_flat_env_cfg
|
||||
from src.tasks.velocity.config.go2.rl_cfg import unitree_go2_ppo_runner_cfg
|
||||
|
||||
cfg = TrainConfig(env=unitree_go2_flat_env_cfg(), agent=unitree_go2_ppo_runner_cfg())
|
||||
self.assertIsNone(_load_pretrained(cfg))
|
||||
self.assertIsNone(_load_pretrained(replace(cfg, resume_checkpoint="trial.pt")))
|
||||
with self.assertRaisesRegex(ValueError, "mutually exclusive"):
|
||||
_load_pretrained(
|
||||
replace(cfg, resume_checkpoint="trial.pt", pretrained_checkpoint="source.pt")
|
||||
)
|
||||
cfg.agent.resume = True
|
||||
with self.assertRaisesRegex(ValueError, "mutually exclusive"):
|
||||
_load_pretrained(replace(cfg, pretrained_checkpoint="source.pt"))
|
||||
with self.assertRaisesRegex(ValueError, "ONNX alone"):
|
||||
_load_pretrained(replace(cfg, pretrained_onnx="policy.onnx"))
|
||||
|
||||
|
||||
@unittest.skipUnless(
|
||||
HAS_STACK and os.environ.get("GO2_PRETRAINED_SOURCE"),
|
||||
"set GO2_PRETRAINED_SOURCE to explicit allowed source directory",
|
||||
)
|
||||
class RealPretrainedTest(PretrainedTest):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
super().setUpClass()
|
||||
from pretrained import read_pretrained_source
|
||||
|
||||
directory = Path(os.environ["GO2_PRETRAINED_SOURCE"])
|
||||
cls.source = read_pretrained_source(
|
||||
directory / "model_10000.pt",
|
||||
allowed_roots=[directory],
|
||||
target_env=cls.env_cfg,
|
||||
target_agent=cls.agent_cfg,
|
||||
)
|
||||
cls.evidence = {"source": cls.source.manifest}
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
if os.environ.get("GO2_PRETRAINED_EVIDENCE_DIR"):
|
||||
directory = Path(os.environ["GO2_PRETRAINED_EVIDENCE_DIR"])
|
||||
directory.mkdir(parents=True, exist_ok=True)
|
||||
(directory / "real-source-evidence.json").write_text(
|
||||
json.dumps(cls.evidence, indent=2) + "\n"
|
||||
)
|
||||
|
||||
def test_wrong_checkpoint_actor_rejected_by_onnx(self):
|
||||
from pretrained import PretrainedError, make_reference_actor, verify_onnx
|
||||
|
||||
# No bulk loading: a fresh random actor is sufficient to prove fail-closed identity.
|
||||
onnx = (Path(os.environ["GO2_PRETRAINED_SOURCE"]) / "policy.onnx").read_bytes()
|
||||
with self.assertRaisesRegex(PretrainedError, "does not match"):
|
||||
verify_onnx(onnx, make_reference_actor())
|
||||
|
||||
def test_extended_export_matches_cpu_actor(self):
|
||||
import numpy as np
|
||||
import onnxruntime as ort
|
||||
import torch
|
||||
from pretrained import comparison_observations, make_reference_actor, warm_start_actor
|
||||
|
||||
errors = {}
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
for dim in (81, 97):
|
||||
actor = make_reference_actor(dim).eval()
|
||||
warm_start_actor(actor, self.source)
|
||||
export = actor.as_onnx(verbose=False)
|
||||
path = str(Path(directory) / f"policy-{dim}.onnx")
|
||||
torch.onnx.export(
|
||||
export,
|
||||
export.get_dummy_inputs(),
|
||||
path,
|
||||
input_names=export.input_names,
|
||||
output_names=export.output_names,
|
||||
opset_version=18,
|
||||
dynamo=False,
|
||||
)
|
||||
session = ort.InferenceSession(path, providers=["CPUExecutionProvider"])
|
||||
x = torch.cat((comparison_observations(), torch.rand(48, dim - 47) * 2 - 1), dim=1)
|
||||
with torch.no_grad():
|
||||
expected = export(x).numpy()
|
||||
actual = np.concatenate(
|
||||
[session.run(None, {"obs": row[None].numpy()})[0] for row in x]
|
||||
)
|
||||
np.testing.assert_allclose(actual, expected, atol=2e-5, rtol=2e-5)
|
||||
errors[str(dim)] = float(np.abs(actual - expected).max())
|
||||
self.evidence["extended_onnx_max_abs_error"] = errors
|
||||
|
||||
@unittest.skipUnless(
|
||||
os.environ.get("GO2_PRETRAINED_GPU_SMOKE") == "1", "opt-in <=4env x 1iteration GPU smoke"
|
||||
)
|
||||
def test_real_observations_and_one_ppo_iteration(self):
|
||||
import torch
|
||||
import warp as wp
|
||||
|
||||
if not hasattr(wp, "context"):
|
||||
from warp._src import context
|
||||
|
||||
wp.context = context
|
||||
from mjlab.envs import ManagerBasedRlEnv
|
||||
from mjlab.rl import RslRlVecEnvWrapper
|
||||
from pretrained import initialize_runner, make_reference_actor, validate_runtime_contract
|
||||
from src.tasks.obstacle_avoidance.env_cfg import unitree_go2_obstacle_env_cfg
|
||||
from src.tasks.velocity.config.go2.rl_cfg import unitree_go2_ppo_runner_cfg
|
||||
from src.tasks.velocity.rl.runner import VelocityOnPolicyRunner
|
||||
from task_config import OBSTACLE_TASK, deployment_metadata, validate_task_config
|
||||
from tensordict import TensorDict
|
||||
|
||||
env_cfg = unitree_go2_obstacle_env_cfg()
|
||||
env_cfg.scene.num_envs = 4
|
||||
env_cfg.seed = 42
|
||||
agent_cfg = unitree_go2_ppo_runner_cfg()
|
||||
agent_cfg.logger = "tensorboard"
|
||||
agent_cfg.max_iterations = 1
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
raw = ManagerBasedRlEnv(env_cfg, device="cuda:0")
|
||||
self.evidence["compiled_runtime"] = validate_runtime_contract(raw)
|
||||
raw.platform_deployment = deployment_metadata(
|
||||
OBSTACLE_TASK, validate_task_config(OBSTACLE_TASK, {}, 42), 42
|
||||
)
|
||||
env = RslRlVecEnvWrapper(raw)
|
||||
try:
|
||||
runner = VelocityOnPolicyRunner(env, asdict(agent_cfg), directory, "cuda:0")
|
||||
manifest = initialize_runner(runner, self.source)
|
||||
actor = runner.alg.actor
|
||||
cpu_actor = make_reference_actor().eval()
|
||||
cpu_actor.load_state_dict(self.source.actor_state)
|
||||
obs = env.get_observations()
|
||||
base = obs["actor"][:, :47].cpu()
|
||||
with torch.no_grad():
|
||||
expected = cpu_actor(TensorDict({"actor": base}, batch_size=[4]))
|
||||
before = actor(obs).cpu()
|
||||
torch.testing.assert_close(before, expected, atol=2e-5, rtol=2e-5)
|
||||
old_mean = actor.obs_normalizer._mean.clone()
|
||||
actor.update_normalization(obs)
|
||||
mean_change = (
|
||||
(actor.obs_normalizer._mean[:, :47] - old_mean[:, :47]).abs().max().item()
|
||||
)
|
||||
after = actor(obs).detach().cpu()
|
||||
self.assertTrue(torch.isfinite(after).all())
|
||||
torch.testing.assert_close(before, after, atol=2e-5, rtol=2e-5)
|
||||
loss = actor(obs).square().mean()
|
||||
loss.backward()
|
||||
grad = actor.mlp[0].weight.grad[:, 47:]
|
||||
self.assertTrue(torch.isfinite(grad).all())
|
||||
self.assertGreater(grad.abs().max().item(), 0)
|
||||
grad_max = grad.abs().max().item()
|
||||
runner.alg.optimizer.zero_grad()
|
||||
critic_before = copy.deepcopy(runner.alg.critic.state_dict())
|
||||
runner.learn(num_learning_iterations=1, init_at_random_ep_len=True)
|
||||
self.assertGreater(actor.mlp[0].weight[:, 47:].abs().max().item(), 0)
|
||||
self.assertTrue(runner.alg.optimizer.state)
|
||||
self.assertTrue(
|
||||
any(
|
||||
not torch.equal(v, runner.alg.critic.state_dict()[k])
|
||||
for k, v in critic_before.items()
|
||||
)
|
||||
)
|
||||
saved = torch.load(
|
||||
Path(directory) / "model_0.pt", map_location="cpu", weights_only=True
|
||||
)
|
||||
# Same-trial resume uses the installed PPO loader: restore all states.
|
||||
actor_state = copy.deepcopy(actor.state_dict())
|
||||
# Production promotion starts a fresh process/runner; do not reload
|
||||
# into tensors created by the previous rollout's inference_mode.
|
||||
resumed = VelocityOnPolicyRunner(env, asdict(agent_cfg), directory, "cuda:0")
|
||||
self.assertTrue(resumed.alg.load(saved, None, strict=True))
|
||||
resumed.current_learning_iteration = saved["iter"] + 1
|
||||
self.assertEqual(resumed.current_learning_iteration, 1)
|
||||
self.assertTrue(resumed.alg.optimizer.state)
|
||||
actor = resumed.alg.actor
|
||||
for k, v in actor_state.items():
|
||||
self.assertTrue(torch.equal(v.cpu(), actor.state_dict()[k].cpu()), k)
|
||||
import numpy as np
|
||||
import onnxruntime as ort
|
||||
|
||||
session = ort.InferenceSession(
|
||||
str(Path(directory) / "policy.onnx"), providers=["CPUExecutionProvider"]
|
||||
)
|
||||
final_obs = env.get_observations()
|
||||
with torch.no_grad():
|
||||
expected_final = actor(final_obs).cpu().numpy()
|
||||
actual = np.concatenate(
|
||||
[
|
||||
session.run(None, {"obs": row[None].cpu().numpy()})[0]
|
||||
for row in final_obs["actor"]
|
||||
]
|
||||
)
|
||||
np.testing.assert_allclose(actual, expected_final, atol=2e-5, rtol=2e-5)
|
||||
# Independently run source ORT on observations actually sampled from MuJoCo.
|
||||
source_session = ort.InferenceSession(
|
||||
str(Path(os.environ["GO2_PRETRAINED_SOURCE"]) / "policy.onnx"),
|
||||
providers=["CPUExecutionProvider"],
|
||||
)
|
||||
source_actions = np.concatenate(
|
||||
[source_session.run(None, {"obs": row[None].numpy()})[0] for row in base]
|
||||
)
|
||||
np.testing.assert_allclose(source_actions, before.numpy(), atol=2e-5, rtol=2e-5)
|
||||
self.evidence["gpu_smoke"] = {
|
||||
"num_envs": 4,
|
||||
"iterations": 1,
|
||||
"initialization": manifest,
|
||||
"real_observation_source_onnx_max_abs_error": float(
|
||||
np.abs(source_actions - before.numpy()).max()
|
||||
),
|
||||
"normalization_batch_action_max_abs_change": (after - before)
|
||||
.abs()
|
||||
.max()
|
||||
.item(),
|
||||
"normalization_batch_base_mean_max_abs_change": mean_change,
|
||||
"new_columns_gradient_max": grad_max,
|
||||
"new_columns_weight_max_after_ppo": actor.mlp[0]
|
||||
.weight[:, 47:]
|
||||
.abs()
|
||||
.max()
|
||||
.item(),
|
||||
"export_max_abs_error": float(np.abs(actual - expected_final).max()),
|
||||
"count_restored": actor.obs_normalizer.count.item(),
|
||||
"all_actor_tensors_restored_exactly": True,
|
||||
}
|
||||
finally:
|
||||
env.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,307 @@
|
||||
"""Registered source authority, immutable snapshots, and same-trial continuation."""
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
import threading
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
|
||||
from pretrained_sources import PretrainedSources, SourceError, regular_bytes
|
||||
from server import ApiError, TrainingManager
|
||||
from test_tuning_manager import BASE_METRICS, FakeAdvisor
|
||||
from tuning.manager import TuningError, TuningManager
|
||||
from tuning.process import GpuLease
|
||||
from tuning.schema import BASE_REWARD_CONFIGURATION, RewardConfigError, validate_proposal
|
||||
|
||||
|
||||
def fake_validate(_self, directory, checkpoint, *_args):
|
||||
artifacts = {}
|
||||
for key, relative in {
|
||||
"checkpoint": checkpoint,
|
||||
"onnx": "policy.onnx",
|
||||
"env": "params/env.yaml",
|
||||
"agent": "params/agent.yaml",
|
||||
}.items():
|
||||
data = (directory / relative).read_bytes()
|
||||
artifacts[key] = {
|
||||
"name": Path(relative).name,
|
||||
"sha256": hashlib.sha256(data).hexdigest(),
|
||||
"bytes": len(data),
|
||||
}
|
||||
return {
|
||||
"source_id": hashlib.sha256(json.dumps(artifacts, sort_keys=True).encode()).hexdigest(),
|
||||
"artifacts": artifacts,
|
||||
"source_iteration": 10000,
|
||||
}
|
||||
|
||||
|
||||
class RegisteredSourcesTest(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.temp = tempfile.TemporaryDirectory()
|
||||
self.root = Path(self.temp.name)
|
||||
self.original = self.root / "original"
|
||||
(self.original / "params").mkdir(parents=True)
|
||||
for relative in ("model_10000.pt", "policy.onnx", "params/env.yaml", "params/agent.yaml"):
|
||||
(self.original / relative).write_bytes(relative.encode())
|
||||
self.config = self.root / "sources.json"
|
||||
self.config.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"allowedRoots": [str(self.original)],
|
||||
"sources": [
|
||||
{
|
||||
"id": "base",
|
||||
"label": "基础行走",
|
||||
"checkpoint": str(self.original / "model_10000.pt"),
|
||||
"onnx": str(self.original / "policy.onnx"),
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
)
|
||||
self.validator = patch.object(PretrainedSources, "_validate", fake_validate)
|
||||
self.validator.start()
|
||||
self.registry = PretrainedSources(
|
||||
self.config,
|
||||
self.root / "snapshots",
|
||||
sys.executable,
|
||||
Path(__file__).resolve().parents[1] / "rl",
|
||||
)
|
||||
|
||||
def tearDown(self):
|
||||
self.validator.stop()
|
||||
self.temp.cleanup()
|
||||
|
||||
def test_registered_snapshot_never_rereads_mutated_original_and_checks_sha(self):
|
||||
bound = self.registry.bind(self.registry.catalog()[0]["id"], "Unitree-Go2-Flat")
|
||||
original_sha = bound["manifest"]["artifacts"]["checkpoint"]["sha256"]
|
||||
(self.original / "model_10000.pt").write_bytes(b"changed after registration")
|
||||
self.assertEqual(
|
||||
self.registry.bind(self.registry.catalog()[0]["id"], "Unitree-Go2-Flat"), bound
|
||||
)
|
||||
args = self.registry.arguments(bound)
|
||||
checkpoint = Path(args[args.index("--pretrained-checkpoint") + 1])
|
||||
self.assertEqual(hashlib.sha256(checkpoint.read_bytes()).hexdigest(), original_sha)
|
||||
restored = PretrainedSources(None, self.root / "snapshots", sys.executable, self.root)
|
||||
self.assertEqual(restored.arguments(bound), args)
|
||||
re_registered = PretrainedSources(
|
||||
self.config, self.root / "snapshots", sys.executable, self.root
|
||||
)
|
||||
with self.assertRaisesRegex(SourceError, "已变化"):
|
||||
re_registered.bind(bound["sourceId"], "Unitree-Go2-Flat")
|
||||
self.assertEqual(restored.arguments(bound), args)
|
||||
checkpoint.chmod(0o644)
|
||||
checkpoint.write_bytes(b"tamper")
|
||||
with self.assertRaisesRegex(SourceError, "SHA"):
|
||||
restored.arguments(bound)
|
||||
checkpoint.parent.chmod(0o755)
|
||||
checkpoint.unlink()
|
||||
with self.assertRaises(SourceError):
|
||||
restored.verify(bound)
|
||||
|
||||
def test_rejects_path_symlink_fifo_size_and_invalid_registry(self):
|
||||
with self.assertRaises(SourceError):
|
||||
self.registry.bind(str(self.original / "model_10000.pt"), "Unitree-Go2-Flat")
|
||||
with self.assertRaisesRegex(SourceError, "Rough"):
|
||||
self.registry.bind(self.registry.catalog()[0]["id"], "Unitree-Go2-Rough")
|
||||
(self.original / "link.pt").symlink_to(self.original / "model_10000.pt")
|
||||
with self.assertRaises(SourceError):
|
||||
regular_bytes(self.original / "link.pt", self.original, 256)
|
||||
with self.assertRaises(SourceError):
|
||||
regular_bytes(self.original / ".." / "sources.json", self.original, 256)
|
||||
with self.assertRaises(SourceError):
|
||||
regular_bytes(self.original / "policy.onnx", self.original, 1)
|
||||
os.mkfifo(self.original / "pipe")
|
||||
with self.assertRaises(SourceError):
|
||||
regular_bytes(self.original / "pipe", self.original, 256)
|
||||
self.config.write_text('{"allowedRoots": [], "sources": [], "python": "evil"}')
|
||||
with self.assertRaises(SourceError):
|
||||
PretrainedSources(self.config, self.root / "other", sys.executable, self.root)
|
||||
|
||||
def test_missing_checkpoint_is_visible_failure_not_random_fallback(self):
|
||||
(self.original / "model_10000.pt").unlink()
|
||||
registry = PretrainedSources(self.config, self.root / "other", sys.executable, self.root)
|
||||
self.assertFalse(registry.catalog()[0]["ready"])
|
||||
with self.assertRaisesRegex(SourceError, "pt"):
|
||||
registry.bind("base", "Unitree-Go2-Flat")
|
||||
|
||||
def test_training_request_only_id_and_public_identity(self):
|
||||
manager = TrainingManager(
|
||||
Path(__file__).resolve().parents[1] / "rl",
|
||||
sys.executable,
|
||||
("Unitree-Go2-Flat", "Unitree-Go2-Rough"),
|
||||
check_environment=False,
|
||||
sources=self.registry,
|
||||
)
|
||||
payload = {
|
||||
"taskId": "Unitree-Go2-Flat",
|
||||
"numEnvs": 4,
|
||||
"maxIterations": 1,
|
||||
"seed": 42,
|
||||
"device": "cpu",
|
||||
"pretrainedSourceId": self.registry.catalog()[0]["id"],
|
||||
}
|
||||
config = manager.parse_config(payload)
|
||||
self.assertEqual(
|
||||
config.pretrained,
|
||||
self.registry.bind(self.registry.catalog()[0]["id"], "Unitree-Go2-Flat"),
|
||||
)
|
||||
command = manager.command_for(config)
|
||||
self.assertIn("--pretrained-source-id", command)
|
||||
self.assertNotIn("--resume-checkpoint", command)
|
||||
self.assertNotIn(str(self.original), json.dumps(manager.health()["pretrainedSources"]))
|
||||
for key in ("pretrainedCheckpoint", "allowedRoots", "pretrained", "resumeCheckpoint"):
|
||||
with self.assertRaises(ApiError):
|
||||
manager.parse_config({**payload, key: "/etc/passwd"})
|
||||
with self.assertRaises(ApiError):
|
||||
manager.parse_config({**payload, "pretrainedSourceId": None})
|
||||
with self.assertRaises(ApiError):
|
||||
manager.parse_config({**payload, "taskId": "Unitree-Go2-Rough"})
|
||||
old = manager.parse_config({k: v for k, v in payload.items() if k != "pretrainedSourceId"})
|
||||
self.assertIsNone(old.pretrained)
|
||||
self.assertNotIn("--pretrained-checkpoint", manager.command_for(old))
|
||||
|
||||
def tuning(self):
|
||||
return TuningManager(
|
||||
Path(__file__).resolve().parents[1] / "rl",
|
||||
sys.executable,
|
||||
self.root / "tuning",
|
||||
GpuLease(),
|
||||
advisor=FakeAdvisor(),
|
||||
sources=self.registry,
|
||||
)
|
||||
|
||||
def session(self, manager, mode="approval"):
|
||||
_, config, objective, fallback = manager.parse_create(
|
||||
{
|
||||
"mode": mode,
|
||||
"pretrainedSourceId": self.registry.catalog()[0]["id"],
|
||||
"trialCount": 2,
|
||||
"numEnvs": 4,
|
||||
"initialIterations": 1,
|
||||
"middleIterations": 2,
|
||||
"finalIterations": 3,
|
||||
}
|
||||
)
|
||||
return manager.storage.create_session(mode, config, objective, fallback)
|
||||
|
||||
def test_baseline_new_trial_warmstart_and_rung_resumes_only_own_checkpoint(self):
|
||||
manager = self.tuning()
|
||||
session = self.session(manager)
|
||||
manager.cancel_events[session["id"]] = threading.Event()
|
||||
commands = []
|
||||
|
||||
def run(_session, command, _cwd, _env, _log):
|
||||
commands.append(command)
|
||||
if "--output-dir" in command:
|
||||
directory = Path(command[command.index("--output-dir") + 1])
|
||||
(directory / "model_0.pt").write_bytes(b"trial checkpoint")
|
||||
(directory / "policy.onnx").write_bytes(b"trial policy")
|
||||
(directory / "initialization.json").write_text(
|
||||
json.dumps(session["config"]["pretrained"]["manifest"])
|
||||
)
|
||||
else:
|
||||
Path(command[command.index("--output") + 1]).write_text(
|
||||
json.dumps({"metrics": BASE_METRICS})
|
||||
)
|
||||
return 0
|
||||
|
||||
with patch.object(manager, "_run_command", run):
|
||||
parent = None
|
||||
for number, rung in ((0, 0), (1, 0), (1, 1)):
|
||||
trial = manager.storage.create_trial(
|
||||
session["id"],
|
||||
number,
|
||||
rung,
|
||||
rung + 1,
|
||||
BASE_REWARD_CONFIGURATION,
|
||||
None,
|
||||
f"trial-{number}-{rung}",
|
||||
)
|
||||
result = manager._execute_trial(session, trial, parent if rung else None)
|
||||
parent = manager._session_root(session["id"]) / result["checkpointPath"]
|
||||
train = [c for c in commands if "--output-dir" in c]
|
||||
self.assertEqual(
|
||||
train[0][train[0].index("--pretrained-source-id") + 1],
|
||||
train[1][train[1].index("--pretrained-source-id") + 1],
|
||||
)
|
||||
self.assertIn("--resume-checkpoint", train[2])
|
||||
self.assertNotIn("--pretrained-checkpoint", train[2])
|
||||
self.assertIn("trial-1-0/model_0.pt", train[2][-1])
|
||||
self.assertEqual(len([c for c in commands if "--steps-per-seed=1000" in c]), 3)
|
||||
with self.assertRaises(RewardConfigError):
|
||||
validate_proposal({"pretrainedSourceId": "other"}, BASE_REWARD_CONFIGURATION)
|
||||
|
||||
def test_restart_states_source_constraints_and_explicit_recovery(self):
|
||||
manager = self.tuning()
|
||||
sessions = []
|
||||
for state in ("queued", "paused", "awaiting_approval", "succeeded"):
|
||||
session = self.session(manager)
|
||||
trial = manager.storage.create_trial(
|
||||
session["id"], 0, 0, 1, BASE_REWARD_CONFIGURATION, None, "baseline"
|
||||
)
|
||||
manager.storage.update_trial(
|
||||
trial["id"],
|
||||
state="completed",
|
||||
evaluation={"metrics": BASE_METRICS},
|
||||
eligible=True,
|
||||
score=0,
|
||||
)
|
||||
proposal = manager.storage.create_proposal(
|
||||
session["id"], trial["id"], {"weights": {"pose": 1.1}}, "proposal", {}, 0.8
|
||||
)
|
||||
manager.storage.replace_constraints(
|
||||
session["id"], 0, {"weights.pose": {"kind": "range", "min": 0, "max": 2}}
|
||||
)
|
||||
manager.storage.update_session(session["id"], state=state)
|
||||
sessions.append((session, proposal, state))
|
||||
restored = self.tuning() # Actual SQLite reopening, no worker/API calls on construction.
|
||||
self.assertFalse(restored.workers)
|
||||
for session, _proposal, previous in sessions:
|
||||
value = restored.detail(session["id"])
|
||||
self.assertEqual(
|
||||
value["state"], "succeeded" if previous == "succeeded" else "interrupted"
|
||||
)
|
||||
self.assertEqual(value["config"]["pretrained"], session["config"]["pretrained"])
|
||||
self.assertEqual(value["control"]["constraintsRevision"], 1)
|
||||
self.assertEqual(value["mode"], "approval")
|
||||
self.assertEqual(value["proposals"][0]["state"], "pending")
|
||||
session, proposal, _ = sessions[0]
|
||||
# Stop before dispatch: recovery must invalidate pending first, never execute it.
|
||||
with (
|
||||
patch.object(restored, "_wait_for_dispatch", side_effect=RuntimeError("test boundary")),
|
||||
patch.object(restored.advisor, "propose") as advisor,
|
||||
):
|
||||
restored._run_session(session["id"], True, threading.Event())
|
||||
advisor.assert_not_called()
|
||||
self.assertEqual(restored.storage.get_proposal(proposal["id"])["state"], "rejected")
|
||||
self.assertIn(
|
||||
"recovery_invalidated", restored.storage.get_proposal(proposal["id"])["feedback"]
|
||||
)
|
||||
self.assertEqual(len(restored.storage.list_trials(session["id"])), 1)
|
||||
session = sessions[1][0]
|
||||
with patch.object(restored, "_start_worker") as start:
|
||||
restored.resume(session["id"])
|
||||
with self.assertRaises(TuningError):
|
||||
restored.resume(session["id"])
|
||||
start.assert_called_once()
|
||||
session = sessions[2][0]
|
||||
snapshot = self.registry.verify(session["config"]["pretrained"])
|
||||
file = snapshot / "policy.onnx"
|
||||
file.chmod(0o644)
|
||||
file.write_bytes(b"corrupt")
|
||||
with patch.object(restored, "_start_worker") as start:
|
||||
with self.assertRaises(SourceError):
|
||||
restored.resume(session["id"])
|
||||
start.assert_not_called()
|
||||
self.assertEqual(restored.storage.get_session(session["id"])["state"], "interrupted")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,735 @@
|
||||
"""Single-file upload security, strict graph reconstruction, and shared warm-start seam."""
|
||||
|
||||
import copy
|
||||
import errno
|
||||
import hashlib
|
||||
import http.client
|
||||
import importlib.util
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
import socket
|
||||
import sys
|
||||
import tempfile
|
||||
import threading
|
||||
import unittest
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import asdict
|
||||
from http.server import ThreadingHTTPServer
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[1]
|
||||
for root in (ROOT, ROOT / "rl"):
|
||||
sys.path.insert(0, str(root))
|
||||
|
||||
from pretrained_sources import PretrainedSources, SourceError # noqa: E402
|
||||
from server import TrainingManager, TrainingRequestHandler # noqa: E402
|
||||
|
||||
|
||||
def fixture_onnx(state):
|
||||
import onnx
|
||||
from onnx import helper, numpy_helper
|
||||
from pretrained_upload import SEMANTICS
|
||||
|
||||
tensors = [
|
||||
numpy_helper.from_array(value.numpy(), key)
|
||||
for key, value in state.items()
|
||||
if key.startswith("mlp.") or key == "obs_normalizer._mean"
|
||||
]
|
||||
tensors.append(
|
||||
numpy_helper.from_array((state["obs_normalizer._std"] + 0.01).numpy(), "onnx::Div_24")
|
||||
)
|
||||
nodes = [
|
||||
helper.make_node("Sub", ["obs", "obs_normalizer._mean"], ["centered"]),
|
||||
helper.make_node("Div", ["centered", "onnx::Div_24"], ["normalized"]),
|
||||
]
|
||||
previous = "normalized"
|
||||
for i in (0, 2, 4, 6):
|
||||
output = "actions" if i == 6 else f"linear{i}"
|
||||
nodes.append(
|
||||
helper.make_node(
|
||||
"Gemm", [previous, f"mlp.{i}.weight", f"mlp.{i}.bias"], [output], transB=1
|
||||
)
|
||||
)
|
||||
previous = output
|
||||
if i < 6:
|
||||
previous = f"elu{i}"
|
||||
nodes.append(helper.make_node("Elu", [output], [previous], alpha=1.0))
|
||||
model = helper.make_model(
|
||||
helper.make_graph(
|
||||
nodes,
|
||||
"actor",
|
||||
[helper.make_tensor_value_info("obs", onnx.TensorProto.FLOAT, [1, 47])],
|
||||
[helper.make_tensor_value_info("actions", onnx.TensorProto.FLOAT, [1, 12])],
|
||||
tensors,
|
||||
),
|
||||
opset_imports=[helper.make_opsetid("", 18)],
|
||||
ir_version=8,
|
||||
)
|
||||
helper.set_model_props(
|
||||
model, {key: ",".join(map(str, value)) for key, value in SEMANTICS.items()}
|
||||
)
|
||||
return model
|
||||
|
||||
|
||||
@unittest.skipUnless(
|
||||
all(importlib.util.find_spec(m) is not None for m in ("torch", "mjlab", "onnx", "onnxruntime")),
|
||||
"installed CPU RL/ONNX stack required",
|
||||
)
|
||||
class UploadTest(unittest.TestCase):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
import torch
|
||||
from pretrained import make_reference_actor
|
||||
|
||||
torch.set_num_threads(1)
|
||||
torch.manual_seed(123)
|
||||
actor = make_reference_actor()
|
||||
actor.obs_normalizer.count.fill_(983138304)
|
||||
cls.state = actor.state_dict()
|
||||
stream = io.BytesIO()
|
||||
torch.save({"actor_state_dict": cls.state, "iter": 10000}, stream)
|
||||
cls.pt = stream.getvalue()
|
||||
cls.onnx = fixture_onnx(cls.state).SerializeToString()
|
||||
|
||||
def setUp(self):
|
||||
self.temp = tempfile.TemporaryDirectory()
|
||||
self.root = Path(self.temp.name)
|
||||
self.registry = PretrainedSources(None, self.root / "store", sys.executable, ROOT / "rl")
|
||||
|
||||
def tearDown(self):
|
||||
self.temp.cleanup()
|
||||
|
||||
def upload(self, data=None, fmt="pt", name="policy.pt"):
|
||||
data = self.pt if data is None else data
|
||||
return self.registry.receive_upload(
|
||||
io.BytesIO(data), len(data), fmt, "go2-legacy47-v1", name
|
||||
)
|
||||
|
||||
def test_single_file_persistence_dedup_and_binding_cli(self):
|
||||
record = self.upload(name="../../etc/evil\n.pt")
|
||||
self.assertEqual(record["label"], "evil_.pt")
|
||||
bound = record["initialization"]
|
||||
directory = self.registry.verify(bound)
|
||||
self.assertEqual(
|
||||
{p.name for p in directory.iterdir()},
|
||||
{"upload.pt", "actor.pt", "upload.json", "label.json"},
|
||||
)
|
||||
self.assertEqual(self.upload(name="renamed.pt"), record)
|
||||
restored = PretrainedSources(None, self.root / "store", sys.executable, ROOT / "rl")
|
||||
self.assertEqual(restored.catalog(), [record])
|
||||
self.assertEqual(restored.bind(record["id"], "Unitree-Go2-Flat"), bound)
|
||||
args = restored.arguments(bound)
|
||||
self.assertIn("--pretrained-upload-manifest", args)
|
||||
self.assertNotIn("--resume-checkpoint", args)
|
||||
manager = TrainingManager(
|
||||
ROOT / "rl",
|
||||
sys.executable,
|
||||
("Unitree-Go2-Flat",),
|
||||
check_environment=False,
|
||||
sources=restored,
|
||||
)
|
||||
cfg = manager.parse_config(
|
||||
{
|
||||
"taskId": "Unitree-Go2-Flat",
|
||||
"numEnvs": 4096,
|
||||
"maxIterations": 1,
|
||||
"seed": 42,
|
||||
"device": "cpu",
|
||||
"pretrainedSourceId": record["id"],
|
||||
}
|
||||
)
|
||||
self.assertEqual(cfg.pretrained, bound)
|
||||
self.assertIn("--pretrained-upload-manifest", manager.command_for(cfg))
|
||||
from scripts.train import TrainConfig, _load_pretrained
|
||||
from src.tasks.velocity.config.go2.env_cfgs import unitree_go2_flat_env_cfg
|
||||
from src.tasks.velocity.config.go2.rl_cfg import unitree_go2_ppo_runner_cfg
|
||||
|
||||
train = TrainConfig(
|
||||
env=unitree_go2_flat_env_cfg(),
|
||||
agent=unitree_go2_ppo_runner_cfg(),
|
||||
pretrained_checkpoint=str(directory / "actor.pt"),
|
||||
pretrained_upload_manifest=str(directory / "upload.json"),
|
||||
pretrained_allowed_roots=[str(directory)],
|
||||
pretrained_source_id=record["id"],
|
||||
)
|
||||
source = _load_pretrained(train)
|
||||
self.assertEqual(source.manifest, bound["manifest"])
|
||||
from pretrained import public_initialization_metadata
|
||||
|
||||
public = public_initialization_metadata(source.manifest)
|
||||
self.assertEqual(public["uploadSha256"], hashlib.sha256(self.pt).hexdigest())
|
||||
self.assertEqual(public["templateId"], "go2-legacy47-v1")
|
||||
self.assertNotIn(str(self.root), json.dumps(public))
|
||||
# Changing upload content never changes the previously bound job.
|
||||
newer = self.upload(self.onnx, "onnx")
|
||||
self.assertNotEqual(newer["id"], record["id"])
|
||||
self.assertEqual(restored.arguments(cfg.pretrained), args)
|
||||
(directory / "actor.pt").chmod(0o644)
|
||||
(directory / "actor.pt").write_bytes(b"tamper")
|
||||
with self.assertRaisesRegex(SourceError, "SHA"):
|
||||
restored.arguments(bound)
|
||||
broken = PretrainedSources(None, self.root / "store", sys.executable, ROOT / "rl")
|
||||
self.assertFalse(next(r for r in broken.catalog() if r["id"] == record["id"])["ready"])
|
||||
with self.assertRaises(SourceError):
|
||||
broken.bind(record["id"], "Unitree-Go2-Flat")
|
||||
|
||||
def test_tuning_same_binding_new_trial_and_own_rung_resume(self):
|
||||
from test_tuning_manager import BASE_METRICS, FakeAdvisor
|
||||
from tuning.manager import TuningManager
|
||||
from tuning.process import GpuLease
|
||||
from tuning.schema import BASE_REWARD_CONFIGURATION
|
||||
|
||||
record = self.upload(self.onnx, "onnx")
|
||||
manager = TuningManager(
|
||||
ROOT / "rl",
|
||||
sys.executable,
|
||||
self.root / "tuning",
|
||||
GpuLease(),
|
||||
advisor=FakeAdvisor(),
|
||||
sources=self.registry,
|
||||
)
|
||||
_, config, objective, fallback = manager.parse_create(
|
||||
{
|
||||
"mode": "approval",
|
||||
"pretrainedSourceId": record["id"],
|
||||
"trialCount": 2,
|
||||
"numEnvs": 4096,
|
||||
"initialIterations": 1,
|
||||
"middleIterations": 2,
|
||||
"finalIterations": 3,
|
||||
}
|
||||
)
|
||||
session = manager.storage.create_session("approval", config, objective, fallback)
|
||||
self.assertEqual(config["pretrained"], record["initialization"])
|
||||
from pretrained import public_initialization_metadata
|
||||
|
||||
public = public_initialization_metadata(config["pretrained"]["manifest"])
|
||||
self.assertEqual(public["sourceFormat"], "onnx")
|
||||
self.assertIsNone(public["sourceIteration"])
|
||||
self.assertEqual(public["derivedFields"]["normalizer_count"]["value"], 1_000_000)
|
||||
self.assertNotIn(
|
||||
"onnxSha256", public
|
||||
) # Do not confuse generated checkpoint with user upload.
|
||||
manager.cancel_events[session["id"]] = threading.Event()
|
||||
commands = []
|
||||
|
||||
def run(_session, command, _cwd, _env, _log):
|
||||
commands.append(command)
|
||||
if "--output-dir" in command:
|
||||
directory = Path(command[command.index("--output-dir") + 1])
|
||||
(directory / "model_0.pt").write_bytes(b"own checkpoint with updated statistics")
|
||||
(directory / "policy.onnx").write_bytes(b"trial policy")
|
||||
(directory / "initialization.json").write_text(
|
||||
json.dumps(record["initialization"]["manifest"])
|
||||
)
|
||||
else:
|
||||
Path(command[command.index("--output") + 1]).write_text(
|
||||
json.dumps({"metrics": BASE_METRICS})
|
||||
)
|
||||
return 0
|
||||
|
||||
with patch.object(manager, "_run_command", run):
|
||||
parent = None
|
||||
for number, rung in ((0, 0), (1, 0), (1, 1)):
|
||||
trial = manager.storage.create_trial(
|
||||
session["id"],
|
||||
number,
|
||||
rung,
|
||||
rung + 1,
|
||||
BASE_REWARD_CONFIGURATION,
|
||||
None,
|
||||
f"trial-{number}-{rung}",
|
||||
)
|
||||
result = manager._execute_trial(session, trial, parent if rung else None)
|
||||
parent = manager._session_root(session["id"]) / result["checkpointPath"]
|
||||
train = [c for c in commands if "--output-dir" in c]
|
||||
for command in train[:2]:
|
||||
self.assertIn("--pretrained-upload-manifest", command)
|
||||
self.assertEqual(command[command.index("--pretrained-source-id") + 1], record["id"])
|
||||
self.assertNotIn("--pretrained-upload-manifest", train[2])
|
||||
self.assertNotIn("--pretrained-checkpoint", train[2])
|
||||
self.assertIn("trial-1-0/model_0.pt", train[2][-1])
|
||||
restored_sources = PretrainedSources(None, self.registry.store, sys.executable, ROOT / "rl")
|
||||
restored = TuningManager(
|
||||
ROOT / "rl",
|
||||
sys.executable,
|
||||
self.root / "tuning",
|
||||
GpuLease(),
|
||||
advisor=FakeAdvisor(),
|
||||
sources=restored_sources,
|
||||
)
|
||||
self.assertEqual(
|
||||
restored.detail(session["id"])["config"]["pretrained"], record["initialization"]
|
||||
)
|
||||
self.assertFalse(restored.workers)
|
||||
|
||||
def test_input_limits_disconnect_timeout_capacity_and_cleanup(self):
|
||||
for length, fmt, template in (
|
||||
(0, "pt", "go2-legacy47-v1"),
|
||||
(256 * 1024**2 + 1, "pt", "go2-legacy47-v1"),
|
||||
(64 * 1024**2 + 1, "onnx", "go2-legacy47-v1"),
|
||||
(1, "zip", "go2-legacy47-v1"),
|
||||
(1, "pt", ""),
|
||||
):
|
||||
with (
|
||||
self.subTest(length=length, fmt=fmt, template=template),
|
||||
self.assertRaises(SourceError),
|
||||
):
|
||||
self.registry.receive_upload(io.BytesIO(b"x"), length, fmt, template, "x")
|
||||
with self.assertRaisesRegex(SourceError, "中断"):
|
||||
self.registry.receive_upload(io.BytesIO(b"x"), 100, "pt", "go2-legacy47-v1", "x")
|
||||
with (
|
||||
patch("pretrained_sources.time.monotonic", side_effect=[0, 61]),
|
||||
self.assertRaisesRegex(SourceError, "超时"),
|
||||
):
|
||||
self.upload()
|
||||
trickle = unittest.mock.Mock()
|
||||
trickle.read.side_effect = AssertionError("buffered read could bypass total deadline")
|
||||
trickle.read1.return_value = b"x"
|
||||
with (
|
||||
patch("pretrained_sources.time.monotonic", side_effect=[0, 1, 61]),
|
||||
self.assertRaisesRegex(SourceError, "超时"),
|
||||
):
|
||||
self.registry.receive_upload(trickle, 10, "pt", "go2-legacy47-v1", "x")
|
||||
self.assertEqual(trickle.read1.call_count, 1)
|
||||
stream = unittest.mock.Mock()
|
||||
stream.read1.side_effect = TimeoutError()
|
||||
with self.assertRaises(SourceError):
|
||||
self.registry.receive_upload(stream, 10, "pt", "go2-legacy47-v1", "x")
|
||||
with self.registry.upload_slot(1), self.assertRaisesRegex(SourceError, "正在上传"):
|
||||
self.upload()
|
||||
self.assertEqual(list(self.registry.store.iterdir()), [])
|
||||
import subprocess
|
||||
|
||||
with (
|
||||
patch(
|
||||
"pretrained_sources.subprocess.run",
|
||||
side_effect=subprocess.TimeoutExpired("validator", 60),
|
||||
),
|
||||
self.assertRaisesRegex(SourceError, "超时"),
|
||||
):
|
||||
self.upload()
|
||||
self.assertEqual(list(self.registry.store.iterdir()), [])
|
||||
huge = self.registry.store / "quota"
|
||||
with huge.open("wb") as stream:
|
||||
stream.truncate(2 * 1024**3)
|
||||
with self.assertRaisesRegex(SourceError, "上限"):
|
||||
self.upload()
|
||||
huge.unlink()
|
||||
for n in range(32):
|
||||
(self.registry.store / f"{n:064x}").mkdir()
|
||||
with self.assertRaisesRegex(SourceError, "上限"):
|
||||
self.upload()
|
||||
|
||||
def test_unsafe_pickle_invalid_checkpoint_and_zip_never_publish(self):
|
||||
import torch
|
||||
|
||||
marker = self.root / "unsafe-executed"
|
||||
|
||||
class Unsafe:
|
||||
def __reduce__(self):
|
||||
return (os.system, (f"touch {marker}",))
|
||||
|
||||
bad = dict(self.state)
|
||||
bad["mlp.0.weight"] = torch.full((512, 47), float("nan"))
|
||||
wrong = dict(self.state)
|
||||
wrong["mlp.0.weight"] = torch.zeros(512, 97)
|
||||
cases = [b"PK\x03\x04bad zip", b"not a model"]
|
||||
for payload in (
|
||||
{"actor_state_dict": bad},
|
||||
{"actor_state_dict": wrong},
|
||||
{"actor_state_dict": self.state, "metadata": {"joint_names": ["wrong"]}},
|
||||
{"actor_state_dict": self.state, "evil": Unsafe()},
|
||||
):
|
||||
stream = io.BytesIO()
|
||||
torch.save(payload, stream)
|
||||
cases.append(stream.getvalue())
|
||||
for data in cases:
|
||||
with self.subTest(size=len(data)), self.assertRaises(SourceError):
|
||||
self.upload(data)
|
||||
self.assertEqual(self.registry.catalog(), [])
|
||||
self.assertEqual(list(self.registry.store.iterdir()), [])
|
||||
self.assertFalse((self.root / "unsafe-executed").exists())
|
||||
|
||||
def test_strict_graph_rejects_bad_edges_ops_attributes_external_and_metadata(self):
|
||||
import onnx
|
||||
from pretrained import PretrainedError
|
||||
from pretrained_upload import _onnx
|
||||
|
||||
base = onnx.load_model_from_string(self.onnx)
|
||||
|
||||
def external(m):
|
||||
m.graph.initializer[0].data_location = onnx.TensorProto.EXTERNAL
|
||||
m.graph.initializer[0].external_data.add(key="location", value="/etc/passwd")
|
||||
|
||||
def attrs(m):
|
||||
m.graph.node[2].attribute.add(name="transA", type=onnx.AttributeProto.INT, i=1)
|
||||
|
||||
def metadata(m):
|
||||
next(p for p in m.metadata_props if p.key == "joint_names").value = "wrong"
|
||||
|
||||
def extra(m):
|
||||
m.graph.node.append(onnx.helper.make_node("Identity", ["actions"], ["branch"]))
|
||||
|
||||
def shape(m):
|
||||
m.graph.input[0].type.tensor_type.shape.dim[1].dim_value = 97
|
||||
|
||||
edits = [
|
||||
external,
|
||||
attrs,
|
||||
metadata,
|
||||
extra,
|
||||
shape,
|
||||
lambda m: setattr(m.graph.node[2], "domain", "evil"),
|
||||
lambda m: setattr(m.graph.node[3], "op_type", "Relu"),
|
||||
lambda m: m.graph.node[4].input.__setitem__(0, "normalized"),
|
||||
lambda m: m.graph.node[0].input.reverse(),
|
||||
lambda m: m.graph.node[4].input.__setitem__(1, "mlp.0.weight"),
|
||||
lambda m: setattr(m.graph.output[0], "name", "linear0"),
|
||||
lambda m: setattr(m.graph.initializer[0], "data_type", onnx.TensorProto.DOUBLE),
|
||||
]
|
||||
for edit in edits:
|
||||
model = copy.deepcopy(base)
|
||||
edit(model)
|
||||
with self.subTest(edit=edit), self.assertRaises((PretrainedError, ValueError)):
|
||||
_onnx(model.SerializeToString())
|
||||
bad = copy.deepcopy(base)
|
||||
external(bad)
|
||||
with self.assertRaisesRegex(SourceError, "external"):
|
||||
self.upload(bad.SerializeToString(), "onnx")
|
||||
self.assertEqual(self.registry.catalog(), [])
|
||||
|
||||
def test_http_upload_storage_errors_are_503_but_read_failures_are_400(self):
|
||||
manager = TrainingManager(
|
||||
ROOT / "rl",
|
||||
sys.executable,
|
||||
("Unitree-Go2-Flat",),
|
||||
check_environment=False,
|
||||
sources=self.registry,
|
||||
)
|
||||
reset_timeouts = []
|
||||
|
||||
class Handler(TrainingRequestHandler):
|
||||
access_token = "upload-test-token"
|
||||
read_error = None
|
||||
|
||||
def log_message(self, *_args):
|
||||
pass
|
||||
|
||||
def _upload(self):
|
||||
original_stream = self.rfile
|
||||
if self.read_error is not None:
|
||||
self.rfile = unittest.mock.Mock(wraps=original_stream)
|
||||
self.rfile.read1.side_effect = self.read_error
|
||||
try:
|
||||
super()._upload()
|
||||
finally:
|
||||
self.rfile = original_stream
|
||||
reset_timeouts.append(self.connection.gettimeout())
|
||||
|
||||
Handler.manager = manager
|
||||
server = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
|
||||
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
||||
thread.start()
|
||||
original_open = Path.open
|
||||
url = (
|
||||
"/api/training/pretrained-sources/upload"
|
||||
"?format=pt&template=go2-legacy47-v1&name=policy.pt"
|
||||
)
|
||||
|
||||
try:
|
||||
for phase, error in (
|
||||
("open", OSError(errno.ENOSPC, "disk full", "/private/upload.pt")),
|
||||
("open", OSError(errno.EACCES, "permission denied", "/private/upload.pt")),
|
||||
("write", OSError(errno.ENOSPC, "disk full", "/private/upload.pt")),
|
||||
("close", OSError(errno.ENOSPC, "flush failed", "/private/upload.pt")),
|
||||
("read", ConnectionResetError(errno.ECONNRESET, "connection reset")),
|
||||
("read", TimeoutError("read timeout")),
|
||||
("truncated", None),
|
||||
):
|
||||
with self.subTest(phase=phase, error=error):
|
||||
Handler.read_error = error if phase == "read" else None
|
||||
|
||||
@contextmanager
|
||||
def faulty_output(path, *args, phase=phase, error=error, **kwargs):
|
||||
with original_open(path, *args, **kwargs) as output:
|
||||
if phase == "write":
|
||||
output = unittest.mock.Mock(wraps=output)
|
||||
output.write.side_effect = error
|
||||
yield output
|
||||
if phase == "close":
|
||||
raise error
|
||||
|
||||
def open_file(
|
||||
path,
|
||||
*args,
|
||||
phase=phase,
|
||||
error=error,
|
||||
output_factory=faulty_output,
|
||||
**kwargs,
|
||||
):
|
||||
if path.name == "upload.pt" and args == ("xb",):
|
||||
if phase == "open":
|
||||
raise error
|
||||
return output_factory(path, *args, **kwargs)
|
||||
return original_open(path, *args, **kwargs)
|
||||
|
||||
conn = http.client.HTTPConnection(*server.server_address, timeout=10)
|
||||
try:
|
||||
with patch.object(Path, "open", open_file):
|
||||
conn.request(
|
||||
"POST",
|
||||
url,
|
||||
body=b"x",
|
||||
headers={
|
||||
"Authorization": "Bearer upload-test-token",
|
||||
"Content-Type": "application/octet-stream",
|
||||
"Content-Length": "2" if phase == "truncated" else "1",
|
||||
},
|
||||
)
|
||||
if phase == "truncated":
|
||||
conn.sock.shutdown(socket.SHUT_WR)
|
||||
response = conn.getresponse()
|
||||
message = json.loads(response.read())["error"]
|
||||
finally:
|
||||
conn.close()
|
||||
storage_error = phase in ("open", "write", "close")
|
||||
self.assertEqual(response.status, 503 if storage_error else 400)
|
||||
self.assertIn("磁盘空间/权限" if storage_error else "中断", message)
|
||||
self.assertNotIn("/private", message)
|
||||
self.assertNotIn(str(self.root), message)
|
||||
self.assertEqual(reset_timeouts[-1], 10)
|
||||
self.assertEqual(list(self.registry.store.iterdir()), [])
|
||||
self.assertEqual(self.registry.catalog(), [])
|
||||
self.assertEqual(manager.jobs, {})
|
||||
finally:
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
thread.join()
|
||||
|
||||
def test_http_auth_json_limit_binary_headers_and_incomplete_body(self):
|
||||
manager = TrainingManager(
|
||||
ROOT / "rl",
|
||||
sys.executable,
|
||||
("Unitree-Go2-Flat",),
|
||||
check_environment=False,
|
||||
sources=self.registry,
|
||||
)
|
||||
|
||||
class Handler(TrainingRequestHandler):
|
||||
access_token = "upload-test-token"
|
||||
|
||||
def log_message(self, *_args):
|
||||
pass
|
||||
|
||||
Handler.manager = manager
|
||||
server = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
|
||||
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
||||
thread.start()
|
||||
url = (
|
||||
"/api/training/pretrained-sources/upload"
|
||||
"?format=pt&template=go2-legacy47-v1&name=..%2Ftest.pt"
|
||||
)
|
||||
headers = {
|
||||
"Authorization": "Bearer upload-test-token",
|
||||
"Content-Type": "application/octet-stream",
|
||||
}
|
||||
|
||||
def request(path=url, data=b"x", extra=None):
|
||||
conn = http.client.HTTPConnection(*server.server_address, timeout=30)
|
||||
conn.request("POST", path, body=data, headers={**headers, **(extra or {})})
|
||||
response = conn.getresponse()
|
||||
status, body = response.status, response.read()
|
||||
conn.close()
|
||||
return status, json.loads(body)
|
||||
|
||||
try:
|
||||
self.assertEqual(request(extra={"Authorization": "Bearer wrong"})[0], 401)
|
||||
self.assertEqual(request(extra={"Origin": "https://evil.example"})[0], 403)
|
||||
self.assertEqual(request(extra={"Host": "evil.example"})[0], 403)
|
||||
self.assertEqual(request(extra={"Content-Type": "application/json"})[0], 400)
|
||||
self.assertEqual(request(extra={"Content-Length": str(256 * 1024**2 + 1)})[0], 413)
|
||||
self.assertEqual(request(extra={"Content-Encoding": "gzip"})[0], 400)
|
||||
self.assertEqual(request(extra={"Transfer-Encoding": "chunked"})[0], 400)
|
||||
self.assertEqual(
|
||||
request(path=url.replace("template=go2-legacy47-v1", "template="))[0], 400
|
||||
)
|
||||
self.assertEqual(
|
||||
request(path="/api/training/jobs", data=b" " * (128 * 1024 + 1))[0], 413
|
||||
)
|
||||
# A blocked decoder is isolated from health/catalog requests.
|
||||
import subprocess
|
||||
from types import SimpleNamespace
|
||||
|
||||
Handler.tuning_manager = SimpleNamespace(capability=lambda: {})
|
||||
started, release = threading.Event(), threading.Event()
|
||||
original_run = subprocess.run
|
||||
|
||||
def blocked(*args, **kwargs):
|
||||
started.set()
|
||||
release.wait(10)
|
||||
return original_run(*args, **kwargs)
|
||||
|
||||
outcomes = []
|
||||
with patch("pretrained_sources.subprocess.run", side_effect=blocked):
|
||||
worker = threading.Thread(target=lambda: outcomes.append(request(data=self.pt)))
|
||||
worker.start()
|
||||
self.assertTrue(started.wait(5))
|
||||
health = http.client.HTTPConnection(*server.server_address, timeout=3)
|
||||
health.request("GET", "/api/training/health", headers=headers)
|
||||
response = health.getresponse()
|
||||
self.assertEqual(response.status, 200)
|
||||
self.assertTrue(json.loads(response.read())["pretrainedUpload"]["enabled"])
|
||||
health.close()
|
||||
release.set()
|
||||
worker.join(30)
|
||||
status, record = outcomes[0]
|
||||
self.assertEqual(status, 201)
|
||||
self.assertEqual(record["label"], "test.pt")
|
||||
status, onnx_record = request(
|
||||
path=url.replace("format=pt", "format=onnx"), data=self.onnx
|
||||
)
|
||||
self.assertEqual(status, 201)
|
||||
self.assertEqual(onnx_record["initialization"]["manifest"]["sourceFormat"], "onnx")
|
||||
raw = socket.create_connection(server.server_address, timeout=10)
|
||||
raw.sendall(
|
||||
(
|
||||
f"POST {url} HTTP/1.0\r\nHost: localhost\r\n"
|
||||
"Authorization: Bearer upload-test-token\r\n"
|
||||
"Content-Type: application/octet-stream\r\nContent-Length: 1000\r\n\r\nx"
|
||||
).encode()
|
||||
)
|
||||
raw.shutdown(socket.SHUT_WR)
|
||||
self.assertIn(b"400", raw.recv(4096))
|
||||
raw.close()
|
||||
self.assertFalse(list(self.registry.store.glob("upload-*")))
|
||||
finally:
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
thread.join()
|
||||
|
||||
def test_onnx_expansion_and_checkpoint_preserve_updated_count(self):
|
||||
import torch
|
||||
from pretrained import (
|
||||
ValidatedSource,
|
||||
comparison_observations,
|
||||
make_reference_actor,
|
||||
warm_start_actor,
|
||||
)
|
||||
from pretrained_upload import _onnx
|
||||
from tensordict import TensorDict
|
||||
|
||||
state, _, identity = _onnx(self.onnx)
|
||||
self.assertLess(identity["max_abs_error"], 2e-5)
|
||||
base = comparison_observations()
|
||||
reference = make_reference_actor().eval()
|
||||
reference.load_state_dict(state)
|
||||
expected = reference.mlp(reference.obs_normalizer(base)).detach()
|
||||
for dim in (81, 97):
|
||||
actor = make_reference_actor(dim)
|
||||
warm_start_actor(actor, ValidatedSource(state, {}))
|
||||
self.assertEqual(actor.mlp[0].weight[:, 47:].count_nonzero().item(), 0)
|
||||
values = torch.cat((base, torch.ones(len(base), dim - 47)), dim=1)
|
||||
torch.testing.assert_close(
|
||||
actor.mlp(actor.obs_normalizer(values)), expected, atol=2e-5, rtol=2e-5
|
||||
)
|
||||
batch = TensorDict({"actor": values}, batch_size=[len(base)])
|
||||
actor.update_normalization(batch)
|
||||
self.assertEqual(actor.obs_normalizer.count.item(), 1_000_048)
|
||||
self.assertGreater(actor.obs_normalizer._mean[:, 47:].min().item(), 0)
|
||||
actor(batch).square().mean().backward()
|
||||
grad = actor.mlp[0].weight.grad[:, 47:]
|
||||
self.assertTrue(torch.isfinite(grad).all())
|
||||
self.assertGreater(grad.abs().max().item(), 0)
|
||||
torch.optim.Adam(actor.parameters(), lr=1e-4).step()
|
||||
buffer = io.BytesIO()
|
||||
torch.save(actor.state_dict(), buffer)
|
||||
buffer.seek(0)
|
||||
restored = make_reference_actor(dim)
|
||||
restored.load_state_dict(torch.load(buffer, weights_only=True))
|
||||
for key, tensor in actor.state_dict().items():
|
||||
self.assertTrue(torch.equal(tensor, restored.state_dict()[key]), key)
|
||||
|
||||
|
||||
@unittest.skipUnless(
|
||||
os.environ.get("GO2_UPLOAD_REAL_DIR"), "opt-in read-only real single-file probes"
|
||||
)
|
||||
class RealSingleFileUploadTest(UploadTest):
|
||||
"""Inherited tests run against real uploads without ever requesting a sidecar."""
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
import torch
|
||||
|
||||
torch.set_num_threads(1)
|
||||
root = Path(os.environ["GO2_UPLOAD_REAL_DIR"])
|
||||
# Explicit two files read separately as bytes; each upload receives only one.
|
||||
cls.pt = (root / "model_10000.pt").read_bytes()
|
||||
cls.onnx = (root / "policy.onnx").read_bytes()
|
||||
cls.state = torch.load(io.BytesIO(cls.pt), map_location="cpu", weights_only=True)[
|
||||
"actor_state_dict"
|
||||
]
|
||||
|
||||
def test_real_independent_uploads_ort_and_batch_drift(self):
|
||||
import torch
|
||||
from pretrained import comparison_observations, make_reference_actor, verify_onnx
|
||||
from pretrained_upload import read_uploaded_source
|
||||
from src.tasks.velocity.config.go2.env_cfgs import unitree_go2_flat_env_cfg
|
||||
from src.tasks.velocity.config.go2.rl_cfg import unitree_go2_ppo_runner_cfg
|
||||
|
||||
evidence = {}
|
||||
for fmt, data in (("pt", self.pt), ("onnx", self.onnx)):
|
||||
with tempfile.TemporaryDirectory() as empty:
|
||||
registry = PretrainedSources(None, Path(empty), sys.executable, ROOT / "rl")
|
||||
record = registry.receive_upload(
|
||||
io.BytesIO(data), len(data), fmt, "go2-legacy47-v1", f"single.{fmt}"
|
||||
)
|
||||
directory = registry.verify(record["initialization"])
|
||||
source = read_uploaded_source(
|
||||
directory / "actor.pt",
|
||||
allowed_roots=[directory],
|
||||
manifest_path=directory / "upload.json",
|
||||
target_env=asdict(unitree_go2_flat_env_cfg()),
|
||||
target_agent=asdict(unitree_go2_ppo_runner_cfg()),
|
||||
)
|
||||
actor = make_reference_actor().eval()
|
||||
actor.load_state_dict(source.actor_state)
|
||||
identity = verify_onnx(self.onnx, actor)
|
||||
self.assertLess(identity["max_abs_error"], 2e-5)
|
||||
evidence[fmt] = {
|
||||
"sourceId": record["id"],
|
||||
"sha256": hashlib.sha256(data).hexdigest(),
|
||||
"identity": identity,
|
||||
"normalizer_count": actor.obs_normalizer.count.item(),
|
||||
"updates": {},
|
||||
}
|
||||
probes = comparison_observations()[32:]
|
||||
before = actor.mlp(actor.obs_normalizer(probes)).detach()
|
||||
# Ordinary UI default=4096 and service maximum=16384 environments.
|
||||
for batch_size in (4096, 16384):
|
||||
updated = copy.deepcopy(actor).train()
|
||||
updated.obs_normalizer.update(probes.repeat(batch_size // len(probes), 1))
|
||||
after = updated.mlp(updated.obs_normalizer(probes)).detach()
|
||||
evidence[fmt]["updates"][str(batch_size)] = {
|
||||
"rate": batch_size / updated.obs_normalizer.count.item(),
|
||||
"max_action_drift": (after - before).abs().max().item(),
|
||||
"max_mean_drift": (
|
||||
updated.obs_normalizer._mean - actor.obs_normalizer._mean
|
||||
)
|
||||
.abs()
|
||||
.max()
|
||||
.item(),
|
||||
"count_after": updated.obs_normalizer.count.item(),
|
||||
}
|
||||
self.assertTrue(torch.isfinite(after).all())
|
||||
self.assertEqual(
|
||||
PretrainedSources(None, Path(empty), sys.executable, ROOT / "rl").catalog(),
|
||||
[record],
|
||||
)
|
||||
print("REAL_UPLOAD_EVIDENCE=" + json.dumps(evidence, sort_keys=True))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,163 @@
|
||||
"""Opt-in real CPU runner integration: four environments, one ONNX-derived PPO iteration."""
|
||||
|
||||
import copy
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
from dataclasses import asdict
|
||||
from pathlib import Path
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[1]
|
||||
for root in (ROOT, ROOT / "rl"):
|
||||
sys.path.insert(0, str(root))
|
||||
|
||||
|
||||
@unittest.skipUnless(os.environ.get("GO2_UPLOAD_RUNNER_OUTPUT"), "opt-in four-env CPU runner smoke")
|
||||
class UploadedRunnerTest(unittest.TestCase):
|
||||
def test_real_upload_runner_ppo_export_and_same_trial_restore(self):
|
||||
import numpy as np
|
||||
import onnxruntime as ort
|
||||
import torch
|
||||
import warp as wp
|
||||
|
||||
if not hasattr(wp, "context"):
|
||||
from warp._src import context
|
||||
|
||||
wp.context = context
|
||||
from mjlab.envs import ManagerBasedRlEnv
|
||||
from mjlab.rl import RslRlVecEnvWrapper
|
||||
from pretrained import initialize_runner, make_reference_actor, validate_runtime_contract
|
||||
from pretrained_sources import PretrainedSources
|
||||
from pretrained_upload import read_uploaded_source
|
||||
from src.tasks.obstacle_avoidance.env_cfg import unitree_go2_obstacle_env_cfg
|
||||
from src.tasks.velocity.config.go2.rl_cfg import unitree_go2_ppo_runner_cfg
|
||||
from src.tasks.velocity.rl.runner import VelocityOnPolicyRunner
|
||||
from task_config import OBSTACLE_TASK, deployment_metadata, validate_task_config
|
||||
|
||||
torch.set_num_threads(1)
|
||||
source_root = Path(os.environ["GO2_UPLOAD_REAL_DIR"])
|
||||
output = Path(os.environ["GO2_UPLOAD_RUNNER_OUTPUT"])
|
||||
output.mkdir(parents=True, exist_ok=True)
|
||||
# Only these explicit files; each decoder sees its independent empty root.
|
||||
original_onnx = (source_root / "policy.onnx").read_bytes()
|
||||
options = ort.SessionOptions()
|
||||
options.intra_op_num_threads = options.inter_op_num_threads = 1
|
||||
original = ort.InferenceSession(original_onnx, options, providers=["CPUExecutionProvider"])
|
||||
evidence = {}
|
||||
for fmt, data in (
|
||||
("pt", (source_root / "model_10000.pt").read_bytes()),
|
||||
("onnx", original_onnx),
|
||||
):
|
||||
with tempfile.TemporaryDirectory() as store:
|
||||
registry = PretrainedSources(None, Path(store), sys.executable, ROOT / "rl")
|
||||
record = registry.receive_upload(
|
||||
io.BytesIO(data), len(data), fmt, "go2-legacy47-v1", f"single.{fmt}"
|
||||
)
|
||||
bound = record["initialization"]
|
||||
directory = registry.verify(bound)
|
||||
cfg = unitree_go2_obstacle_env_cfg()
|
||||
cfg.scene.num_envs = 4
|
||||
cfg.seed = 42
|
||||
agent = unitree_go2_ppo_runner_cfg()
|
||||
agent.logger = "tensorboard"
|
||||
agent.max_iterations = 1
|
||||
source = read_uploaded_source(
|
||||
directory / "actor.pt",
|
||||
allowed_roots=[directory],
|
||||
manifest_path=directory / "upload.json",
|
||||
target_env=asdict(cfg),
|
||||
target_agent=asdict(agent),
|
||||
)
|
||||
raw = ManagerBasedRlEnv(cfg, device="cpu")
|
||||
env = RslRlVecEnvWrapper(raw)
|
||||
try:
|
||||
validate_runtime_contract(raw)
|
||||
raw.platform_deployment = deployment_metadata(
|
||||
OBSTACLE_TASK, validate_task_config(OBSTACLE_TASK, {}, 42), 42
|
||||
)
|
||||
log = output / fmt
|
||||
log.mkdir(exist_ok=True)
|
||||
runner = VelocityOnPolicyRunner(env, asdict(agent), str(log), "cpu")
|
||||
raw.platform_initialization = initialize_runner(runner, source)
|
||||
(log / "initialization.json").write_text(
|
||||
json.dumps(raw.platform_initialization)
|
||||
)
|
||||
actor = runner.alg.actor
|
||||
obs = env.get_observations()
|
||||
reference = make_reference_actor().eval()
|
||||
reference.load_state_dict(source.actor_state)
|
||||
with torch.no_grad():
|
||||
before = actor(obs)
|
||||
expected = reference.mlp(reference.obs_normalizer(obs["actor"][:, :47]))
|
||||
torch.testing.assert_close(before, expected, atol=2e-5, rtol=2e-5)
|
||||
onnx_actions = np.concatenate(
|
||||
[
|
||||
original.run(None, {"obs": row[None].numpy()})[0]
|
||||
for row in obs["actor"][:, :47]
|
||||
]
|
||||
)
|
||||
np.testing.assert_allclose(before.numpy(), onnx_actions, atol=2e-5, rtol=2e-5)
|
||||
self.assertEqual(actor.mlp[0].weight[:, 47:].count_nonzero().item(), 0)
|
||||
self.assertFalse(runner.alg.optimizer.state)
|
||||
actor(obs).square().mean().backward()
|
||||
gradient = actor.mlp[0].weight.grad[:, 47:]
|
||||
self.assertTrue(torch.isfinite(gradient).all())
|
||||
self.assertGreater(gradient.abs().max().item(), 0)
|
||||
evidence[fmt] = {
|
||||
"source_id": record["id"],
|
||||
"real_observation_ort_error": float(
|
||||
np.abs(before.numpy() - onnx_actions).max()
|
||||
),
|
||||
"new_column_gradient_max": gradient.abs().max().item(),
|
||||
}
|
||||
runner.alg.optimizer.zero_grad()
|
||||
if fmt == "pt":
|
||||
continue # Only ONNX-derived branch runs the one approved PPO iteration.
|
||||
runner.learn(num_learning_iterations=1, init_at_random_ep_len=True)
|
||||
self.assertGreater(actor.mlp[0].weight[:, 47:].abs().max().item(), 0)
|
||||
saved = torch.load(log / "model_0.pt", map_location="cpu", weights_only=True)
|
||||
updated = copy.deepcopy(actor.state_dict())
|
||||
resumed = VelocityOnPolicyRunner(env, asdict(agent), str(log), "cpu")
|
||||
self.assertTrue(resumed.alg.load(saved, None, strict=True))
|
||||
self.assertTrue(resumed.alg.optimizer.state)
|
||||
for key, tensor in updated.items():
|
||||
self.assertTrue(
|
||||
torch.equal(tensor, resumed.alg.actor.state_dict()[key]), key
|
||||
)
|
||||
self.assertGreater(resumed.alg.actor.obs_normalizer.count.item(), 1_000_000)
|
||||
exported = ort.InferenceSession(
|
||||
str(log / "policy.onnx"), options, providers=["CPUExecutionProvider"]
|
||||
)
|
||||
metadata = json.loads(
|
||||
exported.get_modelmeta().custom_metadata_map["pretrained_initialization"]
|
||||
)
|
||||
self.assertEqual(metadata["sourceFormat"], "onnx")
|
||||
self.assertEqual(
|
||||
metadata["uploadSha256"], source.manifest["artifacts"]["upload"]["sha256"]
|
||||
)
|
||||
self.assertNotIn(str(source_root), json.dumps(metadata))
|
||||
final_obs = env.get_observations()
|
||||
with torch.no_grad():
|
||||
expected = resumed.alg.actor(final_obs).numpy()
|
||||
actual = np.concatenate(
|
||||
[
|
||||
exported.run(None, {"obs": row[None].numpy()})[0]
|
||||
for row in final_obs["actor"]
|
||||
]
|
||||
)
|
||||
np.testing.assert_allclose(actual, expected, atol=2e-5, rtol=2e-5)
|
||||
evidence[fmt].update(
|
||||
ppo_iterations=1,
|
||||
num_envs=4,
|
||||
count_restored=resumed.alg.actor.obs_normalizer.count.item(),
|
||||
export_max_error=float(np.abs(actual - expected).max()),
|
||||
new_columns_max_after_ppo=actor.mlp[0].weight[:, 47:].abs().max().item(),
|
||||
source_metadata=metadata,
|
||||
)
|
||||
finally:
|
||||
env.close()
|
||||
(output / "evidence.json").write_text(json.dumps(evidence, indent=2))
|
||||
print("REAL_RUNNER_UPLOAD_EVIDENCE=" + json.dumps(evidence))
|
||||
@@ -0,0 +1,137 @@
|
||||
"""Persisted preset identity -> resolver -> training validation; no training/API calls."""
|
||||
|
||||
import json
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
from server import ApiError, TrainingManager # noqa: E402
|
||||
from tuning.manager import TuningManager # noqa: E402
|
||||
from tuning.schema import ( # noqa: E402
|
||||
FLAT_TASK,
|
||||
OBSTACLE_TASK,
|
||||
RewardConfigError,
|
||||
base_configuration,
|
||||
)
|
||||
from tuning.storage import TuningStorage # noqa: E402
|
||||
|
||||
|
||||
class RewardPresetTaskTest(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.temp = tempfile.TemporaryDirectory()
|
||||
root = Path(self.temp.name)
|
||||
(root / "scripts").mkdir()
|
||||
(root / "scripts/train.py").write_text("raise AssertionError('must not train')")
|
||||
self.storage = TuningStorage(root / "tuning.sqlite3")
|
||||
# Real resolver with real storage; no advisor, processes or training workers are needed.
|
||||
self.tuning = object.__new__(TuningManager)
|
||||
self.tuning.storage = self.storage
|
||||
self.training = TrainingManager(
|
||||
root, sys.executable, (FLAT_TASK, OBSTACLE_TASK), check_environment=False
|
||||
)
|
||||
self.training.preset_resolver = self.tuning.preset_config
|
||||
|
||||
def tearDown(self):
|
||||
self.storage.connection().close()
|
||||
self.temp.cleanup()
|
||||
|
||||
def save(self, config, reward=None):
|
||||
session = self.storage.create_session("approval", config, {}, False)
|
||||
reward = reward or base_configuration(config.get("taskId", FLAT_TASK))
|
||||
trial = self.storage.create_trial(session["id"], 0, 0, 1, reward, None, "trial")
|
||||
return self.storage.save_preset(session["id"], session["id"], trial["id"], reward)
|
||||
|
||||
def payload(self, preset):
|
||||
return dict(
|
||||
taskId=FLAT_TASK,
|
||||
rewardPresetId=preset["id"],
|
||||
numEnvs=2,
|
||||
maxIterations=1,
|
||||
seed=42,
|
||||
device="cpu",
|
||||
)
|
||||
|
||||
def assert_rejected(self, preset):
|
||||
for entry in (self.training.parse_config, self.training.start):
|
||||
with self.assertRaises(ApiError) as raised:
|
||||
entry(self.payload(preset))
|
||||
self.assertEqual(raised.exception.status, 400)
|
||||
self.assertEqual(self.training.jobs, {})
|
||||
self.assertIsNone(self.training.lease.public())
|
||||
|
||||
def test_obstacle_list_identity_and_cross_task_rejected_before_job_creation(self):
|
||||
preset = self.save({"taskId": OBSTACLE_TASK})
|
||||
self.assertEqual(preset["taskId"], OBSTACLE_TASK)
|
||||
self.assertEqual(self.storage.list_presets(), [preset])
|
||||
self.assertEqual(self.storage.get_preset(preset["id"]), preset)
|
||||
self.assertEqual(
|
||||
self.tuning.preset_config(preset["id"], OBSTACLE_TASK), preset["rewardConfig"]
|
||||
)
|
||||
with self.assertRaisesRegex(RewardConfigError, "任务"):
|
||||
self.tuning.preset_config(preset["id"])
|
||||
self.assert_rejected(preset)
|
||||
|
||||
def test_historical_flat_without_task_id_and_explicit_flat_remain_valid(self):
|
||||
for config in ({}, {"taskId": FLAT_TASK}):
|
||||
with self.subTest(config=config):
|
||||
preset = self.save(config)
|
||||
self.assertEqual(preset["taskId"], FLAT_TASK)
|
||||
self.assertEqual(self.tuning.preset_config(preset["id"]), preset["rewardConfig"])
|
||||
parsed = self.training.parse_config(self.payload(preset))
|
||||
self.assertEqual(parsed.reward_config, base_configuration())
|
||||
self.assertIn("--reward-config-json", self.training.command_for(parsed))
|
||||
self.assertEqual({p["taskId"] for p in self.storage.list_presets()}, {FLAT_TASK})
|
||||
|
||||
def test_malformed_preset_and_corrupt_or_missing_source_fail_closed(self):
|
||||
preset = self.save({})
|
||||
connection = self.storage.connection()
|
||||
valid_reward = json.dumps(preset["rewardConfig"])
|
||||
for bad_reward in (
|
||||
'{"weights":{},"params":{}}',
|
||||
"null",
|
||||
"{broken",
|
||||
json.dumps(base_configuration(OBSTACLE_TASK)),
|
||||
):
|
||||
with self.subTest(reward=bad_reward):
|
||||
connection.execute(
|
||||
"UPDATE presets SET reward_config_json=? WHERE id=?", (bad_reward, preset["id"])
|
||||
)
|
||||
with self.assertRaises(RewardConfigError):
|
||||
self.storage.list_presets()
|
||||
self.assert_rejected(preset)
|
||||
connection.execute(
|
||||
"UPDATE presets SET reward_config_json=? WHERE id=?", (valid_reward, preset["id"])
|
||||
)
|
||||
for source in (
|
||||
"null",
|
||||
"[]",
|
||||
"{broken",
|
||||
'{"taskId":null}',
|
||||
'{"taskId":"unknown"}',
|
||||
json.dumps({"taskId": OBSTACLE_TASK}),
|
||||
):
|
||||
with self.subTest(source=source):
|
||||
connection.execute(
|
||||
"UPDATE sessions SET config_json=? WHERE id=?", (source, preset["sessionId"])
|
||||
)
|
||||
with self.assertRaises(RewardConfigError):
|
||||
self.storage.get_preset(preset["id"])
|
||||
self.assert_rejected(preset)
|
||||
connection.execute("DELETE FROM sessions WHERE id=?", (preset["sessionId"],))
|
||||
with self.assertRaises(RewardConfigError):
|
||||
self.storage.list_presets()
|
||||
self.assert_rejected(preset)
|
||||
|
||||
def test_save_and_service_both_validate_complete_task_schema(self):
|
||||
with self.assertRaises(RewardConfigError):
|
||||
self.save({}, base_configuration(OBSTACLE_TASK))
|
||||
self.assertEqual(self.storage.list_presets(), [])
|
||||
# The service itself must reject partial/malformed configs even from a broken resolver.
|
||||
self.training.preset_resolver = lambda _id, _task: {"weights": {"pose": 1}, "params": {}}
|
||||
self.assert_rejected({"id": "f" * 32})
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -10,6 +10,7 @@ from unittest.mock import patch
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
from server import ( # noqa: E402
|
||||
DEFAULT_TASKS,
|
||||
MAX_JOBS,
|
||||
ApiError,
|
||||
TrainingJob,
|
||||
@@ -79,6 +80,70 @@ out.write_bytes(b'onnx')
|
||||
with self.assertRaises(ApiError):
|
||||
self.manager.parse_config(self.payload(wandbMode="login"))
|
||||
|
||||
def test_custom_task_metadata_and_validated_deployment(self):
|
||||
self.manager.tasks = DEFAULT_TASKS
|
||||
health = self.manager.health()
|
||||
task = next(
|
||||
item for item in health["taskMetadata"] if item["id"] == "Unitree-Go2-ObstacleAvoidance"
|
||||
)
|
||||
self.assertEqual(task["sensorTypes"], ["raycast"])
|
||||
config = self.manager.parse_config(
|
||||
self.payload(
|
||||
taskId=task["id"],
|
||||
terrainPreset="discrete_obstacles",
|
||||
terrainParams={"obstacle_count": 8, "friction": 0.9},
|
||||
sensorCfg={"type": "raycast", "fov": 100, "maxDistance": 5},
|
||||
)
|
||||
)
|
||||
self.assertEqual(config.deployment["observationSize"], 81)
|
||||
self.assertEqual(config.deployment["terrain"]["actualObstacleCount"], 8)
|
||||
self.assertEqual(config.deployment["sensorCfg"]["rayCount"], 32)
|
||||
self.assertEqual(
|
||||
TrainingJob(id="a" * 32, config=config).public()["deployment"], config.deployment
|
||||
)
|
||||
rough = self.manager.parse_config(self.payload(taskId="Unitree-Go2-Rough"))
|
||||
self.assertFalse(rough.deployment["browserCompatible"])
|
||||
self.assertEqual(rough.deployment["observationSize"], 234)
|
||||
|
||||
def test_custom_config_rejects_unknown_nonfinite_and_incompatible_fields(self):
|
||||
self.manager.tasks = DEFAULT_TASKS
|
||||
for fields in (
|
||||
{"terrainPreset": "plane;touch /tmp/injected"},
|
||||
{"terrainPreset": ["plane"]},
|
||||
{"terrainParams": {"friction": float("nan")}},
|
||||
{"terrainParams": {"size": float("inf")}},
|
||||
{"terrainParams": {"size": 10**400}},
|
||||
{"terrainParams": {"obstacle_count": True}},
|
||||
{"terrainParams": {"obstacle_count": 1.5}},
|
||||
{"terrainParams": {"mjcf": "<include/>"}},
|
||||
{"terrainParams": {"obstacle_height_min": 1, "obstacle_height_max": 0.1}},
|
||||
{"sensorCfg": {"rayCount": 1_000_000}},
|
||||
{"sensorCfg": {"maxDistance": 1, "safetyDistance": 1}},
|
||||
{"sensorCfg": {"fov": float("nan")}},
|
||||
{"sensorType": "camera_depth"},
|
||||
{"rewardPresetId": "a" * 32},
|
||||
):
|
||||
with self.subTest(fields=fields), self.assertRaises(ApiError):
|
||||
self.manager.parse_config(
|
||||
self.payload(taskId="Unitree-Go2-ObstacleAvoidance", **fields)
|
||||
)
|
||||
with self.assertRaises(ApiError):
|
||||
self.manager.parse_config(self.payload(sensorCfg={"fov": 90}))
|
||||
|
||||
def test_custom_job_passes_server_owned_json_file_without_shell(self):
|
||||
self.manager.tasks = DEFAULT_TASKS
|
||||
config = self.manager.parse_config(self.payload(taskId="Unitree-Go2-ObstacleAvoidance"))
|
||||
job = TrainingJob(id="b" * 32, config=config)
|
||||
self.manager.jobs[job.id] = job
|
||||
with patch("server.subprocess.Popen", wraps=subprocess.Popen) as popen:
|
||||
self.manager._run(job)
|
||||
args = popen.call_args.args[0]
|
||||
self.assertNotIn("shell", popen.call_args.kwargs)
|
||||
path = Path(args[args.index("--task-config") + 1])
|
||||
self.assertTrue(path.is_relative_to(self.root))
|
||||
self.assertEqual(json.loads(path.read_text()), config.task_config)
|
||||
self.assertEqual(job.state, "succeeded")
|
||||
|
||||
def test_builds_argument_array_without_shell(self):
|
||||
config = self.manager.parse_config(self.payload(device="gpu", gpuIds=[0, 2]))
|
||||
command = self.manager.command_for(config)
|
||||
@@ -89,8 +154,13 @@ out.write_bytes(b'onnx')
|
||||
|
||||
def test_resolves_reward_preset_to_inline_validated_trainer_argument(self):
|
||||
preset_id = "f" * 32
|
||||
reward_config = {"weights": {"pose": 1.2}, "params": {}}
|
||||
self.manager.preset_resolver = lambda value: reward_config if value == preset_id else None
|
||||
from tuning.schema import base_configuration
|
||||
|
||||
reward_config = base_configuration()
|
||||
reward_config["weights"]["pose"] = 1.2
|
||||
self.manager.preset_resolver = lambda value, task: (
|
||||
reward_config if value == preset_id and task == "Unitree-Go2-Flat" else None
|
||||
)
|
||||
config = self.manager.parse_config(self.payload(rewardPresetId=preset_id))
|
||||
command = self.manager.command_for(config)
|
||||
index = command.index("--reward-config-json")
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
import json
|
||||
import math
|
||||
import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
|
||||
from task_config import ( # noqa: E402
|
||||
FLAT_TASK,
|
||||
OBSTACLE_TASK,
|
||||
TERRAIN_PRESETS,
|
||||
build_terrain_layout,
|
||||
deployment_metadata,
|
||||
navigation_candidates,
|
||||
validate_task_config,
|
||||
)
|
||||
|
||||
|
||||
class TaskConfigTest(unittest.TestCase):
|
||||
def test_flat_legacy_request_does_not_override_environment(self):
|
||||
self.assertIsNone(validate_task_config(FLAT_TASK, {}, 42))
|
||||
self.assertNotIn("terrain", deployment_metadata(FLAT_TASK, None, 42))
|
||||
|
||||
def test_layouts_are_bounded_finite_and_leave_flat_spawn_goal(self):
|
||||
for preset in (p for p in TERRAIN_PRESETS if p != "custom_boxes"):
|
||||
with self.subTest(preset=preset):
|
||||
custom = validate_task_config(OBSTACLE_TASK, {"terrainPreset": preset}, 42)
|
||||
layout = build_terrain_layout(custom)
|
||||
self.assertLessEqual(len(layout["boxes"]), 257)
|
||||
self.assertLess(len(json.dumps(layout)), 64 * 1024)
|
||||
self.assertEqual(layout["spawn"], [-5, 0, 0.32])
|
||||
self.assertEqual(layout["target"], [5, 0])
|
||||
for box in layout["boxes"]:
|
||||
self.assertTrue(all(math.isfinite(v) for v in box["pos"] + box["size"]))
|
||||
self.assertTrue(all(v > 0 for v in box["size"]))
|
||||
self.assertEqual(box["yaw"], 0)
|
||||
for box in layout["boxes"][1:]:
|
||||
self.assertGreaterEqual(box["pos"][0] - box["size"][0], -4)
|
||||
self.assertLessEqual(box["pos"][0] + box["size"][0], 4)
|
||||
|
||||
def test_layout_is_deterministic_and_export_is_authoritative(self):
|
||||
custom = validate_task_config(OBSTACLE_TASK, {}, 42)
|
||||
original = build_terrain_layout(custom)
|
||||
self.assertEqual(original, build_terrain_layout(custom))
|
||||
self.assertNotEqual(original, build_terrain_layout({**custom, "seed": 43}))
|
||||
metadata = deployment_metadata(OBSTACLE_TASK, custom, 42)
|
||||
self.assertEqual(metadata["terrain"], original)
|
||||
self.assertEqual(metadata["observationSize"], 47 + 32 + 2)
|
||||
self.assertEqual(metadata["jointNames"][0], "FL_hip_joint")
|
||||
self.assertEqual(metadata["navigation"]["distanceScale"], 12)
|
||||
self.assertTrue(metadata["sensorCfg"]["includeGround"])
|
||||
self.assertEqual(metadata["navigation"]["trainingReset"], "random-connected-free-pair")
|
||||
|
||||
def test_navigation_candidates_are_safe_connected_and_have_distant_fallbacks(self):
|
||||
layout = build_terrain_layout(validate_task_config(OBSTACLE_TASK, {}, 42))
|
||||
candidates = navigation_candidates(layout)
|
||||
self.assertGreater(len(candidates["points"]), 2)
|
||||
self.assertEqual(len(candidates["componentStarts"]), len(candidates["fallbackPairs"]))
|
||||
for point in candidates["points"]:
|
||||
self.assertTrue(all(abs(value) <= layout["size"] / 2 - 0.55 for value in point))
|
||||
for box in layout["boxes"][1:]:
|
||||
distance = math.hypot(
|
||||
max(abs(point[0] - box["pos"][0]) - box["size"][0], 0),
|
||||
max(abs(point[1] - box["pos"][1]) - box["size"][1], 0),
|
||||
)
|
||||
self.assertGreater(distance, candidates["clearance"])
|
||||
for first, second in candidates["fallbackPairs"]:
|
||||
self.assertGreaterEqual(math.dist(first, second), candidates["minDistance"])
|
||||
|
||||
def test_obstacle_count_is_capped_by_spacing_capacity(self):
|
||||
custom = validate_task_config(
|
||||
OBSTACLE_TASK,
|
||||
{
|
||||
"terrainParams": {"size": 8, "spacing": 3, "obstacle_count": 100},
|
||||
},
|
||||
0,
|
||||
)
|
||||
layout = build_terrain_layout(custom)
|
||||
self.assertEqual(layout["actualObstacleCount"], 2)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -12,7 +12,8 @@ from .schema import validate_proposal
|
||||
|
||||
SYSTEM_PROMPT = """你是 Unitree Go2 强化学习奖励调参专家。
|
||||
只根据提供的数值配置、训练曲线摘要和固定评估结果提出下一轮稀疏修改。
|
||||
必须优先保持速度跟踪与跌倒安全门槛;每轮最多修改四个白名单标量,不得改变符号、函数、传感器或结构。
|
||||
必须优先保持当前任务的客观安全门槛:Flat速度跟踪/跌倒,Obstacle无碰撞到达/跌倒。
|
||||
每轮最多修改四个白名单标量,不得改变符号、函数、传感器、评估协议或结构。
|
||||
不要建议 Python 代码、命令、文件路径或白名单外参数。输出必须符合 RewardProposal schema。
|
||||
"""
|
||||
|
||||
@@ -101,7 +102,9 @@ class DeepSeekAdvisor:
|
||||
result = self._agent().run_sync(prompt)
|
||||
output = result.output
|
||||
patch = validate_proposal(
|
||||
{"weights": dict(output.weights), "params": dict(output.params)}, previous
|
||||
{"weights": dict(output.weights), "params": dict(output.params)},
|
||||
previous,
|
||||
task_id=context.get("task", "Unitree-Go2-Flat"),
|
||||
)
|
||||
try:
|
||||
usage = result.usage()
|
||||
|
||||
@@ -13,11 +13,20 @@ from copy import deepcopy
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from pretrained_sources import PretrainedSources, SourceError
|
||||
from task_config import TaskConfigError, validate_task_config
|
||||
|
||||
from . import obstacle_scoring
|
||||
from .advisor import AdvisorUnavailable, DeepSeekAdvisor
|
||||
from .process import GpuLease, ResourceBusyError, terminate_process
|
||||
from .schema import (
|
||||
BASE_REWARD_CONFIGURATION,
|
||||
FLAT_TASK,
|
||||
OBSTACLE_TASK,
|
||||
RewardConfigError,
|
||||
base_configuration,
|
||||
merge_proposal,
|
||||
task_specs,
|
||||
validate_configuration,
|
||||
validate_configuration_constraints,
|
||||
validate_constraints,
|
||||
validate_proposal,
|
||||
@@ -51,6 +60,7 @@ class TuningManager:
|
||||
data_root: Path,
|
||||
lease: GpuLease,
|
||||
advisor: Any | None = None,
|
||||
sources: PretrainedSources | None = None,
|
||||
):
|
||||
self.trainer_root = trainer_root.expanduser().resolve()
|
||||
self.python = python
|
||||
@@ -66,10 +76,22 @@ class TuningManager:
|
||||
self.workers: dict[str, threading.Thread] = {}
|
||||
self.processes: dict[str, subprocess.Popen[str]] = {}
|
||||
self.cancel_events: dict[str, threading.Event] = {}
|
||||
self.sources = sources
|
||||
|
||||
def capability(self) -> dict[str, Any]:
|
||||
capability = self.advisor.capability()
|
||||
capability.update({"ready": (self.trainer_root / "scripts" / "evaluate.py").is_file()})
|
||||
capability.update(
|
||||
{
|
||||
"ready": (self.trainer_root / "scripts" / "evaluate.py").is_file(),
|
||||
"pretrainedSources": self.sources.catalog() if self.sources else [],
|
||||
"pretrainedUpload": {
|
||||
"enabled": self.sources is not None,
|
||||
"templateId": "go2-legacy47-v1",
|
||||
"formats": {"pt": 256 * 1024**2, "onnx": 64 * 1024**2},
|
||||
"endpoint": "/api/training/pretrained-sources/upload",
|
||||
},
|
||||
}
|
||||
)
|
||||
return capability
|
||||
|
||||
@staticmethod
|
||||
@@ -85,8 +107,34 @@ class TuningManager:
|
||||
mode = payload.get("mode", "automatic")
|
||||
if mode not in ("automatic", "approval"):
|
||||
raise TuningError("mode 必须是 automatic 或 approval")
|
||||
if payload.get("taskId", "Unitree-Go2-Flat") != "Unitree-Go2-Flat":
|
||||
raise TuningError("第一版只支持 Unitree-Go2-Flat")
|
||||
task_id = payload.get("taskId", "Unitree-Go2-Flat")
|
||||
task_specs(task_id)
|
||||
allowed = {
|
||||
"taskId",
|
||||
"mode",
|
||||
"runName",
|
||||
"gpuIds",
|
||||
"trialCount",
|
||||
"initialIterations",
|
||||
"middleIterations",
|
||||
"finalIterations",
|
||||
"numEnvs",
|
||||
"seed",
|
||||
"evalNumEnvs",
|
||||
"evalSteps",
|
||||
"earlyStopPatience",
|
||||
"objectiveWeights",
|
||||
"fallbackEnabled",
|
||||
"terrainPreset",
|
||||
"terrainParams",
|
||||
"sensorType",
|
||||
"sensorCfg",
|
||||
"customTerrainBoxes",
|
||||
"taskConfig",
|
||||
"pretrainedSourceId",
|
||||
}
|
||||
if payload.keys() - allowed:
|
||||
raise TuningError("未知session字段")
|
||||
run_name = payload.get("runName", "auto-tune")
|
||||
if not isinstance(run_name, str) or not RUN_NAME.fullmatch(run_name):
|
||||
raise TuningError("runName 格式无效")
|
||||
@@ -105,7 +153,7 @@ class TuningManager:
|
||||
rung1 = self._integer(payload, "middleIterations", 900, rung0, 1000000)
|
||||
rung2 = self._integer(payload, "finalIterations", 2000, rung1, 1000000)
|
||||
config = {
|
||||
"taskId": "Unitree-Go2-Flat",
|
||||
"taskId": task_id,
|
||||
"numEnvs": self._integer(payload, "numEnvs", 4096, 1, 16384),
|
||||
"seed": self._integer(payload, "seed", 42, 0, 2147483647),
|
||||
"runName": run_name,
|
||||
@@ -117,12 +165,78 @@ class TuningManager:
|
||||
"evalSteps": self._integer(payload, "evalSteps", 1000, 10, 100000),
|
||||
"earlyStopPatience": self._integer(payload, "earlyStopPatience", 4, 1, 20),
|
||||
}
|
||||
objective = validate_objective_weights(
|
||||
payload.get("objectiveWeights", DEFAULT_OBJECTIVE_WEIGHTS)
|
||||
)
|
||||
if task_id == OBSTACLE_TASK:
|
||||
if config["evalSteps"] != obstacle_scoring.STEPS:
|
||||
raise TuningError("避障评估固定1000步,不能调整")
|
||||
objective = deepcopy(obstacle_scoring.WEIGHTS)
|
||||
if "objectiveWeights" in payload and payload["objectiveWeights"] != objective:
|
||||
raise TuningError("避障评估权重固定")
|
||||
task_payload = payload.get("taskConfig", payload)
|
||||
if "taskConfig" in payload:
|
||||
if not isinstance(task_payload, dict) or task_payload.keys() - {
|
||||
"terrainPreset",
|
||||
"terrainParams",
|
||||
"sensorCfg",
|
||||
"sensorType",
|
||||
"seed",
|
||||
"customTerrainBoxes",
|
||||
}:
|
||||
raise TuningError("taskConfig字段无效")
|
||||
if any(
|
||||
k in payload
|
||||
for k in (
|
||||
"terrainPreset",
|
||||
"terrainParams",
|
||||
"sensorCfg",
|
||||
"sensorType",
|
||||
"customTerrainBoxes",
|
||||
)
|
||||
):
|
||||
raise TuningError("taskConfig不能与顶层场景配置混用")
|
||||
task_seed = task_payload.get("seed", config["seed"])
|
||||
if (
|
||||
isinstance(task_seed, bool)
|
||||
or not isinstance(task_seed, int)
|
||||
or task_seed != config["seed"]
|
||||
):
|
||||
raise TuningError("taskConfig seed必须与session一致")
|
||||
try:
|
||||
config["taskConfig"] = validate_task_config(task_id, task_payload, config["seed"])
|
||||
except TaskConfigError as error:
|
||||
raise TuningError(str(error)) from error
|
||||
base = base_configuration(task_id)
|
||||
base["weights"]["avoidance_weight"] = config["taskConfig"]["sensorCfg"][
|
||||
"avoidanceWeight"
|
||||
]
|
||||
validate_configuration_constraints(base, {}, task_id)
|
||||
else:
|
||||
if any(
|
||||
k in payload
|
||||
for k in (
|
||||
"terrainPreset",
|
||||
"terrainParams",
|
||||
"sensorCfg",
|
||||
"sensorType",
|
||||
"customTerrainBoxes",
|
||||
"taskConfig",
|
||||
)
|
||||
):
|
||||
raise TuningError("Flat调参不接受避障场景字段")
|
||||
objective = validate_objective_weights(
|
||||
payload.get("objectiveWeights", DEFAULT_OBJECTIVE_WEIGHTS)
|
||||
)
|
||||
fallback = payload.get("fallbackEnabled", False)
|
||||
if not isinstance(fallback, bool):
|
||||
raise TuningError("fallbackEnabled 必须是布尔值")
|
||||
if "pretrainedSourceId" in payload:
|
||||
if self.sources is None:
|
||||
raise TuningError("服务尚未注册基础策略,请配置--pretrained-sources")
|
||||
try:
|
||||
config["pretrained"] = self.sources.bind(
|
||||
payload["pretrainedSourceId"], task_id, config.get("taskConfig")
|
||||
)
|
||||
except SourceError as error:
|
||||
raise TuningError(str(error)) from error
|
||||
return mode, config, objective, fallback
|
||||
|
||||
def create(self, payload: Any) -> dict:
|
||||
@@ -262,8 +376,36 @@ class TuningManager:
|
||||
"--reward-config",
|
||||
str(reward_path),
|
||||
]
|
||||
task_path = None
|
||||
if config["taskId"] == OBSTACLE_TASK:
|
||||
task_path = run_dir / "task_config.json"
|
||||
task_path.write_text(
|
||||
json.dumps(config["taskConfig"], allow_nan=False), encoding="utf-8"
|
||||
)
|
||||
command.extend(("--task-config", str(task_path)))
|
||||
pretrained = config.get("pretrained")
|
||||
if pretrained is not None:
|
||||
if self.sources is None:
|
||||
raise TuningError("基础策略快照服务未配置,拒绝随机初始化")
|
||||
self.sources.verify(pretrained)
|
||||
if resume_checkpoint is not None:
|
||||
if pretrained is not None:
|
||||
origin_path = resume_checkpoint.parent / "initialization.json"
|
||||
if (
|
||||
not origin_path.is_file()
|
||||
or origin_path.is_symlink()
|
||||
or origin_path.stat().st_size > 64 * 1024
|
||||
):
|
||||
raise TuningError("续训checkpoint缺少有效基础策略来源记录,拒绝丢失来源")
|
||||
origin = json.loads(origin_path.read_text(encoding="utf-8"))
|
||||
if (
|
||||
origin.get("source_id") != pretrained["sourceId"]
|
||||
or origin.get("artifacts") != pretrained["manifest"]["artifacts"]
|
||||
):
|
||||
raise TuningError("续训checkpoint基础来源身份不匹配")
|
||||
command.extend(("--resume-checkpoint", str(resume_checkpoint)))
|
||||
elif pretrained is not None:
|
||||
command.extend(self.sources.arguments(pretrained))
|
||||
environment = os.environ.copy()
|
||||
environment["WANDB_MODE"] = "disabled"
|
||||
environment["WANDB_SILENT"] = "true"
|
||||
@@ -304,6 +446,8 @@ class TuningManager:
|
||||
"--gpu-ids",
|
||||
json.dumps(config["gpuIds"], separators=(",", ":")),
|
||||
]
|
||||
if task_path is not None:
|
||||
eval_command.extend(("--task-config", str(task_path)))
|
||||
self.storage.update_trial(trial_id, state="evaluating", message="正在固定协议评估")
|
||||
self.storage.update_session(
|
||||
session_id, state="evaluating", message=f"正在评估 trial {trial['number']}"
|
||||
@@ -316,7 +460,15 @@ class TuningManager:
|
||||
raise TuningError(f"评估失败(返回码 {return_code})")
|
||||
evaluation = json.loads(eval_output.read_text(encoding="utf-8"))
|
||||
baseline_trial = self.storage.list_trials(session_id)[0]
|
||||
if baseline_trial["evaluation"] is None:
|
||||
if config["taskId"] == OBSTACLE_TASK:
|
||||
expected = obstacle_scoring.protocol(config["taskConfig"], config["evalNumEnvs"])
|
||||
metrics = obstacle_scoring.validate_evaluation(evaluation, expected)
|
||||
baseline = baseline_trial["evaluation"]
|
||||
baseline_metrics = (
|
||||
obstacle_scoring.validate_evaluation(baseline, expected) if baseline else None
|
||||
)
|
||||
scored = obstacle_scoring.score_evaluation(metrics, baseline_metrics)
|
||||
elif baseline_trial["evaluation"] is None:
|
||||
scored = {
|
||||
"eligible": True,
|
||||
"score": 0.0,
|
||||
@@ -394,7 +546,28 @@ class TuningManager:
|
||||
return {
|
||||
"task": session["config"]["taskId"],
|
||||
"objectiveWeights": session["objectiveWeights"],
|
||||
"allowlist": "服务端将验证固定 schema;最多四项修改",
|
||||
"allowlist": {
|
||||
section: {
|
||||
key: {"min": spec.minimum, "max": spec.maximum} for key, spec in specs.items()
|
||||
}
|
||||
for section, specs in zip(
|
||||
("weights", "params"), task_specs(session["config"]["taskId"]), strict=True
|
||||
)
|
||||
},
|
||||
"taskContext": (
|
||||
(
|
||||
"97维目标导航:48条前视ray,3x16层pitch=[0,-20,-45]deg;仍有侧后/层间/坑盲区。"
|
||||
if session["config"]["taskConfig"]["sensorCfg"].get("sensorMode")
|
||||
== "multi_ring_raycast"
|
||||
else "81维目标导航:32条前视ray仅单高度切片,存在侧后方/矮障碍/跌落盲区。"
|
||||
)
|
||||
+ "避障权重对应近障平方惩罚,collision_penalty为非足端>10N接触惩罚;"
|
||||
"target_velocity是真实导航command速度而非reward权重。关注擦碰、绕行、目标到达和动作抖动。"
|
||||
"评估固定seed/1000步/客观权重,不允许修改地图、传感器、起终点、协议或以训练奖励代替指标。"
|
||||
)
|
||||
if session["config"]["taskId"] == OBSTACLE_TASK
|
||||
else "Flat速度跟踪与步态稳定;保留原六目标评估。",
|
||||
"taskConfig": session["config"].get("taskConfig"),
|
||||
"parameterConstraints": self.storage.get_control(session["id"])["constraints"],
|
||||
"rejectedFeedback": rejected_feedback,
|
||||
"trials": [
|
||||
@@ -410,12 +583,14 @@ class TuningManager:
|
||||
],
|
||||
}
|
||||
|
||||
def _fallback_patch(self, previous: dict, index: int) -> dict:
|
||||
def _fallback_patch(self, previous: dict, index: int, task_id="Unitree-Go2-Flat") -> dict:
|
||||
names = ("track_linear_velocity", "action_rate_l2", "body_orientation_l2", "foot_slip")
|
||||
if task_id == OBSTACLE_TASK:
|
||||
names = ("avoidance_weight", "collision_penalty")
|
||||
name = names[index % len(names)]
|
||||
old = previous["weights"][name]
|
||||
factor = 1.1 if index % 2 == 0 else 0.9
|
||||
return validate_proposal({"weights": {name: old * factor}}, previous)
|
||||
return validate_proposal({"weights": {name: old * factor}}, previous, task_id=task_id)
|
||||
|
||||
def _request_proposal(
|
||||
self, session: dict, previous: dict, base_trial_id: str, index: int
|
||||
@@ -427,7 +602,7 @@ class TuningManager:
|
||||
if not session["fallbackEnabled"]:
|
||||
raise AdvisorUnavailable(str(error)) from error
|
||||
result = {
|
||||
"patch": self._fallback_patch(previous, index),
|
||||
"patch": self._fallback_patch(previous, index, session["config"]["taskId"]),
|
||||
"rationale": f"Agent 不可用,显式 fallback:{error}",
|
||||
"expectedImpact": {},
|
||||
"confidence": 0.2,
|
||||
@@ -437,7 +612,9 @@ class TuningManager:
|
||||
}
|
||||
source = "fallback"
|
||||
constraints = self.storage.get_control(session["id"])["constraints"]
|
||||
result["patch"] = validate_proposal(result["patch"], previous, constraints)
|
||||
result["patch"] = validate_proposal(
|
||||
result["patch"], previous, constraints, session["config"]["taskId"]
|
||||
)
|
||||
proposal = self.storage.create_proposal(
|
||||
session["id"],
|
||||
base_trial_id,
|
||||
@@ -471,11 +648,34 @@ class TuningManager:
|
||||
self.condition.wait(timeout=1.0)
|
||||
raise TuningError("session 已取消")
|
||||
|
||||
@staticmethod
|
||||
def _base_configuration(session):
|
||||
base = base_configuration(session["config"]["taskId"])
|
||||
if session["config"]["taskId"] == OBSTACLE_TASK:
|
||||
base["weights"]["avoidance_weight"] = session["config"]["taskConfig"]["sensorCfg"][
|
||||
"avoidanceWeight"
|
||||
]
|
||||
return base
|
||||
|
||||
def _run_session(self, session_id: str, resume: bool, cancel: threading.Event) -> None:
|
||||
try:
|
||||
session = self.storage.get_session(session_id)
|
||||
trials = self.storage.list_trials(session_id)
|
||||
if resume:
|
||||
for proposal in self.storage.list_proposals(session_id):
|
||||
if proposal["state"] == "pending" and self.storage.decide_proposal(
|
||||
proposal["id"],
|
||||
"rejected",
|
||||
"service_restart/recovery_invalidated:服务重启,需重新提案并审批(非用户拒绝)",
|
||||
):
|
||||
self.storage.audit(
|
||||
session_id,
|
||||
"recovery_proposal_invalidated",
|
||||
{
|
||||
"proposalId": proposal["id"],
|
||||
"reason": "service_restart/recovery_invalidated",
|
||||
},
|
||||
)
|
||||
root = self._session_root(session_id)
|
||||
for interrupted in [trial for trial in trials if trial["state"] == "interrupted"]:
|
||||
run_dir = (root / interrupted["runDir"]).resolve()
|
||||
@@ -499,7 +699,7 @@ class TuningManager:
|
||||
0,
|
||||
0,
|
||||
session["config"]["rungs"][0],
|
||||
deepcopy(BASE_REWARD_CONFIGURATION),
|
||||
self._base_configuration(session),
|
||||
None,
|
||||
baseline_dir,
|
||||
)
|
||||
@@ -551,7 +751,12 @@ class TuningManager:
|
||||
proposal = self.storage.get_proposal(proposal["id"])
|
||||
self._claim_dispatch(session_id, cancel)
|
||||
constraints = self.storage.get_control(session_id)["constraints"]
|
||||
reward_config = merge_proposal(base["rewardConfig"], proposal["patch"], constraints)
|
||||
reward_config = merge_proposal(
|
||||
base["rewardConfig"],
|
||||
proposal["patch"],
|
||||
constraints,
|
||||
session["config"]["taskId"],
|
||||
)
|
||||
trial = self.storage.create_trial(
|
||||
session_id,
|
||||
next_number,
|
||||
@@ -698,7 +903,12 @@ class TuningManager:
|
||||
patch = payload["patch"]
|
||||
if feedback is not None and (not isinstance(feedback, str) or len(feedback) > 2000):
|
||||
raise TuningError("feedback 无效")
|
||||
patch = validate_proposal(patch, base["rewardConfig"], constraints)
|
||||
patch = validate_proposal(
|
||||
patch,
|
||||
base["rewardConfig"],
|
||||
constraints,
|
||||
self.storage.get_session(session_id)["config"]["taskId"],
|
||||
)
|
||||
if not self.storage.decide_proposal(proposal_id, "approved", feedback, patch):
|
||||
raise TuningError("proposal 已处理")
|
||||
self.storage.audit(
|
||||
@@ -743,7 +953,12 @@ class TuningManager:
|
||||
continue
|
||||
base = self.storage.get_trial(proposal["baseTrialId"])
|
||||
try:
|
||||
validate_proposal(proposal["patch"], base["rewardConfig"], constraints)
|
||||
validate_proposal(
|
||||
proposal["patch"],
|
||||
base["rewardConfig"],
|
||||
constraints,
|
||||
session["config"]["taskId"],
|
||||
)
|
||||
except Exception as error:
|
||||
self.storage.decide_proposal(
|
||||
proposal["id"], "rejected", f"参数护栏已变化:{error}"
|
||||
@@ -771,8 +986,8 @@ class TuningManager:
|
||||
revision = payload["revision"]
|
||||
if isinstance(revision, bool) or not isinstance(revision, int) or revision < 0:
|
||||
raise TuningError("constraints revision 必须是非负整数")
|
||||
constraints = validate_constraints(payload["constraints"])
|
||||
session = self.storage.get_session(session_id)
|
||||
constraints = validate_constraints(payload["constraints"], session["config"]["taskId"])
|
||||
if session["state"] not in ACTIVE_SESSION_STATES:
|
||||
raise TuningError("终态 session 不能修改参数护栏")
|
||||
control = self.storage.get_control(session_id)
|
||||
@@ -784,7 +999,9 @@ class TuningManager:
|
||||
else:
|
||||
base = self._best_highest_rung(session_id)
|
||||
if base is not None:
|
||||
validate_configuration_constraints(base["rewardConfig"], constraints)
|
||||
validate_configuration_constraints(
|
||||
base["rewardConfig"], constraints, session["config"]["taskId"]
|
||||
)
|
||||
try:
|
||||
updated = self.storage.replace_constraints(session_id, revision, constraints)
|
||||
except StorageConflict as error:
|
||||
@@ -796,7 +1013,12 @@ class TuningManager:
|
||||
continue
|
||||
proposal_base = self.storage.get_trial(proposal["baseTrialId"])
|
||||
try:
|
||||
validate_proposal(proposal["patch"], proposal_base["rewardConfig"], constraints)
|
||||
validate_proposal(
|
||||
proposal["patch"],
|
||||
proposal_base["rewardConfig"],
|
||||
constraints,
|
||||
session["config"]["taskId"],
|
||||
)
|
||||
except Exception as error:
|
||||
message = f"参数护栏 revision {updated['constraintsRevision']}:{error}"
|
||||
if self.storage.decide_proposal(proposal["id"], "rejected", message):
|
||||
@@ -914,7 +1136,16 @@ class TuningManager:
|
||||
return self.detail(session_id)
|
||||
|
||||
def resume(self, session_id: str) -> dict:
|
||||
# Serialize state transition and worker registration against duplicate requests.
|
||||
with self.condition:
|
||||
return self._resume_locked(session_id)
|
||||
|
||||
def _resume_locked(self, session_id: str) -> dict:
|
||||
session = self.storage.get_session(session_id)
|
||||
if session["config"].get("pretrained") is not None:
|
||||
if self.sources is None:
|
||||
raise TuningError("基础策略快照服务未配置,无法恢复")
|
||||
self.sources.verify(session["config"]["pretrained"])
|
||||
if session["state"] == "paused":
|
||||
self.storage.set_run_policy(session_id, "continuous")
|
||||
has_pending = any(
|
||||
@@ -982,8 +1213,11 @@ class TuningManager:
|
||||
raise TuningError("最佳策略文件不存在")
|
||||
return path
|
||||
|
||||
def preset_config(self, preset_id: str) -> dict:
|
||||
return self.storage.get_preset(preset_id)["rewardConfig"]
|
||||
def preset_config(self, preset_id: str, task_id: str = FLAT_TASK) -> dict:
|
||||
preset = self.storage.get_preset(preset_id)
|
||||
if preset["taskId"] != task_id:
|
||||
raise RewardConfigError("奖励 preset 来源任务与训练任务不匹配")
|
||||
return validate_configuration(preset["rewardConfig"], task_id)
|
||||
|
||||
def test_agent(self) -> dict:
|
||||
try:
|
||||
|
||||
@@ -0,0 +1,130 @@
|
||||
"""Fixed obstacle protocol; objective measurements never consume training rewards."""
|
||||
|
||||
import math
|
||||
from statistics import fmean
|
||||
|
||||
from task_config import build_terrain_layout, validate_task_config
|
||||
|
||||
from .schema import OBSTACLE_TASK
|
||||
from .scoring import EvaluationError
|
||||
|
||||
SEEDS = (101, 202, 303)
|
||||
STEPS = 1000
|
||||
WEIGHTS = {"success": 0.4, "time": 0.2, "clearance": 0.2, "smooth": 0.1, "no_fall": 0.1}
|
||||
METRICS = (*WEIGHTS, "collision_rate", "fall_rate", "arrival_rate", "ray_hit_rate")
|
||||
|
||||
|
||||
def evaluation_scenarios(custom):
|
||||
scenarios = []
|
||||
for seed in SEEDS:
|
||||
value = validate_task_config(OBSTACLE_TASK, custom, seed)
|
||||
scenarios.append(
|
||||
{"seed": seed, "taskConfig": value, "terrain": build_terrain_layout(value)}
|
||||
)
|
||||
return scenarios
|
||||
|
||||
|
||||
def protocol(custom, num_envs):
|
||||
return {
|
||||
"protocolVersion": "obstacle-v1",
|
||||
"seeds": list(SEEDS),
|
||||
"stepsPerSeed": STEPS,
|
||||
"numEnvs": num_envs,
|
||||
"objectiveWeights": WEIGHTS,
|
||||
"sceneMode": "fixed-custom-map"
|
||||
if custom["terrainPreset"] == "custom_boxes"
|
||||
else "three-seed-layouts",
|
||||
"scenarios": evaluation_scenarios(custom),
|
||||
"episodePolicy": "first-episode-only; terminal snapshot before auto-reset; fixed horizon",
|
||||
}
|
||||
|
||||
|
||||
def score_trajectory(samples, horizon=STEPS):
|
||||
"""One first-episode trajectory, one sample per policy step including first terminal."""
|
||||
if not samples or len(samples) > horizon:
|
||||
raise EvaluationError("轨迹样本数不完整")
|
||||
keys = {"distance", "clearance", "action_delta", "ray_hit", "collision", "fall", "terminal"}
|
||||
for sample in samples:
|
||||
if set(sample) != keys:
|
||||
raise EvaluationError("轨迹指标缺失/未知")
|
||||
for key in keys:
|
||||
value = sample[key]
|
||||
if not isinstance(value, (int, float)) or not math.isfinite(value) or value < 0:
|
||||
raise EvaluationError(f"无效轨迹指标 {key}")
|
||||
if key in {"ray_hit", "collision", "fall", "terminal"} and value > 1:
|
||||
raise EvaluationError(f"无效标志 {key}")
|
||||
if any(s["terminal"] for s in samples[:-1]) or (
|
||||
len(samples) != horizon and not samples[-1]["terminal"]
|
||||
):
|
||||
raise EvaluationError("首episode轨迹不完整")
|
||||
collision = any(s["collision"] for s in samples)
|
||||
fall = any(s["fall"] for s in samples)
|
||||
arrival = next((i + 1 for i, s in enumerate(samples) if s["distance"] < 0.5), None)
|
||||
success = arrival is not None and not collision and not fall
|
||||
# Missing steps after an early terminal earn zero clearance/smoothness, not a bonus.
|
||||
return {
|
||||
"success": float(success),
|
||||
"time": 1 - arrival / horizon if success else 0.0,
|
||||
"clearance": 0.0 if fall else sum(min(s["clearance"] / 0.5, 1) for s in samples) / horizon,
|
||||
"smooth": sum(1 - min(s["action_delta"], 1) for s in samples) / horizon,
|
||||
"no_fall": float(not fall),
|
||||
"collision_rate": float(collision),
|
||||
"fall_rate": float(fall),
|
||||
"arrival_rate": float(arrival is not None),
|
||||
"ray_hit_rate": fmean(s["ray_hit"] for s in samples),
|
||||
}
|
||||
|
||||
|
||||
def validate_metrics(value):
|
||||
if not isinstance(value, dict) or set(value) != set(METRICS):
|
||||
raise EvaluationError("避障评估指标缺失/未知")
|
||||
for key, number in value.items():
|
||||
if (
|
||||
isinstance(number, bool)
|
||||
or not isinstance(number, (int, float))
|
||||
or not math.isfinite(number)
|
||||
or not 0 <= number <= 1
|
||||
):
|
||||
raise EvaluationError(f"避障指标 {key} 必须在0–1且有限")
|
||||
if not math.isclose(value["no_fall"] + value["fall_rate"], 1, abs_tol=1e-8):
|
||||
raise EvaluationError("跌倒指标不一致")
|
||||
return value
|
||||
|
||||
|
||||
def validate_evaluation(value, expected_protocol):
|
||||
if value.get("protocol") != expected_protocol:
|
||||
raise EvaluationError("避障评估协议/场景不一致")
|
||||
seeds = value.get("seedMetrics", [])
|
||||
if len(seeds) != 3:
|
||||
raise EvaluationError("避障评估必须完整包含3seed")
|
||||
for item, seed in zip(seeds, SEEDS, strict=True):
|
||||
if (
|
||||
item.get("seed") != seed
|
||||
or item.get("episodes") != expected_protocol["numEnvs"]
|
||||
or item.get("rolloutSteps") != STEPS
|
||||
):
|
||||
raise EvaluationError("避障评估样本数不完整")
|
||||
validate_metrics(item.get("metrics"))
|
||||
metrics = validate_metrics(value.get("metrics"))
|
||||
if any(
|
||||
not math.isclose(metrics[k], fmean(s["metrics"][k] for s in seeds), abs_tol=1e-8)
|
||||
for k in METRICS
|
||||
):
|
||||
raise EvaluationError("避障评估聚合不一致")
|
||||
return metrics
|
||||
|
||||
|
||||
def score_evaluation(current, baseline=None):
|
||||
current = validate_metrics(current)
|
||||
baseline = validate_metrics(baseline) if baseline is not None else current
|
||||
eligible = (
|
||||
current["fall_rate"] <= baseline["fall_rate"] + 0.02 + 1e-12
|
||||
and current["success"] >= baseline["success"] - 0.02 - 1e-12
|
||||
)
|
||||
total = sum(WEIGHTS[k] * current[k] for k in WEIGHTS)
|
||||
return {
|
||||
"eligible": eligible,
|
||||
"score": total if eligible else -1.0,
|
||||
"components": {k: current[k] for k in WEIGHTS},
|
||||
"metrics": current,
|
||||
}
|
||||
@@ -68,6 +68,32 @@ BASE_REWARD_CONFIGURATION: dict[str, dict[str, float]] = {
|
||||
}
|
||||
|
||||
|
||||
FLAT_TASK = "Unitree-Go2-Flat"
|
||||
OBSTACLE_TASK = "Unitree-Go2-ObstacleAvoidance"
|
||||
OBSTACLE_WEIGHT_SPECS = {
|
||||
"avoidance_weight": NumericSpec(0.5, 5.0, 2.0, False),
|
||||
"collision_penalty": NumericSpec(-10.0, -0.5, -5.0, False),
|
||||
"action_smoothness": NumericSpec(-0.05, -0.001, -0.05, False),
|
||||
}
|
||||
OBSTACLE_PARAMETER_SPECS = {"target_velocity": NumericSpec(0.3, 1.2, 0.6, False)}
|
||||
|
||||
|
||||
def task_specs(task_id=FLAT_TASK):
|
||||
if task_id == FLAT_TASK:
|
||||
return WEIGHT_SPECS, PARAMETER_SPECS
|
||||
if task_id == OBSTACLE_TASK:
|
||||
return OBSTACLE_WEIGHT_SPECS, OBSTACLE_PARAMETER_SPECS
|
||||
raise RewardConfigError(f"只支持Flat或Obstacle调参任务:{task_id}")
|
||||
|
||||
|
||||
def base_configuration(task_id=FLAT_TASK):
|
||||
weights, params = task_specs(task_id)
|
||||
return {
|
||||
"weights": {k: v.default for k, v in weights.items()},
|
||||
"params": {k: v.default for k, v in params.items()},
|
||||
}
|
||||
|
||||
|
||||
def _number(name: str, value: Any, spec: NumericSpec) -> float:
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||
raise RewardConfigError(f"{name} 必须是数值")
|
||||
@@ -87,19 +113,22 @@ def _mapping(value: Any, name: str) -> Mapping[str, Any]:
|
||||
return value
|
||||
|
||||
|
||||
def _cross_validate(config: Mapping[str, Mapping[str, float]]) -> None:
|
||||
def _cross_validate(config: Mapping[str, Mapping[str, float]], task_id=FLAT_TASK) -> None:
|
||||
if task_id == OBSTACLE_TASK:
|
||||
return
|
||||
params = config["params"]
|
||||
if params["pose.walking_threshold"] >= params["pose.running_threshold"]:
|
||||
raise RewardConfigError("pose.walking_threshold 必须小于 pose.running_threshold")
|
||||
|
||||
|
||||
def _path_spec(path: str) -> tuple[str, str, NumericSpec]:
|
||||
def _path_spec(path: str, task_id=FLAT_TASK) -> tuple[str, str, NumericSpec]:
|
||||
weights, params = task_specs(task_id)
|
||||
if path.startswith("weights."):
|
||||
section, name = "weights", path.removeprefix("weights.")
|
||||
spec = WEIGHT_SPECS.get(name)
|
||||
spec = weights.get(name)
|
||||
elif path.startswith("params."):
|
||||
section, name = "params", path.removeprefix("params.")
|
||||
spec = PARAMETER_SPECS.get(name)
|
||||
spec = params.get(name)
|
||||
else:
|
||||
section, name, spec = "", "", None
|
||||
if spec is None:
|
||||
@@ -107,16 +136,16 @@ def _path_spec(path: str) -> tuple[str, str, NumericSpec]:
|
||||
return section, name, spec
|
||||
|
||||
|
||||
def validate_constraints(value: Any) -> dict[str, dict[str, float | str]]:
|
||||
def validate_constraints(value: Any, task_id=FLAT_TASK) -> dict[str, dict[str, float | str]]:
|
||||
"""Validate sparse per-session range/fixed safety constraints."""
|
||||
root = _mapping(value, "constraints")
|
||||
if len(root) > len(WEIGHT_SPECS) + len(PARAMETER_SPECS):
|
||||
if len(root) > sum(len(specs) for specs in task_specs(task_id)):
|
||||
raise RewardConfigError("constraints 数量超过白名单参数总数")
|
||||
result: dict[str, dict[str, float | str]] = {}
|
||||
for raw_path, raw_constraint in root.items():
|
||||
if not isinstance(raw_path, str):
|
||||
raise RewardConfigError("constraint path 必须是字符串")
|
||||
_, _, spec = _path_spec(raw_path)
|
||||
_, _, spec = _path_spec(raw_path, task_id)
|
||||
constraint = _mapping(raw_constraint, raw_path)
|
||||
kind = constraint.get("kind")
|
||||
if kind == "fixed":
|
||||
@@ -137,12 +166,12 @@ def validate_constraints(value: Any) -> dict[str, dict[str, float | str]]:
|
||||
return result
|
||||
|
||||
|
||||
def validate_configuration_constraints(value: Any, constraints: Any) -> None:
|
||||
def validate_configuration_constraints(value: Any, constraints: Any, task_id=FLAT_TASK) -> None:
|
||||
"""Ensure a complete reward configuration satisfies every session constraint."""
|
||||
config = validate_configuration(value)
|
||||
checked = validate_constraints(constraints)
|
||||
config = validate_configuration(value, task_id)
|
||||
checked = validate_constraints(constraints, task_id)
|
||||
for path, constraint in checked.items():
|
||||
section, name, _ = _path_spec(path)
|
||||
section, name, _ = _path_spec(path, task_id)
|
||||
current = config[section][name]
|
||||
if constraint["kind"] == "fixed":
|
||||
if current != constraint["value"]:
|
||||
@@ -155,36 +184,37 @@ def validate_configuration_constraints(value: Any, constraints: Any) -> None:
|
||||
)
|
||||
|
||||
|
||||
def validate_configuration(value: Any) -> dict[str, dict[str, float]]:
|
||||
def validate_configuration(value: Any, task_id=FLAT_TASK) -> dict[str, dict[str, float]]:
|
||||
"""Validate a complete configuration and reject missing/unknown fields."""
|
||||
root = _mapping(value, "rewardConfig")
|
||||
if set(root) != {"weights", "params"}:
|
||||
raise RewardConfigError("rewardConfig 只能包含 weights 和 params")
|
||||
weights, params = task_specs(task_id)
|
||||
raw_weights = _mapping(root["weights"], "weights")
|
||||
raw_params = _mapping(root["params"], "params")
|
||||
if set(raw_weights) != set(WEIGHT_SPECS):
|
||||
if set(raw_weights) != set(weights):
|
||||
raise RewardConfigError("weights 必须完整且不能包含未知奖励项")
|
||||
if set(raw_params) != set(PARAMETER_SPECS):
|
||||
if set(raw_params) != set(params):
|
||||
raise RewardConfigError("params 必须完整且不能包含未知参数")
|
||||
config = {
|
||||
"weights": {
|
||||
name: _number(f"weights.{name}", raw_weights[name], spec)
|
||||
for name, spec in WEIGHT_SPECS.items()
|
||||
for name, spec in weights.items()
|
||||
},
|
||||
"params": {
|
||||
name: _number(f"params.{name}", raw_params[name], spec)
|
||||
for name, spec in PARAMETER_SPECS.items()
|
||||
name: _number(f"params.{name}", raw_params[name], spec) for name, spec in params.items()
|
||||
},
|
||||
}
|
||||
_cross_validate(config)
|
||||
_cross_validate(config, task_id)
|
||||
return config
|
||||
|
||||
|
||||
def validate_proposal(
|
||||
value: Any, previous: Any, constraints: Any | None = None
|
||||
value: Any, previous: Any, constraints: Any | None = None, task_id=FLAT_TASK
|
||||
) -> dict[str, dict[str, float]]:
|
||||
"""Validate a sparse Agent patch relative to a complete previous config and guardrails."""
|
||||
current = validate_configuration(previous)
|
||||
weights, params = task_specs(task_id)
|
||||
current = validate_configuration(previous, task_id)
|
||||
root = _mapping(value, "proposal")
|
||||
if not set(root).issubset({"weights", "params"}):
|
||||
raise RewardConfigError("proposal 只能包含 weights 和 params")
|
||||
@@ -194,8 +224,8 @@ def validate_proposal(
|
||||
raise RewardConfigError("proposal 至少需要一项修改")
|
||||
if len(raw_weights) + len(raw_params) > MAX_PROPOSAL_CHANGES:
|
||||
raise RewardConfigError(f"proposal 每轮最多修改 {MAX_PROPOSAL_CHANGES} 项")
|
||||
unknown_weights = set(raw_weights) - set(WEIGHT_SPECS)
|
||||
unknown_params = set(raw_params) - set(PARAMETER_SPECS)
|
||||
unknown_weights = set(raw_weights) - set(weights)
|
||||
unknown_params = set(raw_params) - set(params)
|
||||
if unknown_weights:
|
||||
raise RewardConfigError(f"未知奖励项:{', '.join(sorted(unknown_weights))}")
|
||||
if unknown_params:
|
||||
@@ -203,7 +233,7 @@ def validate_proposal(
|
||||
|
||||
patch: dict[str, dict[str, float]] = {"weights": {}, "params": {}}
|
||||
for name, raw in raw_weights.items():
|
||||
value_number = _number(f"weights.{name}", raw, WEIGHT_SPECS[name])
|
||||
value_number = _number(f"weights.{name}", raw, weights[name])
|
||||
old = current["weights"][name]
|
||||
if old != 0.0 and value_number != 0.0:
|
||||
ratio = abs(value_number / old)
|
||||
@@ -216,7 +246,7 @@ def validate_proposal(
|
||||
raise RewardConfigError(f"weights.{name} 没有发生变化")
|
||||
patch["weights"][name] = value_number
|
||||
for name, raw in raw_params.items():
|
||||
value_number = _number(f"params.{name}", raw, PARAMETER_SPECS[name])
|
||||
value_number = _number(f"params.{name}", raw, params[name])
|
||||
old = current["params"][name]
|
||||
ratio = abs(value_number / old)
|
||||
if ratio < MIN_CHANGE_RATIO or ratio > MAX_CHANGE_RATIO:
|
||||
@@ -230,26 +260,39 @@ def validate_proposal(
|
||||
candidate = deepcopy(current)
|
||||
candidate["weights"].update(patch["weights"])
|
||||
candidate["params"].update(patch["params"])
|
||||
_cross_validate(candidate)
|
||||
_cross_validate(candidate, task_id)
|
||||
if constraints is not None:
|
||||
validate_configuration_constraints(candidate, constraints)
|
||||
validate_configuration_constraints(candidate, constraints, task_id)
|
||||
return patch
|
||||
|
||||
|
||||
def merge_proposal(
|
||||
previous: Any, proposal: Any, constraints: Any | None = None
|
||||
previous: Any, proposal: Any, constraints: Any | None = None, task_id=FLAT_TASK
|
||||
) -> dict[str, dict[str, float]]:
|
||||
current = validate_configuration(previous)
|
||||
patch = validate_proposal(proposal, current, constraints)
|
||||
current = validate_configuration(previous, task_id)
|
||||
patch = validate_proposal(proposal, current, constraints, task_id)
|
||||
merged = deepcopy(current)
|
||||
merged["weights"].update(patch["weights"])
|
||||
merged["params"].update(patch["params"])
|
||||
return validate_configuration(merged)
|
||||
return validate_configuration(merged, task_id)
|
||||
|
||||
|
||||
def apply_reward_configuration(env_cfg: Any, value: Any) -> None:
|
||||
def apply_reward_configuration(env_cfg: Any, value: Any, task_id=FLAT_TASK) -> None:
|
||||
"""Apply a validated full config to a fresh mjlab environment config."""
|
||||
config = validate_configuration(value)
|
||||
config = validate_configuration(value, task_id)
|
||||
if task_id == OBSTACLE_TASK:
|
||||
from mjlab.managers import RewardTermCfg
|
||||
|
||||
contact = env_cfg.terminations["illegal_contact"]
|
||||
env_cfg.rewards["obstacle_collision"] = RewardTermCfg(
|
||||
func=contact.func,
|
||||
params=deepcopy(contact.params),
|
||||
weight=config["weights"]["collision_penalty"],
|
||||
)
|
||||
env_cfg.rewards["obstacle_proximity"].weight = -config["weights"]["avoidance_weight"]
|
||||
env_cfg.rewards["action_rate_l2"].weight = config["weights"]["action_smoothness"]
|
||||
env_cfg.commands["twist"].speed = config["params"]["target_velocity"]
|
||||
return
|
||||
for name, weight in config["weights"].items():
|
||||
if name not in env_cfg.rewards:
|
||||
raise RewardConfigError(f"环境缺少奖励项:{name}")
|
||||
|
||||
@@ -12,6 +12,8 @@ from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from .schema import FLAT_TASK, RewardConfigError, validate_configuration
|
||||
|
||||
SCHEMA_VERSION = 2
|
||||
|
||||
|
||||
@@ -169,16 +171,37 @@ class TuningStorage:
|
||||
def recover_interrupted(self) -> None:
|
||||
at = now_iso()
|
||||
with self.transaction() as connection:
|
||||
previous = connection.execute(
|
||||
"SELECT id,state FROM sessions WHERE state IN "
|
||||
"('queued','running','evaluating','paused','awaiting_approval')"
|
||||
).fetchall()
|
||||
for row in previous:
|
||||
connection.execute(
|
||||
"INSERT INTO audit_events(session_id,event_type,payload_json,created_at) "
|
||||
"VALUES (?,?,?,?)",
|
||||
(
|
||||
row["id"],
|
||||
"service_restart_interrupted",
|
||||
_json(
|
||||
{
|
||||
"previousState": row["state"],
|
||||
"reason": "service_restart",
|
||||
"requiresExplicitResume": True,
|
||||
}
|
||||
),
|
||||
at,
|
||||
),
|
||||
)
|
||||
connection.execute(
|
||||
"UPDATE trials SET state='interrupted', ended_at=?, "
|
||||
"message='服务重启中断,等待显式恢复' "
|
||||
"WHERE state IN ('training','evaluating')",
|
||||
"WHERE state IN ('queued','training','evaluating')",
|
||||
(at,),
|
||||
)
|
||||
connection.execute(
|
||||
"UPDATE sessions SET state='interrupted', updated_at=?, "
|
||||
"message='服务重启中断,可从完整 checkpoint 恢复' "
|
||||
"WHERE state IN ('running','evaluating')",
|
||||
"WHERE state IN ('queued','running','evaluating','paused','awaiting_approval')",
|
||||
(at,),
|
||||
)
|
||||
|
||||
@@ -643,47 +666,51 @@ class TuningStorage:
|
||||
for row in rows
|
||||
]
|
||||
|
||||
def _preset_task(self, session_id: str) -> str:
|
||||
# Identity comes only from the persisted source session, never client/preset labels.
|
||||
try:
|
||||
config = self.get_session(session_id)["config"]
|
||||
except (KeyError, ValueError, TypeError) as error:
|
||||
raise RewardConfigError("奖励 preset 来源 session 丢失或损坏") from error
|
||||
if not isinstance(config, dict):
|
||||
raise RewardConfigError("奖励 preset 来源 config 必须是对象")
|
||||
# Original Flat-only sessions did not require taskId. Validate their preset below.
|
||||
return config.get("taskId", FLAT_TASK)
|
||||
|
||||
def _preset(self, row) -> dict:
|
||||
task_id = self._preset_task(row["session_id"])
|
||||
try:
|
||||
reward_config = validate_configuration(_decode(row["reward_config_json"]), task_id)
|
||||
except (ValueError, TypeError) as error:
|
||||
raise RewardConfigError("奖励 preset 配置无效:" + str(error)) from error
|
||||
return {
|
||||
"id": row["id"],
|
||||
"name": row["name"],
|
||||
"sessionId": row["session_id"],
|
||||
"trialId": row["trial_id"],
|
||||
"taskId": task_id,
|
||||
"rewardConfig": reward_config,
|
||||
"createdAt": row["created_at"],
|
||||
}
|
||||
|
||||
def save_preset(self, name: str, session_id: str, trial_id: str, reward_config: dict) -> dict:
|
||||
reward_config = validate_configuration(reward_config, self._preset_task(session_id))
|
||||
preset_id, at = uuid.uuid4().hex, now_iso()
|
||||
self.connection().execute(
|
||||
"INSERT INTO presets(id,name,session_id,trial_id,reward_config_json,created_at) "
|
||||
"VALUES (?,?,?,?,?,?)",
|
||||
(preset_id, name, session_id, trial_id, _json(reward_config), at),
|
||||
)
|
||||
return {
|
||||
"id": preset_id,
|
||||
"name": name,
|
||||
"sessionId": session_id,
|
||||
"trialId": trial_id,
|
||||
"rewardConfig": reward_config,
|
||||
"createdAt": at,
|
||||
}
|
||||
return self.get_preset(preset_id)
|
||||
|
||||
def get_preset(self, preset_id: str) -> dict:
|
||||
row = self.connection().execute("SELECT * FROM presets WHERE id=?", (preset_id,)).fetchone()
|
||||
if row is None:
|
||||
raise KeyError(preset_id)
|
||||
return {
|
||||
"id": row["id"],
|
||||
"name": row["name"],
|
||||
"sessionId": row["session_id"],
|
||||
"trialId": row["trial_id"],
|
||||
"rewardConfig": _decode(row["reward_config_json"]),
|
||||
"createdAt": row["created_at"],
|
||||
}
|
||||
return self._preset(row)
|
||||
|
||||
def list_presets(self) -> list[dict]:
|
||||
rows = (
|
||||
self.connection().execute("SELECT * FROM presets ORDER BY created_at DESC").fetchall()
|
||||
)
|
||||
return [
|
||||
{
|
||||
"id": row["id"],
|
||||
"name": row["name"],
|
||||
"sessionId": row["session_id"],
|
||||
"trialId": row["trial_id"],
|
||||
"rewardConfig": _decode(row["reward_config_json"]),
|
||||
"createdAt": row["created_at"],
|
||||
}
|
||||
for row in rows
|
||||
]
|
||||
return [self._preset(row) for row in rows]
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
# 自定义训练与前视射线避障
|
||||
|
||||
1. 启动仓库本地训练服务,在「控制台 → 强化学习任务」输入令牌并连接。
|
||||
2. 选择「前视射线避障导航」,设置地形、种子、障碍物参数、FOV、探测/安全距离及避障奖励权重,发起训练。界面显示进度、最近价值/策略/熵损失和原始日志。
|
||||
3. 「同步当前场景地图」从全部**已应用**实例的已编译MuJoCo静态碰撞几何导出权威 `custom_boxes/boxes-v1`,支持多实例、平移/旋转、工程地图路径及嵌套body。成功显示“已将视口中 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不生效,重新按下的正常点击仍可设定。
|
||||
@@ -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);
|
||||
});
|
||||
@@ -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);
|
||||
});
|
||||
@@ -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);
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,9 @@
|
||||
# 测试专用避障契约模型
|
||||
|
||||
`zero-action.onnx`:ONNX opset17/IR9,输入float32[1,81],Constant节点输出float32[1,12]全0,携带seed42避障deployment。仅验浏览器地图/感知/推理闭环,**不是已训练或可导航策略**。由onnx.helper.make_model/make_graph/make_node(Constant)生成,metadata来自src/rl/fixtures/obstacleDeployment.json。
|
||||
|
||||
事务回归fixture:`wrong-graph.onnx`包含合法避障metadata但实际输入47维;`ort-init-failure.onnx`包含合法81维metadata但不存在的算子(故意无法初始化ORT);`legacy-flat.onnx`为无metadata的47→12零动作旧导出;`legacy-wrong-shape.onnx`为无metadata的81→12零动作。均由相同onnx.helper入口生成,仅供成功/失败分支测试。
|
||||
|
||||
`multi-zero-action.onnx`同样为Constant零动作,仅将实际输入shape设为[1,97]并携带`multiRingDeployment.json`。不是训练策略,也没有将81网络冒充97。真实短训导出只保留在/tmp。
|
||||
|
||||
`multi-wrong-graph.onnx`携带合法97维metadata但真实输入81维,用于浏览器ORT明确shape拒绝及旧会话保留回归。
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+155
-52
@@ -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(/安全区/);
|
||||
});
|
||||
});
|
||||
@@ -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);
|
||||
@@ -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();
|
||||
});
|
||||
@@ -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"
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
});
|
||||
@@ -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,
|
||||
};
|
||||
});
|
||||
}
|
||||
@@ -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();
|
||||
}
|
||||
});
|
||||
});
|
||||
@@ -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();
|
||||
});
|
||||
|
||||
@@ -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 }));
|
||||
}
|
||||
@@ -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
Reference in New Issue
Block a user