From deead17a9ad2136416a8481e3faec00b1155906c Mon Sep 17 00:00:00 2001 From: cen617-code <1057290604@qq.com> Date: Wed, 2 Sep 2026 13:49:34 +0800 Subject: [PATCH 1/4] =?UTF-8?q?feat(training):=20release=20V0.8=20?= =?UTF-8?q?=E8=87=AA=E8=B0=83=E5=8F=82=20Agent?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .gitignore | 1 + AGENTS.md | 3 + README.md | 4 +- context.md | 116 +++ package-lock.json | 7 + package.json | 1 + plans/map-v3-design.md | 193 ++++ plans/reward-auto-tuning-agent.md | 104 +++ plans/urdf-studio-ui-components.md | 178 ++++ requirements-dev.txt | 1 + training_server/README.md | 40 +- training_server/requirements.txt | 5 + training_server/rl/scripts/evaluate.py | 222 +++++ training_server/rl/scripts/train.py | 83 +- .../src/tasks/velocity/config/go2/env_cfgs.py | 1 + .../rl/src/tasks/velocity/mdp/__init__.py | 1 + .../rl/src/tasks/velocity/mdp/metrics.py | 88 ++ .../rl/src/tasks/velocity/velocity_env_cfg.py | 28 +- training_server/server.py | 188 +++- training_server/tests/test_server.py | 18 +- training_server/tests/test_tuning.py | 195 ++++ training_server/tests/test_tuning_manager.py | 195 ++++ training_server/tuning/__init__.py | 26 + training_server/tuning/advisor.py | 134 +++ training_server/tuning/manager.py | 715 +++++++++++++++ training_server/tuning/process.py | 46 + training_server/tuning/schema.py | 213 +++++ training_server/tuning/scoring.py | 118 +++ training_server/tuning/storage.py | 521 +++++++++++ training_server/tuning/study.py | 55 ++ training_server/tuning/tensorboard.py | 35 + web_platform/ARCHITECTURE.md | 5 +- web_platform/README.md | 6 +- web_platform/e2e/app.spec.ts | 8 + web_platform/src/test/setup.ts | 21 + .../src/training/LocalTrainingClient.test.ts | 75 ++ .../src/training/LocalTrainingClient.ts | 92 +- .../src/training/LocalTrainingPanel.test.tsx | 23 +- .../src/training/LocalTrainingPanel.tsx | 119 ++- web_platform/src/training/storage.ts | 30 + web_platform/src/training/types.ts | 142 +++ web_platform/src/tuning/ScalarChart.tsx | 81 ++ web_platform/src/tuning/TuningApp.test.tsx | 53 ++ web_platform/src/tuning/TuningApp.tsx | 846 ++++++++++++++++++ web_platform/src/tuning/main.tsx | 14 + web_platform/tuning.html | 26 + web_platform/vite.config.ts | 6 + 47 files changed, 4986 insertions(+), 96 deletions(-) create mode 100644 AGENTS.md create mode 100644 context.md create mode 100644 plans/map-v3-design.md create mode 100644 plans/reward-auto-tuning-agent.md create mode 100644 plans/urdf-studio-ui-components.md create mode 100644 training_server/requirements.txt create mode 100644 training_server/rl/scripts/evaluate.py create mode 100644 training_server/rl/src/tasks/velocity/mdp/metrics.py create mode 100644 training_server/tests/test_tuning.py create mode 100644 training_server/tests/test_tuning_manager.py create mode 100644 training_server/tuning/__init__.py create mode 100644 training_server/tuning/advisor.py create mode 100644 training_server/tuning/manager.py create mode 100644 training_server/tuning/process.py create mode 100644 training_server/tuning/schema.py create mode 100644 training_server/tuning/scoring.py create mode 100644 training_server/tuning/storage.py create mode 100644 training_server/tuning/study.py create mode 100644 training_server/tuning/tensorboard.py create mode 100644 web_platform/src/training/storage.ts create mode 100644 web_platform/src/tuning/ScalarChart.tsx create mode 100644 web_platform/src/tuning/TuningApp.test.tsx create mode 100644 web_platform/src/tuning/TuningApp.tsx create mode 100644 web_platform/src/tuning/main.tsx create mode 100644 web_platform/tuning.html diff --git a/.gitignore b/.gitignore index e264caff..4ce07305 100644 --- a/.gitignore +++ b/.gitignore @@ -31,3 +31,4 @@ MUJOCO_LOG.TXT training_server/rl/logs/ training_server/rl/wandb/ training_server/rl/outputs/ +training_server/tuning-data/ diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 00000000..52914f36 --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,3 @@ +Please always speak chinese. +python 虚拟环境路径在/home/cen/Embodied_Workspace/Mujoco_Projects/mujoco/.venv/bin/activate +系统是Ubuntu 24.04 LTS \ No newline at end of file diff --git a/README.md b/README.md index 94c1aecb..18b2010f 100644 --- a/README.md +++ b/README.md @@ -15,7 +15,7 @@ - V3 地图创作层:认证资产库支持点击添加、拖到画布落位和首个资产自动创建场景,并可通过表单与视口操纵器继续编辑、事务式应用及导出地图 ZIP - 浏览器内 Python 控制器(Pyodide) - ONNX 强化学习策略推理(ONNX Runtime Web) -- 内置 Go2 PPO 任务的本机 mjlab 训练桥接服务 +- 内置 Go2 PPO 任务的本机 mjlab 训练桥接服务,以及 DeepSeek 驱动的奖励函数自调参、固定评估和 TensorBoard 风格独立工作台 - 响应式工作区、源码编辑、性能监控和中文诊断 ## 快速开始 @@ -82,7 +82,7 @@ npm run lint:python ## 数据与安全边界 -模型、资源、Python 控制器和 ONNX 策略默认只在当前浏览器会话中处理,不上传到服务器。训练桥接服务只监听本机回环地址,并仅执行服务端允许列表中的任务。 +模型、资源、Python 控制器和 ONNX 策略默认只在当前浏览器会话中处理,不上传到服务器。训练桥接服务只监听本机回环地址,并仅执行服务端允许列表中的任务。启用自调参时,DeepSeek 只接收脱敏后的奖励参数、曲线摘要和评估数值;API key 仅存在训练服务环境中,不进入浏览器、URL、SQLite 或训练日志。 ## 上游与许可证 diff --git a/context.md b/context.md new file mode 100644 index 00000000..9dd38e4b --- /dev/null +++ b/context.md @@ -0,0 +1,116 @@ +# Code Context + +## Files Retrieved +1. `wasm/web_platform/src/simulation/SimulationSession.ts`(lines 1-107)- 仿真步进、控制器生命周期、模型名称解析和 `mj_step` 前控制写入的核心集成点。 +2. `wasm/web_platform/src/simulation/PhysicsAdapter.ts`(lines 12-110)- UI 与 session 的稳定门面;模型切换、资源释放和 Go2-W URDF 增强发生于此。 +3. `wasm/web_platform/src/app/App.tsx`(lines 37-78, 100-105)- 工程导入、控制器状态、错误反馈、右侧栏 wiring。 +4. `wasm/web_platform/src/app/components/SidebarPanel.tsx`(lines 19-35)- 当前“控制”tab 仅嵌入 Python 控制器,是 ONNX 面板的最小 UI 插槽。 +5. `wasm/web_platform/src/controller/PythonControllerRuntime.ts`(lines 30-120)- 可复用的控制器生命周期/状态设计;同时揭示现有控制回调必须同步。 +6. `wasm/web_platform/src/controller/types.ts`(lines 1-42)- 当前状态被写死为 `language:'python'`,且 step API 可读关节、sensor、body quaternion。 +7. `wasm/web_platform/src/project/importer.ts`(lines 1-48, 184-223)- 所有二进制文件已原样进入 manifest,ONNX 无需进入模型 entry 发现逻辑。 +8. `wasm/web_platform/src/project/types.ts`(lines 1-41)- ProjectFile/大小上限;单文件允许 128 MiB、工程 512 MiB。 +9. `wasm/package.json`(scripts/dependencies)- 已有 Vitest/Playwright,尚无 `onnxruntime-web`。 +10. `wasm/web_platform/vite.config.ts`(lines 1-27)- Pyodide 本地资产复制是 ORT wasm 离线资产处理的直接范例。 +11. `wasm/src/go2_w_balance.py`(lines 1-约480)- 仓库已有 Go2-W 手写 200 Hz 平衡/轮速控制,可作模型命名、IMU 和安全回退参考,但不是 RL 策略。 +12. `/home/cen/Embodied_Workspace/unitree_rl_mjlab/deploy/robots/go2/config/policy/velocity/v0/params/deploy.yaml`(lines 1-56)- 部署时序、动作变换及精确 observation 顺序。 +13. `/home/cen/Embodied_Workspace/unitree_rl_mjlab/deploy/robots/go2/src/State_RLBase.cpp`(lines 6-31)- 部署约定为 `params/deploy.yaml` + `exported/policy.onnx`,输出经 action manager 处理后按 joint map 下发。 +14. `/home/cen/Embodied_Workspace/unitree_rl_mjlab/src/tasks/velocity/config/go2/env_cfgs.py`(lines 1-约160)- flat 任务移除 height scan,Go2 base 为 `base_link`。 + +## Key Code + +### 已有正确插入顺序 +`SimulationSession.singleStep/advance` 均执行:控制器 → 外力 → `mj_step`(`SimulationSession.ts:46-57`)。ONNX 动作必须沿用该位置,而不能在 React frame snapshot 后直接写 ctrl,否则会引入一物理步延迟和不稳定的渲染帧时序。 + +### 参考部署契约 +Go2 deploy YAML 定义: +- `step_dt: 0.02`,即策略 50 Hz;仿真步可保持 0.002 s,动作 hold 10 步。 +- observation 串接顺序和维数:`base_ang_vel(3), projected_gravity(3), velocity_commands(3), gait_phase(2), joint_pos_rel(12), joint_vel_rel(12), last_action(12)`,总计 **47 float32**。 +- action 为 12 维位置目标:`processed = raw * 0.25 + default_joint_pos`;默认位姿见 YAML lines 3-5/18-20。 +- joint 顺序不能取 MuJoCo ID 偶然顺序;部署 `joint_ids_map=[3,4,5,0,1,2,9,10,11,6,7,8]` 表示 RL 顺序与硬件/模型顺序存在显式重排。 +- ONNX 与参数是配套件:`params/deploy.yaml` + `exported/policy.onnx`(`State_RLBase.cpp:12-16`)。仅导入 `.onnx` 无法可靠解释 observation/action。 + +### 最小文件级实现方案(建议,不修改) +1. **新增独立策略层**:`src/policy/types.ts`, `OnnxPolicyRuntime.ts`, `Go2VelocityTask.ts`。 + - `OnnxPolicyRuntime` 只负责 ORT session、输入/输出名称与 shape 校验、generation cancellation、dispose。 + - `Go2VelocityTask` 固化首个 task 的 47→12 契约、命令/phase/lastAction、名称到 joint/actuator 的显式映射、PD 目标或 position actuator 输出。 + - 不应把 ONNX 逻辑塞入 `PythonControllerRuntime`;其 Pyodide proxy 与同步函数假设完全不同。 +2. **`SimulationSession.ts`**:将单一 `pythonController` 最小泛化为互斥 controller slot,或新增 `onnxPolicy` 并在同一个 `runController()` 中互斥调度;load/remove/reset/dispose 复用 generation 防模型切换竞态。添加只读 observation helpers(按 joint/body/sensor 名称解析),避免 UI 从 snapshot 拼向量。禁用/异常时必须 `ctrl.fill(0)`、暂停并保留错误。 +3. **异步风险处理**:`onnxruntime-web` 的 `session.run()` 返回 Promise,不能直接 await 于当前同步 `advance()` while loop。最小侵入方案是“最新完成动作 hold”:到 50 Hz 时复制 observation 发起一次推理(仅允许一个 in-flight),后续物理步继续持有上次动作,完成后原子替换;首帧保持安全默认位姿而非零力矩。状态须显示 inference latency/overrun。若要求训练等价的严格每 20 ms 动作,则需把物理 loop 改为异步/worker,已非最小侵入。 +4. **`PhysicsAdapter.ts`**:接口新增 `loadOnnxPolicy(policyBytes, config)`, `setPolicyEnabled`, `setVelocityCommand`, `removePolicy`;MainThread adapter 只转发,模型 `releaseCurrent()` 自动 dispose。不要把 policy 写入 MuJoCo MEMFS,ORT 可直接吃 `Uint8Array`。 +5. **项目导入**:`importer.ts` 已保留任意扩展名,ZIP/目录中的 `.onnx/.yaml` 会存在 `manifest.files`;无需修改 `discoverEntries`(它只应发现模型)。在 App 中按配套目录约定寻找 `exported/policy.onnx` 及 `params/deploy.yaml`,首版更稳妥的是新增浏览器可直接解析的 JSON policy manifest;不要为 YAML 引入重量解析器后猜测任意训练配置。独立单文件导入应沿用 Python 的 upsert 流程,但 accept 改 `.onnx` 并施加专用上限。 +6. **依赖/构建**:`wasm/package.json` 加 `onnxruntime-web`;`vite.config.ts` 仿照 Pyodide 将 ORT wasm 文件复制到固定相对目录并设置 `ort.env.wasm.wasmPaths`,否则离线 `base:'./'` 构建容易 404。首版使用 wasm CPU 单线程,避免 WebGPU/COOP-COEP 扩大范围。 +7. **UI**:新增 `PolicyPanel.tsx`,插在 `SidebarPanel.tsx:32` 的控制 tab,与 Python 控制器并列;展示模型/config、47×12 shape、50 Hz、启停、速度 x/y/yaw、推理耗时/错误。`App.tsx:40,55,72-77,104` 按现有 controller 状态模式增加 policy 状态及 callbacks;加载另一控制器时明确移除/禁用前一个,不能二者同时写 ctrl。 +8. **首个 Go2-W 任务边界**:参考仓库提供的是 **Go2 12-DOF** velocity policy,不是 Go2-W policy;Go2-W 多四个轮关节。首版可把 12 个腿关节按该 policy 控制、轮 actuator 固定阻尼/零速,仅称“Go2-W 腿式站立/速度实验”,不能宣称与训练分布一致。真正轮式平衡/速度需要 Go2-W 专属训练导出(observation/action 很可能含轮速和 16 维动作)。 + +## Architecture +浏览器文件 → `ProjectManifest.files`(二进制 ONNX 已保留)→ `App` 选择配套 policy/config → `PhysicsAdapter` → `SimulationSession` 拥有 policy runtime。每个 MuJoCo 子步前 task 从 `MjData` 组 observation;每 0.02 s 异步提交 ONNX,完成输出经 scale/offset、显式 joint mapping 和限幅后成为 held action;物理 loop 每步应用 held action。snapshot 仅向 UI报告状态,不作为 observation 数据源。 + +## 观测/动作风险与严重度 +- **blocker**:参考资产无 Go2-W ONNX,且仓库参考目录下 Go2 velocity 的 `exported/policy.onnx` 实际不存在;不能验证 input/output 名或数值一致性。 +- **blocker**:Go2 12-DOF policy 与 Go2-W 16 actuator 拓扑不匹配。若按 `model.njnt/nactuator` 顺序拼接会 shape 错误或静默错控。 +- **high**:`projected_gravity` 必须是重力单位向量由世界系旋到 body/IMU 局部系。MuJoCo `xquat` 为 wxyz;符号、quat inverse、IMU 安装姿态任一错误都会使策略立刻摔倒。加速度计并不等价于无噪声 projected gravity。 +- **high**:`base_ang_vel` 需 body/local frame;直接使用 free joint `qvel[3:6]` 或 gyro sensor前必须确认 MuJoCo sensor frame与训练定义一致。 +- **high**:`joint_pos_rel = q-default`,joint_vel 和 action 都必须是 YAML 的 RL 顺序;四足 FL/FR/RL/RR 次序在当前手写 Go2-W controller 与 deploy map 中并不天然一致。 +- **high**:输出是 normalized action,不是 torque;必须 scale+offset 后交给 position actuator/PD。当前 `setActuator` 只写 ctrl,URDF 自动生成的是 motor(`PhysicsAdapter.ts:68`),因此若直接写位置值会被当 N·m。需要 task 内 PD torque,或为该 task 构造 position actuator。 +- **medium**:gait phase 是 0.6 s 周期的 sin/cos,但应确认零点、sin/cos 顺序以及站立零命令时是否冻结;YAML 没编码完整函数语义。 +- **medium**:异步 ORT latency 若超过 20 ms 会跳过策略 tick;必须统计 dropped/late inference,且避免并发堆积。 +- **medium**:reset 应同时清 lastAction、phase、in-flight generation 和 held action;否则旧 Promise 可污染新 episode。 +- **medium**:导入不可信 ONNX 可能造成大内存/计算消耗;现有 128 MiB 文件上限不足以约束 tensor shape/运行时内存。 + +## 测试建议 +1. `src/policy/Go2VelocityTask.test.ts`:固定 qpos/qvel/quaternion/command,逐元素断言 47 维顺序、dtype、projected gravity、phase、joint reorder;固定 raw action 断言 `*0.25+offset` 和 actuator 映射。 +2. `src/policy/OnnxPolicyRuntime.test.ts`:mock ORT,覆盖 input/output shape/name 错误、NaN/Inf、single-flight、超时/overrun、generation cancellation、dispose;用极小 identity ONNX fixture 做一次真实 wasm smoke test。 +3. 新增 `SimulationSession.test.ts`(当前该核心类无测试):mock MainModule 验证 controller 在 `mj_step` 前执行、50 Hz decimation、held action、reset/disable/error 清 ctrl、Python/ONNX 互斥、模型切换后旧推理无效。 +4. `PhysicsAdapter` contract test:load/remove/dispose 转发及失败后 workspace/runtime 都释放。 +5. `importer.test.ts`:ZIP/目录保留 `.onnx` 与 config、路径规范化、重复/超限拒绝;entry discovery 不把 ONNX 当模型。 +6. `PolicyPanel.test.tsx` + `SidebarPanel` 测试:导入、启停、命令范围、shape/错误展示、另一控制器启用时互斥。 +7. Playwright:导入 Go2-W 工程+fixture policy,启用后推进若干秒,断言页面不崩、policy 状态与命令变化;数值验收应另做离线 golden trajectory(相同初态前 N 次 observation/action 与 Python/ORT reference 对齐),容差逐元素约 `1e-5`,并检查 base height/倾角安全阈值。 +8. 常规命令:`npm run typecheck:platform`, `npm run lint:platform`, `npm run test:platform`, `npm run build:platform`, 最后 `npm run test:e2e:platform`。 + +## Start Here +先打开 `wasm/web_platform/src/simulation/SimulationSession.ts:44-93`:这是唯一能保证 observation 采样、控制输出和 `mj_step` 顺序正确,并统一 reset/error/dispose 语义的位置。 + +## Residual Risks +- 未获得 Go2-W 专属 ONNX 与导出 metadata,无法完成真实 shape/name/golden 验证。 +- unitree_rl_mjlab 的 YAML 说明字段顺序与缩放,但 gait phase、frame convention 等函数语义仍需用实际导出/运行器确认。 +- 主线程异步 inference 的 held-action 方案是最小侵入折衷,不提供硬实时或训练环境严格等价。 + +```acceptance-report +{ + "criteriaSatisfied": [ + { + "id": "criterion-1", + "status": "satisfied", + "evidence": "已给出 SimulationSession、PhysicsAdapter、App/Sidebar、导入与测试结构的具体文件/行范围方案,并按 blocker/high/medium 列出观测动作风险。" + } + ], + "changedFiles": [], + "testsAddedOrUpdated": [], + "commandsRun": [ + { + "command": "find/grep/read/nl(只读仓库与 unitree_rl_mjlab 参考配置)", + "result": "passed", + "summary": "确认现有控制步序、项目导入行为、测试布局及 Go2 deploy 的 47维观测/12维动作约定。" + } + ], + "validationOutput": [ + "参考 deploy.yaml: step_dt=0.02,观测 3+3+3+2+12+12+12=47,动作 12,scale=0.25。", + "参考部署约定为 params/deploy.yaml + exported/policy.onnx;当前 Go2 目录未发现实际 ONNX。", + "当前仓库 package dependencies 未包含 onnxruntime-web。" + ], + "residualRisks": [ + "缺少 Go2-W 专属 ONNX/metadata,参考仅为 Go2 12-DOF。", + "异步浏览器 ORT 与同步物理 loop 存在时序折衷。", + "projected gravity、角速度 frame、joint mapping 必须做 golden 对齐。" + ], + "noStagedFiles": true, + "diffSummary": "只读分析;未修改仓库文件,仅写入指定 context.md 报告。", + "reviewFindings": [ + "blocker: /home/cen/Embodied_Workspace/unitree_rl_mjlab/deploy/robots/go2 - 未发现 Go2 velocity policy.onnx,无法验证模型 I/O。", + "blocker: deploy.yaml:1-56 - 参考策略为 12-DOF Go2,与 Go2-W 16 actuator 不匹配。", + "high: wasm/web_platform/src/simulation/SimulationSession.ts:49-57 - 同步物理 while loop 无法直接 await ORT Promise。", + "high: wasm/web_platform/src/simulation/PhysicsAdapter.ts:68-70 - URDF 自动生成 motor,normalized position action 不可直接写 ctrl。" + ], + "manualNotes": "报告已写入权威路径;除 context.md 外未改动项目。" +} +``` diff --git a/package-lock.json b/package-lock.json index e4fbea06..d31f65c8 100644 --- a/package-lock.json +++ b/package-lock.json @@ -18,6 +18,7 @@ "pyodide": "^0.29.4", "react": "^19.2.8", "react-dom": "^19.2.8", + "uplot": "^1.6.32", "zustand": "^5.0.15" }, "devDependencies": { @@ -5091,6 +5092,12 @@ "browserslist": ">= 4.21.0" } }, + "node_modules/uplot": { + "version": "1.6.32", + "resolved": "https://registry.npmjs.org/uplot/-/uplot-1.6.32.tgz", + "integrity": "sha512-KIMVnG68zvu5XXUbC4LQEPnhwOxBuLyW1AHtpm6IKTXImkbLgkMy+jabjLgSLMasNuGGzQm/ep3tOkyTxpiQIw==", + "license": "MIT" + }, "node_modules/uri-js": { "version": "4.4.1", "resolved": "https://registry.npmjs.org/uri-js/-/uri-js-4.4.1.tgz", diff --git a/package.json b/package.json index 3fd8dfdb..97f82da5 100644 --- a/package.json +++ b/package.json @@ -61,6 +61,7 @@ "pyodide": "^0.29.4", "react": "^19.2.8", "react-dom": "^19.2.8", + "uplot": "^1.6.32", "zustand": "^5.0.15" }, "engines": { diff --git a/plans/map-v3-design.md b/plans/map-v3-design.md new file mode 100644 index 00000000..b56e3781 --- /dev/null +++ b/plans/map-v3-design.md @@ -0,0 +1,193 @@ +# MuJoCo Web 地图导入与编辑设计 + +> 状态:V1/V2 与 V3 地图编辑能力已实现。当前阶段只聚焦地图导入、可视化、受约束编辑、事务应用和地图包导出。 + +## 1. 目标 + +平台在不修改 `@mujoco/mujoco` WASM 内核的前提下,支持: + +1. 内置 MJCF 物理地图:平地、坡道、楼梯和随机障碍物。 +2. 工程地图包:`map.json`、简化 MJCF 碰撞层、可选自包含 GLB 视觉层和出生点。 +3. Schema V2 创作层:通过 `map.scene.json` 编辑受支持的静态原语。 +4. 将完全可识别的只读静态 MJCF 显式转换为可编辑副本。 +5. 在浏览器中事务式重新编译,并导出独立地图 ZIP。 + +当前阶段不扩展机器人运动控制或其他地图上层应用。 + +## 2. 坐标与分层 + +所有地图统一使用: + +- 米制; +- Z-up; +- 右手坐标系; +- `+X` 前方; +- 物理层、视觉层、创作层和出生点共享世界原点。 + +地图拆分为三类资产: + +- `physics/world.xml`:MuJoCo 使用的简化静态碰撞层; +- `visuals/scene.glb`:Three.js 使用的高精度视觉层; +- `authoring/map.scene.json`:编辑器的规范源。 + +高面数视觉模型不得直接作为大型地图碰撞网格。 + +## 3. 地图包结构 + +```text +maps/warehouse/ +├── map.json +├── physics/ +│ ├── world.xml +│ └── meshes/*.obj +├── visuals/ +│ └── scene.glb +└── authoring/ + └── map.scene.json +``` + +Schema V1 支持物理层、视觉层和出生点。Schema V2 增加创作层: + +```json +{ + "schemaVersion": 2, + "id": "warehouse", + "name": "仓库", + "coordinateSystem": { "units": "m", "up": "Z", "forward": "+X" }, + "physics": { "source": "physics/world.xml" }, + "visual": { "source": "visuals/scene.glb" }, + "authoring": { "source": "authoring/map.scene.json" }, + "spawnPoints": [ + { "id": "main", "name": "主入口", "position": [0, 0, 0.35], "yawDeg": 0 } + ] +} +``` + +声明 `authoring.source` 时必须同时声明 `physics.source`。 + +## 4. 安全导入 + +所有地图文件必须进入 `ProjectManifest` 和 `MemfsWorkspace`,禁止运行时直接读取任意主机路径。 + +资源引用必须拒绝: + +- HTTP/HTTPS 和其他协议; +- 绝对路径、UNC、盘符路径; +- NUL、编码绕过和路径穿越; +- 不存在的工程资源; +- 重复路径及超出导入配额的文件。 + +视觉层仅支持自包含 GLB 2.0。暂不支持外部 glTF 依赖、Draco、KTX2 或视觉层隐藏坐标修正。 + +## 5. 物理地图约束 + +工程物理层只允许静态 `worldbody` 和必要的基础 asset。拒绝: + +- joint 和动态 body; +- mocap; +- actuator、sensor、tendon、equality; +- include; +- default class; +- 依赖不兼容 compiler 角度语义的姿态; +- 无法安全重写的资源路径。 + +地图碰撞 geom 使用 `group="2"`,并通过稳定命名空间 `__platform_map__...` 与机器人隔离。 + +## 6. 创作文档 + +`map.scene.json` 是唯一编辑规范源,支持: + +- box; +- cylinder; +- capsule; +- ramp; +- stairs; +- 出生点。 + +每个对象包含稳定 ID、名称、类型、世界位姿、原语参数、摩擦、颜色和启用状态。对象不超过 2,000 个,出生点不超过 500 个,生成 geom 不超过 10,000 个。 + +`MapDocumentCompiler` 根据创作文档确定性生成 `physics/world.xml`。生成结果只使用 quaternion,不依赖 Euler 角。 + +## 7. 编辑交互 + +地图面板提供“场景 · 资产库 → 认证资产”入口。用户可以点击添加,或把资产卡片拖到画布的 XY 平面落位;落点按 0.1 m 归一化并根据原语尺寸自动贴地。当前没有可编辑地图时,首个资产会立即创建 `map.json`、空物理层和创作层骨架,切换到新场景并作为未应用草稿进入 Three.js 预览;此阶段不调用 MuJoCo。用户点击“应用并重新编译”后才生成并事务提交实际物理层。 + +React 中的 `MapEditSession` 是草稿唯一来源。Three.js 预览层只负责: + +- 草稿原语显示; +- 射线拾取和高亮; +- 世界坐标平移; +- 绕世界 Z 轴旋转; +- 原语尺寸缩放; +- 出生点预览; +- 操纵器和临时资源生命周期。 + +交互规则: + +- `W`:移动; +- `E`:旋转; +- `S`:缩放; +- `Delete`/`Backspace`:删除; +- `Escape`:取消选择; +- `Ctrl+Z`:撤销; +- `Ctrl+Y` 或 `Ctrl+Shift+Z`:重做。 + +拖动 TransformControls 操纵轴期间禁用 OrbitControls;在画布空白区域按住鼠标左键仍可旋转相机。连续拖动只在 `mouseUp` 时提交一次历史记录。缩放必须写回原语参数,不能把 Three.js 节点 scale 作为持久数据。 + +## 8. 只读地图转换 + +外部 MJCF 默认只读。只有完全可逆的静态原语地图可显式创建可编辑副本。 + +允许转换: + +- box; +- cylinder; +- capsule; +- 由嵌套静态 body 组成的世界位姿; +- pos、quat、size、friction、rgba 和出生点。 + +遇到下列内容必须整体拒绝,禁止静默丢弃: + +- plane、mesh、heightfield; +- asset、材质和碰撞过滤扩展; +- joint、site、light 或其他未知结构; +- euler、axisangle、xyaxes、zaxis、fromto; +- 不能确认无损的 compiler 配置。 + +## 9. 事务边界 + +应用编辑草稿或转换地图时,顺序固定为: + +1. 生成候选创作层; +2. 确定性生成候选物理层; +3. 创建候选 manifest; +4. 在新的 MEMFS 工作区编译 MuJoCo; +5. Viewer 成功 attach 新会话; +6. 最后提交 manifest、地图选择和编辑状态; +7. 释放旧会话和旧工作区。 + +任一步失败都必须恢复旧仿真和旧工程状态,同时保留用户草稿。 + +## 10. 导出 + +浏览器导入文件不保证可写,因此不直接覆盖源目录。地图通过 ZIP 导出,包含: + +- `map.json`; +- 物理层及其显式 asset; +- 可选视觉层; +- 可选创作层。 + +ZIP 内路径保持工程相对结构,并继续执行路径安全校验。 + +## 11. 验证重点 + +- Schema V1/V2 兼容性; +- 路径穿越和协议绕过拒绝; +- 地图物理合成和命名空间隔离; +- GLB 自包含校验及资源释放; +- 编辑文档严格校验; +- 确定性 MJCF 输出; +- Undo/Redo 和视口变换写回; +- 只读转换的白名单与整体拒绝; +- 编译或 Viewer attach 失败后的事务回滚; +- 地图 ZIP 资产完整性。 diff --git a/plans/reward-auto-tuning-agent.md b/plans/reward-auto-tuning-agent.md new file mode 100644 index 00000000..95170471 --- /dev/null +++ b/plans/reward-auto-tuning-agent.md @@ -0,0 +1,104 @@ +# 奖励函数自调参 Agent 实施计划 + +## Context + +当前项目已有完整的本地训练链路:React/Vite 前端通过 `LocalTrainingClient` 调用仅监听 loopback 的 Python `training_server`,服务再启动内置 `training_server/rl` 中的 Go2 + mjlab/RSL-RL 训练器,并产出 `policy.onnx`。现有请求只包含环境数、迭代数、随机种子、设备和 W&B 模式;奖励项与权重固定在任务配置中,服务主要从标准输出解析迭代进度,尚未向前端提供结构化奖励曲线或自动调参循环。 + +第一版范围已确定为 `Unitree-Go2-Flat`,优化优先级依次为:速度跟踪、动作平滑、姿态稳定、减少跌倒、足端滑移、能耗。Agent 可以调整现有奖励权重、启停白名单奖励项以及修改获准的阈值/核宽等参数,但不能生成或执行任意 Python 奖励代码。系统同时提供全自动与逐轮审批模式,在本地 RTX 5080 上串行训练,通过云端 API 调用 Agent;独立打开一个参考 Isaac/TensorBoard 交互方式的监控网页,展示曲线、trial 对比、参数 diff 与最佳策略。 + +目标是在现有安全边界内增加一个可审计、可暂停、可恢复、可回退的调参闭环。云端只接收裁剪后的数值配置、曲线摘要和评估指标,不接收机器人资产、checkpoint、源代码、训练服务访问令牌或本地路径。已确认使用 DeepSeek 官方 OpenAI-compatible API(Base URL `https://api.deepseek.com`)和精确模型标识 `deepseek-v4-flash`;API key 仅从训练服务环境变量读取。目标占比采用 `35/20/15/15/10/5`,逐轮审批时基线自动运行、之后每个建议等待批准。最佳结果保存为命名 preset 并可用于后续训练,不覆盖仓库内 Python 默认值。 + +## Approach + +采用“确定性试验编排 + 独立评估 + 云端 Agent 建议”的分层方案,而不是让 LLM 直接改源码或决定 trial 是否有效: + +1. **白名单参数空间**:为每个 reward term 定义固定符号、默认值、上下界、是否允许置零及单轮最大变化;开放权重,以及 `std`、`command_threshold`、`target_height`、步态 `period/threshold`、姿态分段容差等少量参数。保留当前 15 个奖励项,并新增默认关闭(权重为 0)的 `electrical_power`,以覆盖能耗目标。`track_linear_velocity`、`track_angular_velocity`、`body_orientation_l2`、`is_terminated`、`joint_pos_limits`、`action_rate_l2` 不允许关闭;其余白名单项可置零。每个 proposal 最多改 4 个标量,非零权重幅值单轮限制在前值的 `0.5×–2×`(同时受绝对上下界约束),符号不可翻转,Go2 trot 的足序 offset 不开放。Agent 只能返回结构化 patch,服务端合并并二次校验。 +2. **权重无关的质量指标**:现有 `Episode_Reward/*` 已由 mjlab `RewardManager` 自动记录,但它随权重变化,不能直接作为优化目标。新增固定定义的速度误差、动作加速度、姿态误差、跌倒率、接触足滑移速度和正向机械功率指标;训练曲线用于诊断,最终排名采用固定命令集与固定评估种子得到的这些指标,避免通过放大奖励权重“刷高总奖励”。 +3. **指标与产物管线**:调参 trial 强制使用本地 TensorBoard writer;训练脚本接受服务端生成的明确输出目录与奖励配置文件,保存 `env.yaml`、Agent patch、checkpoint、ONNX 和评估结果。服务通过 TensorBoard EventAccumulator 增量读取 scalar 并写入 SQLite,曲线 API 按 LTTB/桶聚合降采样;不依赖脆弱的控制台正则解析指标。 +4. **独立评估与评分**:增加无探索噪声的评估入口,使用同一组站立、前进/侧移、转向和组合速度命令,对每个 rung 的 checkpoint 运行相同场景。速度目标内部按线速度/角速度误差 `80/20` 合并;六个顶层目标按已确认的 `35% / 20% / 15% / 15% / 10% / 5%` 聚合。对所有“越低越好”的原始指标使用创建 session 时冻结的基线尺度 `scale=max(abs(baseline), physical_floor)`,计算并裁剪相对改善 `(baseline-current)/scale`;跌倒率不得高于基线 `+2%`,速度误差不得恶化超过 `5%`,否则该 trial 不可晋级。输出原始指标、各目标改善、总分和 3-seed 均值/离散度,确保后续新增 trial 不会改变旧 trial 的分数。 +5. **DeepSeek Agent 与数值搜索协作**:采用 PydanticAI 的 OpenAI-compatible provider 连接 `deepseek-v4-flash`,以严格类型的 `RewardProposal` 返回最多 4 项参数 patch、依据、预期影响和置信度;优先使用模型 JSON/structured-output 能力,能力探测失败时退回 PydanticAI 的 prompted JSON + 本地 Pydantic 校验,不授予 Agent 任何 shell、文件或网络工具。Optuna study 记录完整参数/分数并执行 successive-halving;Agent proposal 作为 enqueue/fixed trial 进入 study,只有显式启用 fallback 时才由 Optuna sampler 代提候选。模型使用低温度、60 秒超时和最多 2 次结构化修复;每次仅发送最多 12 个 trial 摘要、每条曲线最多 64 个降采样点,并记录脱敏 prompt hash、模型名、token usage、批准操作与最终 patch。服务对输出执行有限数值、符号、边界、最大步长、重复配置和高风险组合校验,校验失败要求 Agent 修正,不能静默执行。 +6. **session 状态机**:基线自动运行 → 短预算 trial → 固定评估 → Agent 建议 →(自动批准或进入 `awaiting_approval`)→ 下一 trial → successive-halving 晋级 → 最佳配置复核。逐轮审批支持接受、拒绝并附反馈、手动修改后接受;暂停不杀死已完成数据,停止会终止当前进程组。完成后把最佳 reward patch 保存为不可变命名 preset,并提供“从 preset 新建普通训练/新 tuning session”、导出 JSON、下载/导入 ONNX;不写回 `velocity_env_cfg.py`。 +7. **默认 RTX 5080 预算**:12 个唯一配置(含基线)、4096 个并行环境、GPU `0`;所有配置先训练 300 iterations,前 4 名从自身 checkpoint 续训到 900,前 2 名续训到 2000,总量约等于 4.1 次完整 2000-iteration 训练。连续 4 个建议无显著提升时提前停止;trial 数、环境数和各 rung 可在 4–20 / 合法服务范围内调整。搜索期固定训练 seed 控制方差,晋级候选使用 3 个固定评估 seed 复核。 +8. **独立监控网页**:新增 Vite 多页面入口 `tuning.html`,从现有训练面板用新标签页打开。页面采用 TensorBoard 风格的 run 选择、平滑、缩放、悬浮值、标签过滤和多 trial 叠加图,并增加 Agent 决策时间线、API 连通性测试、审批卡片、参数 diff、排行榜、暂停/恢复/停止、preset 导出及最佳 ONNX 下载/导入。新标签页 URL 只携带 session ID;训练服务 token 通过同源、校验 origin 的一次性 `postMessage` 交接并仅存于新标签页 `sessionStorage`,失败时回退到手工输入,绝不放入 query/hash。“导入最佳策略”由 dashboard 向仍打开的 workbench 发送同源消息,workbench 使用自身 client 下载并调用现有 `onPolicyReady`;无 opener 时回退为文件下载。图表使用轻量 `uPlot`,不启动或 iframe 嵌入第二个 TensorBoard 服务。 +9. **API 与落盘边界**:新增 capabilities/Agent 测试、session 创建与详情、trial/metrics 查询、proposal 批准/拒绝、pause/resume/cancel、preset 列表/导出和最佳 artifact 下载接口;继续沿用现有 Bearer Token、Host/Origin 校验与 32 KiB 请求限制。SQLite 使用 WAL 和每线程连接,默认位于已忽略的 `training_server/rl/logs/auto_tuning/tuning.sqlite3`,trial 产物位于同目录的 session 子目录;API 永远只接受 ID,不接受客户端文件路径。启动时把遗留 `training/evaluating` 状态标记为 `interrupted`,从最后完整 checkpoint 显式恢复,不尝试盲目重连旧 PID。 +10. **DeepSeek 配置**:使用 `DEEPSEEK_API_KEY`、`MUJOCO_TUNING_AGENT_BASE_URL`(默认 `https://api.deepseek.com`)和 `MUJOCO_TUNING_AGENT_MODEL`(默认 `deepseek-v4-flash`);可配置 timeout,但 API key 不提供 CLI 参数,避免进入 shell history。健康接口仅返回 `configured/model/baseUrl`,连接测试返回能力与脱敏错误,不返回 key 或完整供应商响应。未配置云端 key 时普通训练保持可用,tuning capability 明确显示不可用;自动模式不静默退回 Optuna,只有用户在 session 中显式勾选 fallback 才允许。 + +初始权重白名单如下;负项只能保持负号,正项只能保持正号。绝对边界与单轮 `0.5×–2×` 限制同时生效,实际启用前用基线 smoke test 校验量纲: + +| Reward term | 当前值 | 允许范围 | 可关闭 | +| --- | ---: | ---: | :---: | +| `track_linear_velocity` | 1.0 | 0.5–3.0 | 否 | +| `track_angular_velocity` | 1.0 | 0.25–2.0 | 否 | +| `body_orientation_l2` | -1.0 | -3.0–-0.1 | 否 | +| `pose` | 1.0 | 0–2.5 | 是 | +| `body_ang_vel` | -0.05 | -0.2–0 | 是 | +| `angular_momentum` | -0.025 | -0.1–0 | 是 | +| `is_terminated` | -200 | -400–-50 | 否 | +| `joint_acc_l2` | -2.5e-7 | -2e-6–0 | 是 | +| `joint_pos_limits` | -10 | -30–-2 | 否 | +| `action_rate_l2` | -0.05 | -0.2–-0.005 | 否 | +| `foot_gait` | 0.5 | 0–1.5 | 是 | +| `foot_clearance` | -1.0 | -3.0–0 | 是 | +| `foot_slip` | -0.25 | -1.0–0 | 是 | +| `soft_landing` | -1e-3 | -5e-3–0 | 是 | +| `stand_still` | -1.0 | -3.0–0 | 是 | +| `electrical_power`(新增) | 0 | -5e-3–0 | 是 | + +参数白名单限制为:线速度 `std=0.25–1.0`、角速度 `std=0.35–1.2`;`pose` 三档 std 使用当前 Go2 数组的 `0.5×–2×` 缩放因子,walking/running threshold 分别为 `0.05–0.5` / `1.0–2.5` 且保持有序;`foot_gait.period=0.4–0.8`、`threshold=0.45–0.65`;`foot_clearance.target_height=0.06–0.16`;各运动相关 `command_threshold=0.02–0.30`。结构对象、函数名、传感器名、asset selector、步态 offset 和终止角度不开放。 + +## Files to modify + +关键修改与新增路径: + +- `training_server/rl/src/tasks/velocity/velocity_env_cfg.py`:挂载固定质量指标与可选能耗奖励项。 +- `training_server/rl/src/tasks/velocity/config/go2/env_cfgs.py`:Go2 参数默认值及白名单配置应用入口。 +- `training_server/rl/src/tasks/velocity/mdp/metrics.py`(新增)与 `mdp/__init__.py`:权重无关的六类质量指标。 +- `training_server/rl/src/tasks/velocity/mdp/rewards.py`:仅在现有函数无法覆盖白名单参数时补充实现;优先复用 mjlab 内置项。 +- `training_server/rl/scripts/train.py`:显式 run 目录、奖励 patch、checkpoint 续训和 TensorBoard 配置。 +- `training_server/rl/scripts/evaluate.py`(新增):固定命令/seed 的 checkpoint 评估与 JSON 结果。 +- `training_server/rl/src/tasks/velocity/rl/runner.py`:保持 checkpoint/ONNX 对应关系,必要时暴露最终 checkpoint 元数据。 +- `training_server/tuning/`(新增):schema/目标函数、SQLite storage、TensorBoard ingest、PydanticAI advisor、Optuna sampler 与 session orchestrator,避免继续膨胀单文件服务。 +- `training_server/server.py`:组合现有训练 manager 与 tuning manager,增加 `/api/tuning/*` 路由和启动配置。 +- `training_server/tests/test_server.py` 及新增 `training_server/tests/test_tuning_*.py`:API、状态机、存储、Agent 校验与 fake trainer 集成测试。 +- `training_server/requirements.txt`(新增)与 `requirements-dev.txt`:加入并固定经 Python 3.12 实测的 `pydantic-ai-slim[openai]`、Optuna 和 TensorBoard;保留 `training_server/rl/requirements.txt` 只承载 mjlab 训练栈,避免职责混杂。 +- `web_platform/tuning.html`、`web_platform/src/tuning/`(新增):独立监控入口、dashboard、uPlot 曲线、审批/参数 diff/排行榜。 +- `web_platform/src/training/types.ts` 与 `LocalTrainingClient.ts`:调参 API 类型与方法;必要时按职责拆出 `TuningClient.ts`。 +- `web_platform/src/training/LocalTrainingPanel.tsx`:创建 session 的基础入口及“在新网页打开”按钮。 +- `web_platform/vite.config.ts`:多页面构建入口;`package.json` / lockfile 增加 `uplot`。 +- `.gitignore`、`README.md`、`training_server/README.md`、`web_platform/README.md`:忽略本地状态/产物并记录安装、安全和使用流程。 + +## Reuse + +- 复用 `training_server/server.py` 的 Bearer Token、loopback/CORS 限制、无 shell 参数数组、训练任务互斥、进程组取消和 ONNX 下载机制;tuning 与普通训练共享同一个 GPU 活动锁。 +- 复用 `training_server/rl/scripts/train.py` 的 tyro/dataclass 配置和现有 checkpoint resume 流程,新增 patch 应用层而不是改写源文件。 +- 复用 mjlab `RewardManager` 自动产生的 `Episode_Reward/`、`MetricsManager` 的 `Episode_Metrics/`、`TerminationManager` 的 `Episode_Termination/`,以及 RSL-RL TensorBoard scalar writer。 +- 复用当前 `mean_action_acc`、`feet_slip`、`body_orientation_l2` 等计算;能耗优先复用 mjlab 的 `electrical_power_cost`,不重复实现扭矩功率算法。 +- 复用 `VelocityOnPolicyRunner.save()` 的 checkpoint + `policy.onnx` 同步导出和元数据附加;晋级 trial 从自己的 checkpoint 恢复。 +- 参考已安装 mjlab 的 `scripts/play.py` / tracking evaluate 结构实现仓库内最小速度任务评估入口。 +- 复用 `LocalTrainingClient` 的鉴权/错误处理、`LocalTrainingPanel` 的连接信息与取消/策略下载交互,以及共享 `Button`、`Badge`、`Tabs`、`ProgressBar` 等 UI。 +- 当前前端没有路由器或图表库,故使用 Vite MPA 而非引入整套路由;图表仅新增面向大量 scalar 的 `uPlot`。 + +## Steps + +- [x] 固化 `Unitree-Go2-Flat` 的 15 个现有 reward term schema:当前值、符号、上下界、启停规则、可调参数、每次最多 4 项、`0.5×–2×` 变化率和跨参数约束,并为能耗加入默认关闭项。 +- [x] 增加权重无关的六类训练/评估指标,定义归一化方向、目标阈值、默认 `35/20/15/15/10/5` 聚合权重、失败/NaN/过早跌倒惩罚。 +- [x] 扩展训练入口以接收服务生成的 patch 文件与明确 run 目录,支持同配置 checkpoint 晋级续训,并强制保存完整配置快照。 +- [x] 实现固定命令集、固定 seed、无探索噪声的独立评估入口,产生可验证 JSON 与 TensorBoard scalars。 +- [x] 实现 TensorBoard scalar 增量采集、SQLite schema/migration、曲线降采样及 session/trial/proposal/audit 持久化。 +- [x] 实现 tuning orchestrator、successive-halving、共享 GPU 互斥、提前停止、崩溃恢复、暂停/取消和最佳产物选择。 +- [x] 通过 PydanticAI OpenAI-compatible provider 接入官方 `deepseek-v4-flash` 和 Optuna study,实现环境变量密钥、连接/能力测试、脱敏上下文、结构化输出、重试/修复、超时、usage 审计与显式 Agent 不可用状态。 +- [x] 实现自动批准与 `awaiting_approval` 两条状态路径,包括接受、拒绝反馈、手动修订后接受及完整审计记录。 +- [x] 扩展 `/api/tuning/*`、健康信息、指标查询和最佳 ONNX 下载,并补齐 TypeScript 类型/客户端。 +- [x] 新增 `tuning.html` TensorBoard 风格 dashboard 与现有面板的新标签页入口,完成曲线、筛选/平滑/缩放、trial 对比、Agent 时间线、审批、diff、排行榜和控制操作。 +- [x] 添加 Python/TypeScript/组件/E2E 测试,更新依赖锁、忽略规则、安全说明和安装/使用文档。 + +## Verification + +- 单元测试:参数 schema 与边界、目标分数、建议约束、状态转换、失败/取消/恢复、最佳 trial 选择、API 鉴权。 +- 集成测试:使用轻量 fake trainer 产生确定性 TensorBoard events/评估 JSON,并用 fake PydanticAI model 返回合法、越界、重复和畸形 proposal,完整跑通多 trial 调参而不依赖 GPU 或真实云端 API。 +- 前端测试:表单校验、轮询、曲线渲染、参数 diff、自动/审批分支、跨标签页无 URL 密钥交接、停止、preset 与最佳结果操作。 +- 真实训练 smoke test:在本地 RTX 5080 / GPU 0 上先用 256–512 environments、10–20 iterations 跑基线与 2 个 trial,验证 TensorBoard 指标、评估 JSON、checkpoint、ONNX、参数快照和数据库记录一一对应;再单独确认 4096 environments 不 OOM 后启用默认预算。 +- 确定性检查:相同 checkpoint + 相同命令/seed 的评估分数在容差内一致;改变 reward 权重不会直接改变固定质量指标的定义或归一化。 +- 恢复检查:分别在训练、评估、等待审批时重启服务;确认 session 从 SQLite 恢复且不会重复启动 trial,活动子进程能被停止。 +- 回归:`npm run typecheck`、`npm run lint`、`npm run test`、`npm run test:training-server`、`npm run build`、`npm run test:e2e`;具备依赖时执行 `npm run lint:python`。 +- 安全检查:非法奖励名、符号翻转、越界/NaN/Inf、超大 patch、路径注入、重复 proposal、并发普通训练/session、伪造 artifact 路径和 Agent 超时均被拒绝或进入明确状态;API key 不出现在响应、日志、SQLite 或浏览器存储中。 diff --git a/plans/urdf-studio-ui-components.md b/plans/urdf-studio-ui-components.md new file mode 100644 index 00000000..d49faf31 --- /dev/null +++ b/plans/urdf-studio-ui-components.md @@ -0,0 +1,178 @@ +# 参考 URDF-Studio 补全前端组件计划 + +## Context + +目标是在不改变现有 MuJoCo 导入、编译、仿真、渲染和资源生命周期逻辑的前提下,参考本机 `/home/cen/Embodied_Workspace/Mujoco_Projects/URDF-Studio` 的前端信息架构与视觉样式,逐项补全 `wasm/web_platform` 缺少的前端组件。 + +关键约束: + +- 只借鉴前端布局、视觉语言与交互组织,不直接迁移 URDF-Studio 的机器人解析、编辑或 Three.js 运行时。 +- 新组件接入前必须先向用户展示候选清单、用途、影响范围并逐项获得确认。 +- 第一批 A–H 已按批准计划实施;后续新增组件仍遵循先确认再接入。 + +初步现状: + +- 本项目 UI 高度集中在 `wasm/web_platform/src/app/App.tsx`,已有顶部工具栏、工程/模型结构左栏、三维视口、模型控制右栏、底部性能状态栏、入口选择、加载态和诊断卡片。 +- MuJoCo 与 Three.js 对象由 `MainThreadPhysicsAdapter`、`MuJoCoViewer` 和 React refs 持有;Zustand 仅保存 UI 可消费快照。后续 UI 重构应保持这一边界。 +- URDF-Studio 已将 Header、WorkspaceSidebars、Viewer overlays、通用 UI 控件、拖动窗口、设置弹窗等拆成独立组件,可作为组件边界和样式参考。 +- 当前资源树和模型结构树已具备纯数据构建函数及单测,不应重写;只需替换字符图标、颜色和行样式,并保留原有树语义。 +- 当前没有 `App.tsx` 级组件测试;后续应通过抽出无状态壳层组件来增加 UI 测试,避免在测试中初始化 WebGL/WASM。 + +已确认产品方向: + +- 第一批聚焦工作台外壳:顶部工具栏、可折叠侧栏、视口工具、底部状态栏、通知与弹窗。 +- 只采用 URDF-Studio 的专业工作台风格,保留 MuJoCo 平台品牌、中文信息架构和现有布局特色,不做像素级复刻。 +- 允许新增 `lucide-react`,用于替代当前字符/Emoji 图标并统一视觉语言。 + +## Approach + +1. 对照两个项目的页面壳层、工具栏、左右侧栏、视口叠层、弹窗/通知和基础控件,形成“已有 / 缺少 / 不适用”的组件差异表。 +2. 将候选组件按纯展示、现有动作封装、需要新增产品能力三类分级;第一目标只推荐前两类,避免触碰底层 MuJoCo 逻辑。 +3. 第一批 A–H 已全部获得用户确认;后续若发现需要新增候选组件,必须再次确认,不能顺带加入。 +4. 先建立精简语义样式令牌和无状态基础组件,再拆分 `App.tsx` 的布局组件;业务回调继续由 `App` 注入,adapter/viewer 生命周期不下沉到展示组件。 +5. 第一批布局采用 40px 紧凑 Header、现有左右固定侧栏和 28px 状态栏;交互模式工具组停靠在 Header 中央。右栏默认展开“模型信息、当前选择、关节”,折叠其他低频分区。 +6. 不照搬 URDF-Studio 的移动端底部工具条、复杂响应式菜单或 Floating UI。保持当前 `min-w-[1024px]` 桌面产品边界。 +7. 每批只接入少量组件并执行回归验证,确保模型导入、仿真和三维交互行为不变。 + +## Files to modify + +- `wasm/web_platform/src/app/App.tsx`:仅保留 viewer/adapter 生命周期、业务状态和事件编排,以 props 连接新 UI 组件。 +- `wasm/web_platform/src/app/components/WorkbenchHeader.tsx`:品牌区、文件导入、仿真传输控制、速度、面板/主题动作与中央工具栏插槽。 +- `wasm/web_platform/src/app/components/ViewerToolDock.tsx`:select/joint/force 模式与相机复位的图标工具组。 +- `wasm/web_platform/src/app/components/SidebarPanel.tsx`:左右侧栏壳层、标题、折叠分区和滚动区域。 +- `wasm/web_platform/src/app/components/StatusBar.tsx`:性能、WASM 和快捷键状态。 +- `wasm/web_platform/src/app/components/EntrySelectionDialog.tsx`:多入口选择的领域弹窗。 +- `wasm/web_platform/src/app/components/DiagnosticNotice.tsx`:诊断摘要、路径和可展开详情。 +- `wasm/web_platform/src/app/components/WorkspaceOverlays.tsx`:加载 HUD、空工作区和中心叠层组织。 +- `wasm/web_platform/src/components/ui/Button.tsx`、`IconButton.tsx`、`Tooltip.tsx`、`Select.tsx`、`Dialog.tsx`、`CollapsibleSection.tsx`、`ToolbarToggleGroup.tsx`:第一批基础控件。 +- `wasm/web_platform/src/components/ui/index.ts`:稳定导出边界。 +- `wasm/web_platform/src/project/ProjectTree.tsx`、`ModelStructureTree.tsx`:仅做令牌化视觉和 Lucide 图标替换,保留构树逻辑及公开 props。 +- `wasm/web_platform/src/app/ErrorBoundary.tsx`:使用新基础控件和语义样式,不改错误边界行为。 +- `wasm/web_platform/src/styles.css`、`wasm/web_platform/tailwind.config.cjs`:颜色、表面、文字、强调色、滚动条、焦点环和组件令牌;移除依赖深色工具类覆盖的主题实现。 +- `wasm/web_platform/src/app/components/*.test.tsx`、`wasm/web_platform/src/components/ui/*.test.tsx`:新增组件测试。 +- `wasm/web_platform/e2e/app.spec.ts`:保持业务断言并更新因图标化/折叠带来的交互定位。 +- `wasm/package.json` / `wasm/package-lock.json`:加入已确认的 `lucide-react`;不引入 `@floating-ui/react`。 + +## Reuse + +本项目继续复用: + +- `wasm/web_platform/src/app/App.tsx` 中的导入、加载、播放/暂停、单步、重置、模式切换及面板状态回调。 +- `wasm/web_platform/src/stores/useAppStore.ts` 的工程、快照、选择、性能和诊断状态。 +- `wasm/web_platform/src/project/ProjectTree.tsx` 与 `ModelStructureTree.tsx`。 +- `wasm/web_platform/src/simulation/PhysicsAdapter.ts` 和 `SimulationSession.ts`(保持不改或仅由现有接口调用)。 +- `wasm/web_platform/src/viewer/MuJoCoViewer.ts`(保持渲染和交互接口不变)。 + +样式/结构参考(只读来源): + +- `/home/cen/Embodied_Workspace/Mujoco_Projects/URDF-Studio/src/app/components/Header.tsx` +- `/home/cen/Embodied_Workspace/Mujoco_Projects/URDF-Studio/src/app/components/workspace/WorkspaceSidebars.tsx` +- `/home/cen/Embodied_Workspace/Mujoco_Projects/URDF-Studio/src/app/components/AppLayoutView.tsx` +- `/home/cen/Embodied_Workspace/Mujoco_Projects/URDF-Studio/src/shared/components/ui/` +- `/home/cen/Embodied_Workspace/Mujoco_Projects/URDF-Studio/src/shared/components/Panel/OptionsPanel.tsx`:参考可折叠分区及紧凑面板标题,不移植拖动/悬浮窗逻辑。 +- `/home/cen/Embodied_Workspace/Mujoco_Projects/URDF-Studio/src/features/urdf-viewer/components/ViewerToolbar.tsx`:参考图标化工具模式组和活动态,不移植 Portal/移动端工具条。 +- `/home/cen/Embodied_Workspace/Mujoco_Projects/URDF-Studio/src/styles/index.css`:仅提炼语义表面、边框、文字、强调色、滚动条和 focus ring;不复制其 AI、编辑器、字体缩放等无关样式。 + +### 第一批组件(已确认) + +| 编号 | 候选组件 | 复用现有能力 | 预期变化 | 底层风险 | +|---|---|---|---|---| +| A ✅ | `WorkbenchHeader` | 导入、播放、单步、重置、速度、主题、面板开关回调 | 40px 紧凑品牌栏,动作分组,Lucide 图标 + 文案/提示 | 低,仅回调透传 | +| B ✅ | `ViewerToolDock` + `ToolbarToggleGroup` | select/joint/force 模式、相机复位 | Header 中央停靠工具组,活动态更清晰 | 低,仅调用现有 mode/viewer API | +| C ✅ | `SidebarPanel` + `CollapsibleSection` | 左右侧栏现有内容和开关 | 统一面板标题、滚动条、分区折叠;保留 resize-x;右栏低频区默认折叠 | 低,内容插槽化 | +| D ✅ | `StatusBar` | 时间、FPS、step、内存、WASM、预算警告 | 图标化状态、语义色和更紧凑层级 | 低,纯展示 | +| E ✅ | `Dialog` + `EntrySelectionDialog` | 多入口 `loadEntry` | Portal、Esc、焦点圈定、焦点恢复和统一弹窗外观 | 低,不改入口判断/加载 | +| F ✅ | `DiagnosticNotice` | `AppDiagnostic` | 摘要以通知样式展示,可展开技术详情;保留手动关闭 | 低,不改错误生成 | +| G ✅ | `LoadingOverlay` + `EmptyWorkspace` | `loading`、空状态、拖放导入 | 统一 HUD、图标、引导层次,不伪造加载百分比 | 低,纯展示 | +| H ✅ | 基础 UI:`Button`、`IconButton`、`Tooltip`、`Select` | 替换 `.btn/.icon-btn/.field` | 统一尺寸、变体、禁用态、focus ring;Tooltip 不增加 Floating UI 依赖 | 低,HTML 原生语义 | + +## Steps + +- [x] 审计 URDF-Studio 页面布局、样式令牌和第一批组件职责。 +- [x] 审计本项目已有组件及其与 MuJoCo/viewer 的耦合边界。 +- [x] 形成第一批组件差异矩阵;排除源码编辑、属性编辑、导出、快照、AI、测量、绘制、撤销/重做、移动端工具条等新增业务能力。 +- [x] 获得用户对第一批 A–H、Header 中央工具组和右栏默认折叠策略的确认。 +- [x] **基础层**:安装 `lucide-react`;定义 light/dark 语义 CSS 变量并映射到 Tailwind;实现 Button、IconButton、Tooltip、Select、Dialog、CollapsibleSection、ToolbarToggleGroup。Tooltip 使用 hover/focus CSS 展示,Dialog 使用 portal、Esc、焦点圈定和焦点恢复。 +- [x] **工作台壳层**:实现 40px `WorkbenchHeader`,左侧保留 MuJoCo 品牌和文件动作,中间放 select/joint/force 工具组,右侧放相机、侧栏和主题动作;保留按钮可访问名称及原有回调。 +- [x] **侧栏拆分**:实现 `SidebarPanel`,将左栏工程资源/模型结构和右栏模型信息、URDF、警告、选择、Actuator、关节、外力内容从 `App.tsx` 搬入展示组件;业务回调和 `adapter.current` 调用仍在 `App` 中生成后传入。右栏“模型信息、当前选择、关节”默认展开,其余默认折叠,警告出现时自动展开。 +- [x] **叠层与反馈**:用 `WorkspaceOverlays` 组织空态/加载态,用 `EntrySelectionDialog` 替换内联入口对话框,用 `DiagnosticNotice` 替换诊断卡;不改变 loading/diagnostic/entries 的状态来源和加载时序。 +- [x] **状态栏与树视觉**:实现 `StatusBar`;为两棵现有树替换 Emoji/字符图标并应用语义令牌,不改 `buildProjectTree`、`buildBodyTree` 或选择/hover 行为。 +- [x] **收敛 App**:删除已迁移的内联 `PanelTitle`、`EmptyImport`、`EntryDialog`、`DiagnosticCard` 等展示函数;将无状态滑杆展示移入侧栏组件,不改数值换算、范围或 onChange 逻辑。 +- [x] **测试与验收**:为基础控件、折叠区、Header 回调、Dialog 焦点/Esc、诊断详情和状态栏增加 Testing Library 测试;更新 E2E 图标/折叠定位并跑完整回归。 + +## Verification + +- 运行 `npm run typecheck:platform --prefix wasm`、`npm run lint:platform --prefix wasm`、`npm run test:platform --prefix wasm`、`npm run build:platform --prefix wasm` 和 `npm run test:e2e:platform --prefix wasm`,全部通过。 +- 新组件测试覆盖:Header 各动作只触发一次;模式按钮 `aria-pressed`;折叠区默认状态;Dialog Esc/Tab/焦点恢复;诊断详情展开与关闭;Loading/Empty 的互斥显示。 +- 用现有 MJCF、URDF、文件夹和 ZIP 路径回归导入及入口切换。 +- 回归播放/暂停、单步、重置、速度、关节、actuator、外力、碰撞显示、选择和相机复位。 +- 检查 adapter/viewer 初始化与销毁次数,确认 UI 拆分未造成重复会话、重复渲染循环或 WASM 资源泄漏。 +- 对照 URDF-Studio 检查布局密度、层级、悬停/选中/禁用态及深浅主题;桌面窄宽度下检查折叠与溢出。 +- 对新增弹窗、菜单和提示执行键盘操作、焦点管理和 ARIA 冒烟检查。 + +## Implementation Result + +- 已完成 A–H 第一批组件和 `lucide-react` 接入;MuJoCo adapter、viewer、manifest 和业务命令仍由 `App.tsx` 持有。 +- `App.tsx` 已收敛为生命周期/事件编排层;工作台 Header、工具组、侧栏、状态栏、弹窗、诊断和加载/空态均拆为展示组件。 +- 深浅主题完成语义令牌化,并在 1440×900 下完成暗色、亮色人工截图检查。 +- 审查发现并修复:必选入口弹窗重渲染抢焦点、原生输入焦点可见性、模型树 ARIA 层级、侧栏重开折叠状态丢失和侧栏按钮展开状态缺失。 +- 验证结果:TypeScript、ESLint、12 个测试文件共 28 个单元/组件测试、5 个 Playwright E2E 和生产构建全部通过。 +- 已知非阻塞项:生产构建仍有约 972 KiB 主 JS chunk 警告,与原 MVP 记录的代码分割遗留项一致。 + +## Second Batch Result + +- [x] `ResizablePanel` / `PanelResizeHandle`:支持鼠标拖动、键盘调整、ARIA 数值和本地宽度持久化。 +- [x] `TreeSearchField`:工程树和模型结构树支持本地过滤,并保留匹配节点的祖先路径。 +- [x] `ViewportHUD`:复用现有仿真、交互模式和选择状态,在视口内提供轻量状态反馈。 +- [x] `ShortcutHelpDialog`:Header 帮助入口集中说明现有键盘与鼠标操作。 +- [x] `Badge`、`Tabs`、`Separator`、`Skeleton`:补齐基础层,并升级模型加载反馈。 +- [x] 侧栏标签页:左栏拆分“工程 / 模型结构”,右栏拆分“属性 / 控制”;非活动面板保持挂载,保留树和折叠状态。 +- [x] 第二批验证:TypeScript、ESLint、14 个测试文件共 36 个测试、生产构建及 5 个 Playwright E2E 全部通过;完成加载模型和快捷键弹窗的人工截图检查。 +- [x] 审查修复:Tabs roving tabindex/方向键导航、侧栏宽度边界与指针取消清理、Dialog 打开时快捷键隔离、搜索结果目录锁定展开,以及重编译失败快照清理和新 Session 速度恢复。 + +## Third Batch Result + +- [x] `ConfirmDialog`:统一替换移除工程的原生确认框。 +- [x] `CommandPalette`:支持 Header 入口、`Ctrl+K`、搜索、方向键和 Enter,复用现有仿真/视口/布局动作。 +- [x] `PerformancePopover`:从状态栏查看 FPS、物理步进、内存与预算状态。 +- [x] `Kbd`、`PropertyRow`、`CopyButton`:统一快捷键和属性展示,并支持复制已有选择信息。 +- [x] `TreeSearchSummary`、`SearchHighlight`、`EmptySearchState`:显示匹配数量、高亮命中并统一无结果反馈。 +- [x] `ViewportFullscreenToggle`:支持浏览器全屏进入、退出及状态同步。 +- [x] 第三批验证:TypeScript、ESLint、16 个测试文件共 42 个测试、生产构建及 5 个 Playwright E2E 全部通过;完成命令面板和完整工作台截图检查。 +- [x] 审查修复:全屏元素内 Dialog Portal、Popover 与命令面板互斥关闭、树搜索计数同源、Clipboard API 降级,以及命令面板 combobox/listbox 活动项关联。 + +## Fourth Batch Result + +- [x] `ToastViewport` / `NotificationCenter`:模型加载成功、兼容提示和编译失败进入会话通知队列,并提供即时 Toast。 +- [x] `LayoutSettingsDialog`:统一左右栏显示、宽度重置和默认/宽视口/工程浏览/控制调试预设。 +- [x] `ProjectBreadcrumb` / `EntrySwitcher`:展示工程入口路径,并可在多入口工程中直接切换。 +- [x] `VirtualTreeViewport`:大型工程文件树及 Body/Joint 树超过阈值时启用窗口化渲染。 +- [x] `SettingsDialog`:集中管理主题、角度单位、碰撞几何、关节高级信息和外力强度。 +- [x] Header 接入通知、布局和设置入口;1024px 产品边界下隐藏品牌长标题以避免工具区重叠。 +- [x] 第四批验证:TypeScript、ESLint、18 个测试文件共 49 个测试、生产构建及 5 个 Playwright E2E 全部通过。 + +- [x] 审查修复:入口切换加载锁与禁用态、1024px Header 轨道约束及 E2E 边界检查、虚拟树 roving active descendant/方向键导航和层级展开折叠。 + +## Fifth Batch Result + +- [x] `ToolbarOverflowMenu`:1024px 工具区通过“更多”菜单承载命令、布局、设置、全屏、帮助与主题动作。 +- [x] 通用 `Popover` / `DropdownMenu`:统一外部点击、Escape、焦点恢复及菜单方向键行为,并重构性能和通知弹层。 +- [x] `ImportProgressPanel` / `ProgressBar`:在不改变转换逻辑的前提下显示读取、资源处理、WASM/编译和视口创建阶段。 +- [x] `DiagnosticsDrawer` / EventLog:按全部、警告和错误分类查看、复制及清空会话事件。 +- [x] `SearchableCombobox`:多入口工程支持路径搜索和键盘选择。 +- [x] `ErrorRecoveryPanel`:诊断反馈提供重试入口、复制详情和返回工程树动作。 +- [x] `LiveRegion`:统一 Toast、加载阶段和诊断变化的辅助技术播报。 +- [x] 第五批验证:TypeScript、ESLint、20 个测试文件共 56 个测试和生产构建通过;Playwright E2E 完整回归覆盖更多菜单、事件日志与错误恢复。 + +- [x] 审查修复:拖放文件收集与导入共享互斥锁、仅模型编译错误允许重试、菜单首项聚焦及触发器恢复、Combobox 完整状态与 Tab 关闭、事件日志按活动标签惰性挂载、Toast 单一 live region。 +- [x] 信息归并:从“模型与控制”侧栏移除 `URDF 兼容处理` 分区,完整兼容处理明细仅保留在通知中心与事件日志。 + +## Confirmed Decisions + +1. 第一批 A–H 全部纳入。 +2. 只采用 URDF-Studio 的专业工作台风格,保留 MuJoCo 品牌、中文界面和当前左右栏布局。 +3. 允许引入 `lucide-react`,不引入 `@floating-ui/react`。 +4. select/joint/force 工具组停靠在 Header 中央。 +5. 右侧默认展开模型信息、当前选择和关节;其他低频分区折叠,警告出现时自动展开。 +6. 第一批不新增编辑、导出、快照、AI、测量、绘制、撤销/重做等产品能力。 +7. 后续任何超出 A–H 的新组件都需再次向用户确认。 diff --git a/requirements-dev.txt b/requirements-dev.txt index 68e63057..2e813688 100644 --- a/requirements-dev.txt +++ b/requirements-dev.txt @@ -1 +1,2 @@ +-r training_server/requirements.txt ruff==0.16.5 diff --git a/training_server/README.md b/training_server/README.md index 557c0922..09535196 100644 --- a/training_server/README.md +++ b/training_server/README.md @@ -1,6 +1,6 @@ # 本地强化学习训练服务 -该服务把 Web 平台发出的受限训练请求转换为本机训练子进程,并提供状态轮询、停止任务和 `policy.onnx` 下载接口。服务只绑定 `127.0.0.1`,不执行前端传入的任意命令。 +该服务把 Web 平台发出的受限训练请求转换为本机训练子进程,并提供状态轮询、停止任务和 `policy.onnx` 下载接口。它还提供 `Unitree-Go2-Flat` 奖励函数自调参:DeepSeek Agent 根据训练曲线与固定评估指标提出受限参数 patch,系统支持全自动或逐轮审批、successive-halving、TensorBoard scalar 查询、最佳 preset 与 ONNX 导出。服务只绑定 `127.0.0.1`,不执行前端传入的任意命令或 Agent 生成的代码。 仓库已在 [`rl/`](rl/) 内置 `Unitree-Go2-Flat` 所需的 PPO 训练代码、Go2 模型资产和 ONNX 导出逻辑,不再要求另外克隆 `unitree_rl_mjlab`。`mjlab`、PyTorch 等大型运行依赖仍需安装在本机训练环境中。 @@ -11,6 +11,7 @@ ```bash source .venv/bin/activate python -m pip install -r training_server/rl/requirements.txt +python -m pip install -r training_server/requirements.txt ``` ## 启动 @@ -29,6 +30,17 @@ MUJOCO_TRAINING_TOKEN='至少十六个字符的随机令牌' npm run training-se --trainer-python /path/to/training-env/bin/python ``` +启用云端自调参 Agent 时,在**服务端环境变量**中配置 DeepSeek;不要把 key 填到浏览器、URL 或命令行参数: + +```bash +export DEEPSEEK_API_KEY='你的 DeepSeek API key' +export MUJOCO_TUNING_AGENT_BASE_URL='https://api.deepseek.com' # 可省略 +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 候选,不会静默降级。 + 默认训练工程是仓库内的 `training_server/rl`。如需使用包含其他已注册任务的外部训练工程,仍可通过 `--trainer-root /path/to/trainer` 或 `UNITREE_RL_MJLAB_ROOT` 覆盖。默认端口是 `8765`。如果前端不是从 `localhost` 或 `127.0.0.1` 提供,可显式添加来源: ```bash @@ -37,7 +49,17 @@ python training_server/server.py \ --allow-origin http://192.168.1.10:5173 ``` -服务一次只运行一个训练任务,最多保留 20 个任务的内存状态,每个任务最多保留 200 行最近日志。停止服务或在界面点击“停止训练”会同步终止整个训练进程组。训练请求的 W&B 模式默认为 `offline`,保留本地指标但不登录;也可以在界面选择完全禁用或在线模式。所有 API 请求都必须携带启动时生成的 Bearer Token。 +普通训练与自调参共享同一个计算资源锁,任何时刻只允许一个训练/评估子进程占用 GPU。普通任务最多保留 20 个内存状态和每个任务 200 行最近日志;调参 session、trial、proposal、审计和 scalar 写入 SQLite/WAL,默认保存在 `training_server/rl/logs/auto_tuning/`。服务重启后等待审批/暂停状态可恢复,正在训练或评估的 trial 标记为 interrupted,只能从已完整保存的 checkpoint 显式恢复。停止服务或点击“停止”会终止整个进程组。 + +普通训练的 W&B 模式默认为 `offline`;调参 trial 强制使用本地 TensorBoard writer,DeepSeek 只接收最多 12 个 trial 的脱敏数值摘要和降采样曲线,不接收源代码、机器人资产、checkpoint、服务 token 或本地路径。所有 HTTP API 请求都必须携带启动时生成的 Bearer Token。 + +## 自调参流程 + +在主工作台连接训练服务后,点击“打开自调参 Agent 工作台”。默认预算为 12 个唯一配置:所有配置先训练 300 iterations,前 4 名续训到 900,前 2 名续训到 2000;默认使用 GPU 0 和 4096 个并行环境。首次使用建议先降低为 256–512 environments 做 smoke test。 + +固定评估使用站立、前进/侧移、转向和组合命令以及 3 个固定 seed。最终分数不直接使用可被权重放大的总 reward,而由速度跟踪 35%、动作平滑 20%、姿态稳定 15%、减少跌倒 15%、足端滑移 10%、能耗 5% 的权重无关指标组成。跌倒率高于基线 2% 或速度误差恶化超过 5% 的 trial 不晋级。逐轮审批模式会自动运行基线,之后每条 Agent 建议都等待批准、修改后批准或拒绝反馈。 + +最佳结果保存为不可变 preset,可在普通训练面板的“奖励配置”中选择,也可导出 JSON;不会覆盖仓库里的 Python 默认奖励配置。 ## 接口 @@ -45,14 +67,22 @@ python training_server/server.py \ - `POST /api/training/jobs`:发起训练; - `GET /api/training/jobs/{id}`:状态、迭代进度和最近日志; - `DELETE /api/training/jobs/{id}`:停止训练; -- `GET /api/training/jobs/{id}/artifacts/policy.onnx`:下载本次生成的策略。 +- `GET /api/training/jobs/{id}/artifacts/policy.onnx`:下载本次生成的策略; +- `GET /api/tuning/capabilities`、`POST /api/tuning/agent/test`:检查/测试 Agent; +- `GET|POST /api/tuning/sessions`、`GET|DELETE /api/tuning/sessions/{id}`:列出、创建、查询、停止 session; +- `POST /api/tuning/sessions/{id}/pause|resume`:暂停后续调度或恢复; +- `POST /api/tuning/sessions/{id}/proposals/{proposalId}/approve|reject`:审批、修改或拒绝建议; +- `GET /api/tuning/sessions/{id}/trials/{trialId}/metrics`:查询降采样 scalar; +- `GET /api/tuning/sessions/{id}/artifacts/best/policy.onnx`:下载最佳策略; +- `GET /api/tuning/presets`:列出可供普通训练复用的最佳奖励 preset。 -任务保存在服务内存中,服务重启后历史任务状态会丢失;训练日志、checkpoint 和 ONNX 产物保留在 `training_server/rl/logs/rsl_rl/`,使用外部训练工程时则保留在对应工程中。 +普通任务状态在服务重启后丢失,但日志、checkpoint 和 ONNX 保留在 `training_server/rl/logs/rsl_rl/`;调参状态及产物持久化在 `logs/auto_tuning/`。API 只接收 32 位资源 ID,不接收客户端文件路径;奖励 patch 受到名称、符号、上下界、每轮最多 4 项及 `0.5×–2×` 变化率校验。 ## 测试 ```bash -python3 -m pip install -r requirements-dev.txt +source .venv/bin/activate +python -m pip install -r requirements-dev.txt npm run lint:python npm run test:training-server ``` diff --git a/training_server/requirements.txt b/training_server/requirements.txt new file mode 100644 index 00000000..176e96a7 --- /dev/null +++ b/training_server/requirements.txt @@ -0,0 +1,5 @@ +# 本地自调参服务(Python 3.12) +pydantic-ai-slim[openai]==2.37.0 +httpx2[socks]==2.12.0 +optuna==4.9.0 +tensorboard==2.21.0 diff --git a/training_server/rl/scripts/evaluate.py b/training_server/rl/scripts/evaluate.py new file mode 100644 index 00000000..6c78bc0b --- /dev/null +++ b/training_server/rl/scripts/evaluate.py @@ -0,0 +1,222 @@ +"""Deterministic, headless evaluation for Unitree Go2 velocity policies.""" + +from __future__ import annotations + +import hashlib +import json +import os +import sys +from dataclasses import asdict, dataclass, field +from pathlib import Path +from statistics import fmean, pstdev +from typing import Literal + +TRAINER_ROOT = Path(__file__).resolve().parents[1] +SERVICE_ROOT = TRAINER_ROOT.parent +for source_root in (TRAINER_ROOT, SERVICE_ROOT): + if str(source_root) not in sys.path: + sys.path.insert(0, str(source_root)) + +import torch +import tyro +import warp as wp + +if not hasattr(wp, "context"): + from warp._src import context as warp_context + + wp.context = warp_context # type: ignore[attr-defined] + +from mjlab.envs import ManagerBasedRlEnv +from mjlab.rl import MjlabOnPolicyRunner, RslRlVecEnvWrapper +from mjlab.tasks.registry import list_tasks, load_env_cfg, load_rl_cfg, load_runner_cls +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 = ( + (0.0, 0.0, 0.0), + (0.5, 0.0, 0.0), + (1.0, 0.0, 0.0), + (1.5, 0.0, 0.0), + (0.0, 0.5, 0.0), + (0.0, -0.5, 0.0), + (0.0, 0.0, 0.5), + (0.0, 0.0, -0.5), + (0.8, 0.25, 0.35), +) +METRIC_NAMES = ( + "linear_velocity_rmse", + "angular_velocity_rmse", + "mean_action_acc", + "orientation_error", + "slip_velocity", + "mechanical_power", +) + + +@dataclass(frozen=True) +class EvaluateConfig: + checkpoint: str + output: str + reward_config: str | None = None + num_envs: int = 256 + steps_per_seed: int = 1000 + seeds: tuple[int, ...] = field(default_factory=lambda: (101, 202, 303)) + device: str | None = None + gpu_ids: list[int] | Literal["all"] | None = field(default_factory=lambda: [0]) + + +def _sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as stream: + for chunk in iter(lambda: stream.read(1024 * 1024), b""): + digest.update(chunk) + return digest.hexdigest() + + +def _fixed_commands(command: torch.Tensor) -> torch.Tensor: + values = torch.as_tensor(SCENARIOS, device=command.device, dtype=command.dtype) + indexes = torch.arange(command.shape[0], device=command.device) % values.shape[0] + command[:] = values[indexes] + return command + + +def _evaluate_seed(task_id: str, cfg: EvaluateConfig, seed: int) -> dict[str, float]: + torch.manual_seed(seed) + env_cfg = load_env_cfg(task_id, play=False) + agent_cfg = load_rl_cfg(task_id) + env_cfg.seed = seed + env_cfg.scene.num_envs = cfg.num_envs + env_cfg.curriculum = {} + env_cfg.observations["actor"].enable_corruption = False + env_cfg.events.pop("push_robot", None) + twist_cfg = env_cfg.commands["twist"] + assert isinstance(twist_cfg, UniformVelocityCommandCfg) + twist_cfg.heading_command = False + twist_cfg.ranges.heading = None + twist_cfg.rel_heading_envs = 0.0 + twist_cfg.rel_standing_envs = 0.0 + twist_cfg.resampling_time_range = (1.0e9, 1.0e9) + if cfg.reward_config: + reward_path = Path(cfg.reward_config).expanduser().resolve(strict=True) + with reward_path.open(encoding="utf-8") as stream: + apply_reward_configuration(env_cfg, validate_configuration(json.load(stream))) + + env = ManagerBasedRlEnv(cfg=env_cfg, device=cfg.device or "cuda:0") + 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) + runner.load( + str(Path(cfg.checkpoint).expanduser().resolve(strict=True)), + load_cfg={"actor": True}, + strict=True, + map_location=str(wrapped.device), + ) + policy = runner.get_inference_policy(device=str(wrapped.device)) + twist = wrapped.unwrapped.command_manager.get_term("twist") + _fixed_commands(twist.command) + obs = wrapped.get_observations() + + sums = {name: 0.0 for name in METRIC_NAMES} + samples = 0 + terminations = 0 + completions = 0 + with torch.inference_mode(): + for _ in range(cfg.steps_per_seed): + _fixed_commands(twist.command) + obs = wrapped.get_observations() + actions = policy(obs) + obs, _rewards, _dones, _extras = wrapped.step(actions) + manager = wrapped.unwrapped.metrics_manager + for index, name in enumerate(manager.active_terms): + if name in sums: + sums[name] += float(torch.sum(manager._step_values[:, index]).item()) + samples += wrapped.num_envs + terminated = wrapped.unwrapped.reset_terminated + timed_out = wrapped.unwrapped.reset_time_outs + terminations += int(torch.count_nonzero(terminated).item()) + completions += int(torch.count_nonzero(terminated | timed_out).item()) + result = {name: sums[name] / max(samples, 1) for name in METRIC_NAMES} + result["fall_rate"] = terminations / max(completions, wrapped.num_envs) + return result + finally: + wrapped.close() + + +def run_evaluation(task_id: str, cfg: EvaluateConfig) -> dict: + configure_torch_backends() + selected = cfg.gpu_ids + if selected is None: + os.environ["CUDA_VISIBLE_DEVICES"] = "" + device = "cpu" + else: + if selected == "all": + selected = [0] + os.environ["CUDA_VISIBLE_DEVICES"] = ",".join(map(str, selected)) + device = cfg.device or "cuda:0" + os.environ["MUJOCO_GL"] = "egl" + cfg = EvaluateConfig(**{**asdict(cfg), "device": device}) + + 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) + ], + } + 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: + for name, value in metrics.items(): + writer.add_scalar(f"Evaluation/{name}", value, 0) + finally: + writer.close() + print("MUJOCO_EVALUATION " + json.dumps({"output": str(output), "metrics": metrics})) + return result + + +def main() -> None: + import mjlab.tasks # noqa: F401 + import src.tasks # noqa: F401 + + chosen_task, remaining = tyro.cli( + tyro.extras.literal_type_from_choices(list_tasks()), + add_help=False, + return_unknown_args=True, + config=mjlab.TYRO_FLAGS, + ) + args = tyro.cli( + EvaluateConfig, + args=remaining, + prog=sys.argv[0] + f" {chosen_task}", + config=mjlab.TYRO_FLAGS, + ) + run_evaluation(chosen_task, args) + + +if __name__ == "__main__": + main() diff --git a/training_server/rl/scripts/train.py b/training_server/rl/scripts/train.py index 8f38282c..5fc7f1e1 100644 --- a/training_server/rl/scripts/train.py +++ b/training_server/rl/scripts/train.py @@ -1,5 +1,6 @@ """Script to train RL agent with RSL-RL.""" +import json import logging import os import sys @@ -10,8 +11,10 @@ from typing import Literal, cast # 训练器作为仓库内置子集直接从 scripts/ 启动,不要求额外执行 pip install -e。 TRAINER_ROOT = Path(__file__).resolve().parents[1] -if str(TRAINER_ROOT) not in sys.path: - sys.path.insert(0, str(TRAINER_ROOT)) +SERVICE_ROOT = TRAINER_ROOT.parent +for source_root in (TRAINER_ROOT, SERVICE_ROOT): + if str(source_root) not in sys.path: + sys.path.insert(0, str(source_root)) import tyro import warp as wp @@ -32,6 +35,8 @@ 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 tuning.schema import apply_reward_configuration, validate_configuration + @dataclass(frozen=True) class TrainConfig: @@ -44,6 +49,10 @@ class TrainConfig: enable_nan_guard: bool = False torchrunx_log_dir: str | None = None gpu_ids: list[int] | Literal["all"] | None = field(default_factory=lambda: [0]) + output_dir: str | None = None + resume_checkpoint: str | None = None + reward_config: str | None = None + reward_config_json: str | None = None @staticmethod def from_task(task_id: str) -> "TrainConfig": @@ -52,7 +61,27 @@ class TrainConfig: return TrainConfig(env=env_cfg, agent=agent_cfg) +def _load_reward_config(path: str | None, inline: str | None) -> 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)) + 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)) + + 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) + cuda_visible = os.environ.get("CUDA_VISIBLE_DEVICES", "") if cuda_visible == "": device = "cpu" @@ -109,11 +138,14 @@ def run_train(task_id: str, cfg: TrainConfig, log_dir: Path) -> None: log_root_path = log_dir.parent # Go up from specific run dir to experiment dir. resume_path: Path | None = None - if cfg.agent.resume: - # Load checkpoint from local filesystem. - resume_path = get_checkpoint_path( - log_root_path, cfg.agent.load_run, cfg.agent.load_checkpoint - ) + explicit_resume = cfg.resume_checkpoint is not None + if explicit_resume: + resume_path = Path(cfg.resume_checkpoint).expanduser().resolve(strict=True) + elif cfg.agent.resume: + # Load checkpoint from local filesystem. + resume_path = get_checkpoint_path( + log_root_path, cfg.agent.load_run, cfg.agent.load_checkpoint + ) # Only record videos on rank 0 to avoid multiple workers writing to the same files. if cfg.video and rank == 0: @@ -141,16 +173,32 @@ def run_train(task_id: str, cfg: TrainConfig, log_dir: Path) -> None: runner.add_git_repo_to_log(__file__) if resume_path is not None: print(f"[INFO]: Loading model checkpoint from: {resume_path}") - runner.load(str(resume_path)) + runner.load(str(resume_path), map_location=device) + 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. + runner.current_learning_iteration += 1 # Only write config files from rank 0 to avoid race conditions. if rank == 0: dump_yaml(log_dir / "params" / "env.yaml", env_cfg) dump_yaml(log_dir / "params" / "agent.yaml", agent_cfg) + if reward_config is not None: + reward_snapshot = log_dir / "params" / "reward_config.json" + reward_snapshot.parent.mkdir(parents=True, exist_ok=True) + reward_snapshot.write_text( + json.dumps(reward_config, ensure_ascii=False, indent=2) + "\n", encoding="utf-8" + ) - runner.learn( - num_learning_iterations=cfg.agent.max_iterations, init_at_random_ep_len=True + iterations = cfg.agent.max_iterations + if explicit_resume: + iterations = max(0, cfg.agent.max_iterations - runner.current_learning_iteration) + print( + f"[INFO] Learning target: current={runner.current_learning_iteration}, " + f"additional={iterations}, target={cfg.agent.max_iterations}", + flush=True, ) + runner.learn(num_learning_iterations=iterations, init_at_random_ep_len=True) env.close() @@ -159,12 +207,15 @@ def launch_training(task_id: str, args: TrainConfig | None = None): args = args or TrainConfig.from_task(task_id) # Create log directory once before launching workers. - log_root_path = Path("logs") / "rsl_rl" / args.agent.experiment_name - log_root_path.resolve() - log_dir_name = datetime.now().strftime("%Y-%m-%d_%H-%M-%S") - if args.agent.run_name: - log_dir_name += f"_{args.agent.run_name}" - log_dir = log_root_path / log_dir_name + if args.output_dir: + log_dir = Path(args.output_dir).expanduser().resolve() + else: + log_root_path = (Path("logs") / "rsl_rl" / args.agent.experiment_name).resolve() + log_dir_name = datetime.now().strftime("%Y-%m-%d_%H-%M-%S") + if args.agent.run_name: + log_dir_name += f"_{args.agent.run_name}" + log_dir = log_root_path / log_dir_name + log_dir.mkdir(parents=True, exist_ok=True) # Select GPUs based on CUDA_VISIBLE_DEVICES and user specification. selected_gpus, num_gpus = select_gpus(args.gpu_ids) diff --git a/training_server/rl/src/tasks/velocity/config/go2/env_cfgs.py b/training_server/rl/src/tasks/velocity/config/go2/env_cfgs.py index 244d72da..e84ff8f9 100644 --- a/training_server/rl/src/tasks/velocity/config/go2/env_cfgs.py +++ b/training_server/rl/src/tasks/velocity/config/go2/env_cfgs.py @@ -106,6 +106,7 @@ def unitree_go2_rough_env_cfg( cfg.rewards["body_ang_vel"].params["asset_cfg"].body_names = ("base_link",) cfg.rewards["foot_clearance"].params["asset_cfg"].site_names = site_names cfg.rewards["foot_slip"].params["asset_cfg"].site_names = site_names + cfg.metrics["slip_velocity"].params["asset_cfg"].site_names = site_names cfg.terminations["illegal_contact"] = TerminationTermCfg( func=mdp.illegal_contact, diff --git a/training_server/rl/src/tasks/velocity/mdp/__init__.py b/training_server/rl/src/tasks/velocity/mdp/__init__.py index 1d3be528..0c83b1ad 100644 --- a/training_server/rl/src/tasks/velocity/mdp/__init__.py +++ b/training_server/rl/src/tasks/velocity/mdp/__init__.py @@ -1,6 +1,7 @@ from mjlab.envs.mdp import * # noqa: F401, F403 from .curriculums import * # noqa: F403 +from .metrics import * # noqa: F403 from .observations import * # noqa: F403 from .rewards import * # noqa: F403 from .terminations import * # noqa: F403 diff --git a/training_server/rl/src/tasks/velocity/mdp/metrics.py b/training_server/rl/src/tasks/velocity/mdp/metrics.py new file mode 100644 index 00000000..937e5e73 --- /dev/null +++ b/training_server/rl/src/tasks/velocity/mdp/metrics.py @@ -0,0 +1,88 @@ +"""Reward-weight-independent quality metrics for velocity tasks.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +import torch + +from mjlab.entity import Entity +from mjlab.managers.scene_entity_config import SceneEntityCfg +from mjlab.sensor import ContactSensor + +if TYPE_CHECKING: + from mjlab.envs import ManagerBasedRlEnv + +_DEFAULT_ASSET_CFG = SceneEntityCfg("robot") + + +def linear_velocity_rmse( + env: ManagerBasedRlEnv, + command_name: str, + asset_cfg: SceneEntityCfg = _DEFAULT_ASSET_CFG, +) -> torch.Tensor: + """Per-step commanded-vs-actual base linear velocity RMSE in body frame.""" + asset: Entity = env.scene[asset_cfg.name] + command = env.command_manager.get_command(command_name) + assert command is not None + actual = asset.data.root_link_lin_vel_b + error = torch.cat((command[:, :2] - actual[:, :2], -actual[:, 2:3]), dim=1) + return torch.sqrt(torch.mean(torch.square(error), dim=1)) + + +def angular_velocity_rmse( + env: ManagerBasedRlEnv, + command_name: str, + asset_cfg: SceneEntityCfg = _DEFAULT_ASSET_CFG, +) -> torch.Tensor: + """Per-step commanded-vs-actual base angular velocity RMSE in body frame.""" + asset: Entity = env.scene[asset_cfg.name] + command = env.command_manager.get_command(command_name) + assert command is not None + actual = asset.data.root_link_ang_vel_b + desired = torch.zeros_like(actual) + desired[:, 2] = command[:, 2] + return torch.sqrt(torch.mean(torch.square(desired - actual), dim=1)) + + +def orientation_error( + env: ManagerBasedRlEnv, + asset_cfg: SceneEntityCfg = _DEFAULT_ASSET_CFG, +) -> torch.Tensor: + """Magnitude of projected gravity in the base x/y plane; zero is upright.""" + asset: Entity = env.scene[asset_cfg.name] + return torch.linalg.vector_norm(asset.data.projected_gravity_b[:, :2], dim=1) + + +def fall_indicator(env: ManagerBasedRlEnv) -> torch.Tensor: + """One on non-timeout terminal steps, otherwise zero.""" + return env.termination_manager.terminated.float() + + +def slip_velocity( + env: ManagerBasedRlEnv, + sensor_name: str, + asset_cfg: SceneEntityCfg = _DEFAULT_ASSET_CFG, +) -> torch.Tensor: + """Mean x/y velocity of feet currently touching the ground.""" + asset: Entity = env.scene[asset_cfg.name] + sensor: ContactSensor = env.scene[sensor_name] + assert sensor.data.found is not None + contact = (sensor.data.found > 0).float() + speed = torch.linalg.vector_norm( + asset.data.site_lin_vel_w[:, asset_cfg.site_ids, :2], dim=-1 + ) + count = torch.clamp(torch.sum(contact, dim=1), min=1.0) + return torch.sum(speed * contact, dim=1) / count + + +def mechanical_power( + env: ManagerBasedRlEnv, + asset_cfg: SceneEntityCfg = _DEFAULT_ASSET_CFG, +) -> torch.Tensor: + """Positive actuator mechanical power in watts (regeneration is ignored).""" + asset: Entity = env.scene[asset_cfg.name] + torque = asset.data.actuator_force[:, asset_cfg.actuator_ids] + velocity = asset.data.joint_vel[:, asset_cfg.joint_ids] + count = min(torque.shape[1], velocity.shape[1]) + return torch.sum(torch.clamp(torque[:, :count] * velocity[:, :count], min=0.0), dim=1) diff --git a/training_server/rl/src/tasks/velocity/velocity_env_cfg.py b/training_server/rl/src/tasks/velocity/velocity_env_cfg.py index feebe378..1cc70db0 100644 --- a/training_server/rl/src/tasks/velocity/velocity_env_cfg.py +++ b/training_server/rl/src/tasks/velocity/velocity_env_cfg.py @@ -140,8 +140,27 @@ def make_velocity_env_cfg() -> ManagerBasedRlEnvCfg: ## metrics = { - "mean_action_acc": MetricsTermCfg( - func=mdp.mean_action_acc, + "linear_velocity_rmse": MetricsTermCfg( + func=mdp.linear_velocity_rmse, + params={"command_name": "twist"}, + ), + "angular_velocity_rmse": MetricsTermCfg( + func=mdp.angular_velocity_rmse, + params={"command_name": "twist"}, + ), + "mean_action_acc": MetricsTermCfg(func=mdp.mean_action_acc), + "orientation_error": MetricsTermCfg(func=mdp.orientation_error), + "fall_indicator": MetricsTermCfg(func=mdp.fall_indicator), + "slip_velocity": MetricsTermCfg( + func=mdp.slip_velocity, + params={ + "sensor_name": "feet_ground_contact", + "asset_cfg": SceneEntityCfg("robot", site_names=()), # Set per-robot. + }, + ), + "mechanical_power": MetricsTermCfg( + func=mdp.mechanical_power, + params={"asset_cfg": SceneEntityCfg("robot", joint_names=(".*",))}, ), } @@ -298,6 +317,11 @@ def make_velocity_env_cfg() -> ManagerBasedRlEnvCfg: params={"sensor_name": "robot/root_angmom"}, ), "is_terminated": RewardTermCfg(func=mdp.is_terminated, weight=-200.0), + "electrical_power": RewardTermCfg( + func=mdp.electrical_power_cost, + weight=0.0, + params={"asset_cfg": SceneEntityCfg("robot", joint_names=(".*",))}, + ), "joint_acc_l2": RewardTermCfg(func=mdp.joint_acc_l2, weight=-2.5e-7), "joint_pos_limits": RewardTermCfg(func=mdp.joint_pos_limits, weight=-10.0), "action_rate_l2": RewardTermCfg(func=mdp.action_rate_l2, weight=-0.05), diff --git a/training_server/server.py b/training_server/server.py index e96e132d..12a78fe6 100644 --- a/training_server/server.py +++ b/training_server/server.py @@ -23,9 +23,12 @@ from http import HTTPStatus from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from pathlib import Path from typing import Any -from urllib.parse import unquote, urlsplit +from urllib.parse import parse_qs, unquote, urlsplit -VERSION = "0.3.0" +from tuning.manager import TuningError, TuningManager +from tuning.process import GpuLease, ResourceBusyError + +VERSION = "0.4.0" # 浏览器当前 ONNX 运行时只实现 Go2 的 47→12 部署契约;其他任务须由服务启动参数显式放行。 DEFAULT_TASKS = ("Unitree-Go2-Flat",) ACTIVE_STATES = {"queued", "running"} @@ -65,6 +68,7 @@ class TrainingConfig: device: str gpu_ids: list[int] wandb_mode: str + reward_config: dict[str, Any] | None = None @dataclass @@ -110,6 +114,7 @@ class TrainingManager: python: str, tasks: tuple[str, ...], check_environment: bool = True, + lease: GpuLease | None = None, ): self.trainer_root = trainer_root.expanduser().resolve() self.python = str(Path(python).expanduser()) if os.sep in python else python @@ -118,6 +123,8 @@ class TrainingManager: self.lock = threading.RLock() self.check_environment = check_environment self._environment_error: str | None | bool = False + self.lease = lease or GpuLease() + self.preset_resolver: Any = None def readiness_error(self) -> str | None: if not self.trainer_root.is_dir(): @@ -204,6 +211,17 @@ class TrainingManager: wandb_mode = payload.get("wandbMode", "offline") if wandb_mode not in ("offline", "disabled", "online"): raise ApiError(HTTPStatus.BAD_REQUEST, "wandbMode 必须是 offline、disabled 或 online") + preset_id = payload.get("rewardPresetId") + reward_config = None + if preset_id is not None: + 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) + except KeyError as error: + raise ApiError(HTTPStatus.BAD_REQUEST, "奖励 preset 不存在") from error return TrainingConfig( task_id=task_id, num_envs=integer("numEnvs", 1, 16384), @@ -213,6 +231,7 @@ class TrainingManager: device=device, gpu_ids=raw_gpu_ids, wandb_mode=wandb_mode, + reward_config=reward_config, ) def start(self, payload: Any) -> dict[str, Any]: @@ -232,10 +251,20 @@ class TrainingManager: raise ApiError(HTTPStatus.CONFLICT, "训练任务历史已满,请稍后重试") del self.jobs[completed] job = TrainingJob(id=uuid.uuid4().hex, config=config) + owner = f"training:{job.id}" + try: + self.lease.acquire(owner) + except ResourceBusyError as error: + raise ApiError(HTTPStatus.CONFLICT, str(error)) from error self.jobs[job.id] = job - threading.Thread( - target=self._run, args=(job,), name=f"training-{job.id[:8]}", daemon=True - ).start() + try: + threading.Thread( + target=self._run, args=(job,), name=f"training-{job.id[:8]}", daemon=True + ).start() + except Exception: + self.jobs.pop(job.id, None) + self.lease.release(owner) + raise return job.public() def get(self, job_id: str) -> dict[str, Any]: @@ -305,6 +334,13 @@ class TrainingManager: f"--agent.seed={config.seed}", f"--agent.run-name={config.run_name}", ] + if config.reward_config is not None: + command.extend( + ( + "--reward-config-json", + json.dumps(config.reward_config, ensure_ascii=False, separators=(",", ":")), + ) + ) if config.device == "cpu": command.extend(("--gpu-ids", "None")) else: @@ -408,13 +444,16 @@ class TrainingManager: job.state = "cancelled" if job.cancel_requested else "failed" job.message = f"启动训练失败:{error}" job.logs.append(job.message) + finally: + self.lease.release(f"training:{job.id}") class TrainingRequestHandler(BaseHTTPRequestHandler): manager: TrainingManager + tuning_manager: TuningManager allowed_origins: tuple[str, ...] = () access_token = "" - server_version = "MuJoCoLocalTraining/0.3" + server_version = "MuJoCoLocalTraining/0.4" def log_message(self, format: str, *args: Any) -> None: sys.stderr.write(f"[{self.log_date_time_string()}] {format % args}\n") @@ -455,6 +494,12 @@ class TrainingRequestHandler(BaseHTTPRequestHandler): def _error(self, error: Exception) -> None: if isinstance(error, ApiError): self._json(error.status, {"error": str(error)}) + elif isinstance(error, KeyError): + self._json(HTTPStatus.NOT_FOUND, {"error": "调参 session、trial 或 proposal 不存在"}) + elif isinstance(error, ResourceBusyError): + self._json(HTTPStatus.CONFLICT, {"error": str(error)}) + elif isinstance(error, TuningError): + self._json(HTTPStatus.BAD_REQUEST, {"error": str(error)}) else: self._json( HTTPStatus.INTERNAL_SERVER_ERROR, {"error": f"本地训练服务内部错误:{error}"} @@ -490,6 +535,18 @@ class TrainingRequestHandler(BaseHTTPRequestHandler): match = re.fullmatch(r"/api/training/jobs/([0-9a-f]{32})(/artifacts/policy\.onnx)?", path) return (unquote(match.group(1)), bool(match.group(2))) if match else (None, False) + def _send_file(self, file_path: Path, filename: str) -> None: + size = file_path.stat().st_size + self.send_response(HTTPStatus.OK) + self._cors() + self.send_header("Content-Type", "application/octet-stream") + self.send_header("Content-Disposition", f'attachment; filename="{filename}"') + self.send_header("Content-Length", str(size)) + self.send_header("Cache-Control", "no-store") + self.end_headers() + with file_path.open("rb") as source: + shutil.copyfileobj(source, self.wfile) + def do_OPTIONS(self) -> None: try: self._ensure_origin() @@ -505,25 +562,57 @@ class TrainingRequestHandler(BaseHTTPRequestHandler): def do_GET(self) -> None: try: self._ensure_request() - path = urlsplit(self.path).path + parsed = urlsplit(self.path) + path = parsed.path if path == "/api/training/health": - self._json(HTTPStatus.OK, self.manager.health()) + health = self.manager.health() + health["tuning"] = self.tuning_manager.capability() + health["resourceOwner"] = self.manager.lease.public() + self._json(HTTPStatus.OK, health) + return + if path == "/api/tuning/capabilities": + self._json(HTTPStatus.OK, self.tuning_manager.capability()) + return + if path == "/api/tuning/sessions": + self._json(HTTPStatus.OK, {"sessions": self.tuning_manager.list()}) + return + if path == "/api/tuning/presets": + self._json(HTTPStatus.OK, {"presets": self.tuning_manager.storage.list_presets()}) + return + match = re.fullmatch(r"/api/tuning/sessions/([0-9a-f]{32})", path) + if match: + self._json(HTTPStatus.OK, self.tuning_manager.detail(match.group(1))) + return + match = re.fullmatch( + r"/api/tuning/sessions/([0-9a-f]{32})/trials/([0-9a-f]{32})/metrics", path + ) + if match: + query = parse_qs(parsed.query) + tags = [tag for value in query.get("tags", []) for tag in value.split(",") if tag] + try: + max_points = int(query.get("maxPoints", ["1000"])[0]) + except ValueError as error: + raise TuningError("maxPoints 必须是整数") from error + if not 10 <= max_points <= 5000: + raise TuningError("maxPoints 必须在 10–5000 之间") + self._json( + HTTPStatus.OK, + self.tuning_manager.metrics( + match.group(1), match.group(2), tags or None, max_points + ), + ) + return + match = re.fullmatch( + r"/api/tuning/sessions/([0-9a-f]{32})/artifacts/best/policy\.onnx", path + ) + if match: + self._send_file(self.tuning_manager.best_artifact(match.group(1)), "policy.onnx") return job_id, artifact = self._route(path) if not job_id: raise ApiError(HTTPStatus.NOT_FOUND, "接口不存在") if artifact: - file_path = self.manager.artifact(job_id) - size = file_path.stat().st_size - self.send_response(HTTPStatus.OK) - self._cors() - self.send_header("Content-Type", "application/octet-stream") - self.send_header("Content-Disposition", 'attachment; filename="policy.onnx"') - self.send_header("Content-Length", str(size)) - self.send_header("Cache-Control", "no-store") - self.end_headers() - with file_path.open("rb") as source: - shutil.copyfileobj(source, self.wfile) + self._send_file(self.manager.artifact(job_id), "policy.onnx") else: self._json(HTTPStatus.OK, self.manager.get(job_id)) except Exception as error: @@ -532,16 +621,52 @@ class TrainingRequestHandler(BaseHTTPRequestHandler): def do_POST(self) -> None: try: self._ensure_request() - if urlsplit(self.path).path != "/api/training/jobs": - raise ApiError(HTTPStatus.NOT_FOUND, "接口不存在") - self._json(HTTPStatus.ACCEPTED, self.manager.start(self._payload())) + path = urlsplit(self.path).path + if path == "/api/training/jobs": + self._json(HTTPStatus.ACCEPTED, self.manager.start(self._payload())) + return + if path == "/api/tuning/agent/test": + self._json(HTTPStatus.OK, self.tuning_manager.test_agent()) + return + if path == "/api/tuning/sessions": + self._json(HTTPStatus.ACCEPTED, self.tuning_manager.create(self._payload())) + return + match = re.fullmatch(r"/api/tuning/sessions/([0-9a-f]{32})/(pause|resume)", path) + if match: + action = ( + self.tuning_manager.pause + if match.group(2) == "pause" + else self.tuning_manager.resume + ) + self._json(HTTPStatus.ACCEPTED, action(match.group(1))) + return + match = re.fullmatch( + r"/api/tuning/sessions/([0-9a-f]{32})/proposals/([0-9a-f]{32})/(approve|reject)", + path, + ) + if match: + action = ( + self.tuning_manager.approve + if match.group(3) == "approve" + else self.tuning_manager.reject + ) + self._json( + HTTPStatus.ACCEPTED, action(match.group(1), match.group(2), self._payload()) + ) + return + raise ApiError(HTTPStatus.NOT_FOUND, "接口不存在") except Exception as error: self._error(error) def do_DELETE(self) -> None: try: self._ensure_request() - job_id, artifact = self._route(urlsplit(self.path).path) + path = urlsplit(self.path).path + match = re.fullmatch(r"/api/tuning/sessions/([0-9a-f]{32})", path) + if match: + self._json(HTTPStatus.ACCEPTED, self.tuning_manager.cancel(match.group(1))) + return + job_id, artifact = self._route(path) if not job_id or artifact: raise ApiError(HTTPStatus.NOT_FOUND, "接口不存在") self._json(HTTPStatus.ACCEPTED, self.manager.cancel(job_id)) @@ -574,6 +699,12 @@ def parse_args() -> argparse.Namespace: parser.add_argument( "--trainer-python", default=sys.executable, help="已安装 mjlab/torch 的 Python 解释器" ) + parser.add_argument( + "--tuning-data-root", + type=Path, + default=None, + help="调参 SQLite 与 trial 产物目录;默认位于训练工程 logs/auto_tuning", + ) parser.add_argument( "--task", action="append", dest="tasks", help="允许前端启动的任务 ID;可重复" ) @@ -593,10 +724,18 @@ def main() -> None: token = args.token or secrets.token_urlsafe(24) if len(token) < 16: raise SystemExit("训练服务访问令牌至少需要 16 个字符") + lease = GpuLease() manager = TrainingManager( - args.trainer_root, args.trainer_python, tuple(args.tasks or DEFAULT_TASKS) + args.trainer_root, + args.trainer_python, + tuple(args.tasks or DEFAULT_TASKS), + lease=lease, ) + 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 TrainingRequestHandler.allowed_origins = tuple(args.allow_origin) TrainingRequestHandler.access_token = token server = ThreadingHTTPServer((args.host, args.port), TrainingRequestHandler) @@ -612,6 +751,7 @@ def main() -> None: except KeyboardInterrupt: print("\n正在停止本地训练服务…") finally: + tuning_manager.shutdown() manager.shutdown() server.server_close() diff --git a/training_server/tests/test_server.py b/training_server/tests/test_server.py index 46d2ab35..b0d7f8cf 100644 --- a/training_server/tests/test_server.py +++ b/training_server/tests/test_server.py @@ -1,3 +1,4 @@ +import json import subprocess import sys import tempfile @@ -64,13 +65,7 @@ out.write_bytes(b'onnx') self.assertTrue((trainer_root / "scripts" / "train.py").is_file()) self.assertTrue( ( - trainer_root - / "src" - / "assets" - / "robots" - / "unitree_go2" - / "xmls" - / "go2.xml" + trainer_root / "src" / "assets" / "robots" / "unitree_go2" / "xmls" / "go2.xml" ).is_file() ) @@ -92,6 +87,15 @@ out.write_bytes(b'onnx') ) self.assertEqual(command[-2:], ["--gpu-ids", "[0,2]"]) + 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 + config = self.manager.parse_config(self.payload(rewardPresetId=preset_id)) + command = self.manager.command_for(config) + index = command.index("--reward-config-json") + self.assertEqual(json.loads(command[index + 1]), reward_config) + def test_requires_local_host_origin_and_bearer_token(self): handler = object.__new__(TrainingRequestHandler) handler.access_token = "secret-token-1234" diff --git a/training_server/tests/test_tuning.py b/training_server/tests/test_tuning.py new file mode 100644 index 00000000..88d06a31 --- /dev/null +++ b/training_server/tests/test_tuning.py @@ -0,0 +1,195 @@ +import math +import sys +import tempfile +import unittest +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from tuning.advisor import AdvisorConfig, DeepSeekAdvisor # noqa: E402 +from tuning.schema import ( # noqa: E402 + BASE_REWARD_CONFIGURATION, + RewardConfigError, + merge_proposal, + validate_configuration, + validate_proposal, +) +from tuning.scoring import ( # noqa: E402 + DEFAULT_OBJECTIVE_WEIGHTS, + EvaluationError, + score_evaluation, +) +from tuning.storage import TuningStorage # noqa: E402 + + +class RewardSchemaTest(unittest.TestCase): + def test_baseline_is_complete_and_energy_is_disabled(self): + config = validate_configuration(BASE_REWARD_CONFIGURATION) + self.assertEqual(len(config["weights"]), 16) + self.assertEqual(config["weights"]["electrical_power"], 0.0) + + def test_sparse_proposal_constraints(self): + patch = validate_proposal( + { + "weights": {"track_linear_velocity": 1.2, "foot_slip": -0.3}, + "params": {"foot_gait.period": 0.65}, + }, + BASE_REWARD_CONFIGURATION, + ) + merged = merge_proposal(BASE_REWARD_CONFIGURATION, patch) + self.assertEqual(merged["weights"]["track_linear_velocity"], 1.2) + with self.assertRaises(RewardConfigError): + validate_proposal({"weights": {"track_linear_velocity": 0}}, BASE_REWARD_CONFIGURATION) + with self.assertRaises(RewardConfigError): + validate_proposal({"weights": {"foot_slip": 0.2}}, BASE_REWARD_CONFIGURATION) + with self.assertRaises(RewardConfigError): + validate_proposal( + {"weights": {"track_linear_velocity": 3.0}}, BASE_REWARD_CONFIGURATION + ) + with self.assertRaises(RewardConfigError): + validate_proposal({"params": {"foot_gait.period": math.nan}}, BASE_REWARD_CONFIGURATION) + with self.assertRaises(RewardConfigError): + validate_proposal( + { + "weights": { + "pose": 1.1, + "foot_gait": 0.6, + "foot_slip": -0.3, + "soft_landing": -0.002, + }, + "params": {"foot_gait.period": 0.65}, + }, + BASE_REWARD_CONFIGURATION, + ) + + def test_cross_parameter_order(self): + with self.assertRaises(RewardConfigError): + validate_proposal( + {"params": {"pose.walking_threshold": 0.5, "pose.running_threshold": 0.4}}, + BASE_REWARD_CONFIGURATION, + ) + + +class ScoringTest(unittest.TestCase): + baseline = { + "linear_velocity_rmse": 0.3, + "angular_velocity_rmse": 0.2, + "mean_action_acc": 0.1, + "orientation_error": 0.2, + "fall_rate": 0.1, + "slip_velocity": 0.2, + "mechanical_power": 100.0, + } + + def test_improvement_and_safety_gate(self): + better = {key: value * 0.8 for key, value in self.baseline.items()} + scored = score_evaluation(self.baseline, better, DEFAULT_OBJECTIVE_WEIGHTS) + self.assertTrue(scored["eligible"]) + self.assertGreater(scored["score"], 0) + unsafe = dict(better, fall_rate=0.2) + scored = score_evaluation(self.baseline, unsafe) + self.assertFalse(scored["eligible"]) + self.assertEqual(scored["score"], -1.0) + + def test_rejects_missing_and_nonfinite_metrics(self): + with self.assertRaises(EvaluationError): + score_evaluation(self.baseline, {"fall_rate": 0.1}) + bad = dict(self.baseline, mechanical_power=math.inf) + with self.assertRaises(EvaluationError): + score_evaluation(self.baseline, bad) + + +class AdvisorTest(unittest.TestCase): + class Output: + weights = {"pose": 1.1} + params = {} + rationale = "improve posture" + expected_impact = {"posture": "better"} + confidence = 0.7 + + class Result: + output = None + + @staticmethod + def usage(): + return type("Usage", (), {"requests": 1, "input_tokens": 10, "output_tokens": 5})() + + class Agent: + def __init__(self, output): + self.output = output + + def run_sync(self, _prompt): + result = AdvisorTest.Result() + result.output = self.output + return result + + def test_structured_result_is_revalidated_locally(self): + advisor = DeepSeekAdvisor(AdvisorConfig("fake")) + advisor._cached_agent = self.Agent(self.Output()) + proposal = advisor.propose({"trials": []}, BASE_REWARD_CONFIGURATION) + self.assertEqual(proposal["patch"]["weights"]["pose"], 1.1) + self.assertEqual(proposal["usage"]["input_tokens"], 10) + + def test_invalid_model_patch_is_rejected(self): + output = self.Output() + output.weights = {"track_linear_velocity": -1.0} + advisor = DeepSeekAdvisor(AdvisorConfig("fake")) + advisor._cached_agent = self.Agent(output) + with self.assertRaises(RewardConfigError): + advisor.propose({"trials": []}, BASE_REWARD_CONFIGURATION) + + +class StorageTest(unittest.TestCase): + def setUp(self): + self.temporary = tempfile.TemporaryDirectory() + self.storage = TuningStorage(Path(self.temporary.name) / "state.sqlite3") + + def tearDown(self): + self.temporary.cleanup() + + def test_persists_session_trial_proposal_and_metrics(self): + session = self.storage.create_session( + "approval", {"taskId": "Unitree-Go2-Flat"}, DEFAULT_OBJECTIVE_WEIGHTS, False + ) + trial = self.storage.create_trial( + session["id"], 0, 0, 10, BASE_REWARD_CONFIGURATION, None, "trial-000-rung-0" + ) + proposal = self.storage.create_proposal( + session["id"], + trial["id"], + {"weights": {"pose": 1.1}, "params": {}}, + "test", + {}, + 0.8, + ) + self.assertTrue(self.storage.decide_proposal(proposal["id"], "approved", None)) + self.assertFalse(self.storage.decide_proposal(proposal["id"], "approved", None)) + points = [ + ("Train/reward", step, float(step), 50.0 if step == 50 else float(step % 7)) + for step in range(100) + ] + self.storage.insert_metrics(trial["id"], points) + sampled = self.storage.metrics(trial["id"], max_points=10)[0]["points"] + self.assertEqual(len(sampled), 10) + self.assertEqual(sampled[0]["step"], 0) + self.assertEqual(sampled[-1]["step"], 99) + self.assertIn(50.0, [point["value"] for point in sampled]) + reopened = TuningStorage(self.storage.path) + self.assertEqual(reopened.get_session(session["id"])["mode"], "approval") + + def test_recovery_marks_inflight_records(self): + session = self.storage.create_session( + "automatic", {"taskId": "Unitree-Go2-Flat"}, DEFAULT_OBJECTIVE_WEIGHTS, True + ) + trial = self.storage.create_trial( + session["id"], 0, 0, 10, BASE_REWARD_CONFIGURATION, None, "trial-000-rung-0" + ) + self.storage.update_session(session["id"], state="running") + self.storage.update_trial(trial["id"], state="training") + self.storage.recover_interrupted() + self.assertEqual(self.storage.get_session(session["id"])["state"], "interrupted") + self.assertEqual(self.storage.get_trial(trial["id"])["state"], "interrupted") + + +if __name__ == "__main__": + unittest.main() diff --git a/training_server/tests/test_tuning_manager.py b/training_server/tests/test_tuning_manager.py new file mode 100644 index 00000000..61b66e37 --- /dev/null +++ b/training_server/tests/test_tuning_manager.py @@ -0,0 +1,195 @@ +import sys +import tempfile +import time +import unittest +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from tuning.manager import TuningManager # noqa: E402 +from tuning.process import GpuLease # noqa: E402 +from tuning.schema import BASE_REWARD_CONFIGURATION # noqa: E402 +from tuning.scoring import score_evaluation # noqa: E402 +from tuning.storage import now_iso # noqa: E402 + +BASE_METRICS = { + "linear_velocity_rmse": 0.3, + "angular_velocity_rmse": 0.2, + "mean_action_acc": 0.1, + "orientation_error": 0.2, + "fall_rate": 0.1, + "slip_velocity": 0.2, + "mechanical_power": 100.0, +} + + +class FakeAdvisor: + def capability(self): + return { + "configured": True, + "apiKeyConfigured": True, + "frameworkInstalled": True, + "model": "fake", + "baseUrl": "https://example.invalid", + } + + def propose(self, _context, previous): + value = min(2.4, previous["weights"]["pose"] * 1.05) + return { + "patch": {"weights": {"pose": value}, "params": {}}, + "rationale": "fake", + "expectedImpact": {}, + "confidence": 0.8, + "promptHash": "abc", + "usage": {}, + "model": "fake", + } + + def test_connection(self): + return {"ok": True, "model": "fake", "outputType": "fake"} + + +class FakeTuningManager(TuningManager): + def _execute_trial(self, session, trial, resume_checkpoint=None): + del resume_checkpoint + factor = max(0.5, 1.0 - 0.03 * trial["number"] - 0.01 * trial["rung"]) + metrics = {key: value * factor for key, value in BASE_METRICS.items()} + trials = self.storage.list_trials(session["id"]) + baseline = next( + (item for item in trials if item["number"] == 0 and item["rung"] == 0), None + ) + if baseline and baseline["evaluation"]: + scored = score_evaluation( + baseline["evaluation"]["metrics"], metrics, session["objectiveWeights"] + ) + else: + scored = {"score": 0.0, "eligible": True, "components": {}} + root = self._session_root(session["id"]) + run = root / trial["runDir"] + run.mkdir(parents=True, exist_ok=True) + (run / "model_1.pt").write_bytes(b"checkpoint") + (run / "policy.onnx").write_bytes(b"onnx") + evaluation = {"metrics": metrics, "score": scored} + self.storage.update_trial( + trial["id"], + state="completed", + started_at=now_iso(), + ended_at=now_iso(), + message="fake complete", + checkpoint_path=str((run / "model_1.pt").relative_to(root)), + policy_path=str((run / "policy.onnx").relative_to(root)), + evaluation=evaluation, + score=scored["score"], + eligible=scored["eligible"], + ) + return self.storage.get_trial(trial["id"]) + + +class TuningManagerTest(unittest.TestCase): + def setUp(self): + self.temporary = tempfile.TemporaryDirectory() + self.root = Path(self.temporary.name) + (self.root / "trainer" / "scripts").mkdir(parents=True) + (self.root / "trainer" / "scripts" / "evaluate.py").write_text("", encoding="utf-8") + self.manager = FakeTuningManager( + self.root / "trainer", + sys.executable, + self.root / "data", + GpuLease(), + advisor=FakeAdvisor(), + ) + + def tearDown(self): + self.manager.shutdown() + self.temporary.cleanup() + + @staticmethod + def payload(mode="automatic"): + return { + "taskId": "Unitree-Go2-Flat", + "mode": mode, + "runName": "test", + "numEnvs": 16, + "gpuIds": [0], + "trialCount": 4, + "initialIterations": 1, + "middleIterations": 2, + "finalIterations": 3, + "evalNumEnvs": 8, + "evalSteps": 10, + } + + def wait_terminal(self, session_id, timeout=5): + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + session = self.manager.detail(session_id) + if session["state"] in {"succeeded", "failed", "cancelled"}: + return session + time.sleep(0.01) + self.fail("session did not finish") + + def test_automatic_session_runs_rungs_and_persists_best_artifact(self): + session = self.manager.create(self.payload()) + completed = self.wait_terminal(session["id"]) + self.assertEqual(completed["state"], "succeeded", completed["message"]) + self.assertGreaterEqual(len(completed["trials"]), 7) + self.assertTrue(self.manager.best_artifact(session["id"]).is_file()) + self.assertEqual(len(self.manager.storage.list_presets()), 1) + + def test_approval_session_waits_and_accepts_modified_patch(self): + session = self.manager.create(self.payload("approval")) + deadline = time.monotonic() + 3 + while time.monotonic() < deadline: + detail = self.manager.detail(session["id"]) + if detail["state"] == "awaiting_approval": + break + time.sleep(0.01) + else: + self.fail("session did not wait for approval") + proposal = detail["proposals"][-1] + patch = {"weights": {"pose": 1.1}, "params": {}} + approved = self.manager.approve( + session["id"], proposal["id"], {"feedback": "ok", "patch": patch} + ) + self.assertEqual(approved["proposals"][-1]["state"], "approved") + self.manager.cancel(session["id"]) + self.assertEqual(self.wait_terminal(session["id"])["state"], "cancelled") + + def test_resume_discards_only_interrupted_trial_and_continues(self): + mode, config, objective, fallback = self.manager.parse_create(self.payload()) + session = self.manager.storage.create_session(mode, config, objective, fallback) + baseline = self.manager.storage.create_trial( + session["id"], + 0, + 0, + 1, + BASE_REWARD_CONFIGURATION, + None, + "trial-000-rung-0", + ) + self.manager._execute_trial(self.manager.storage.get_session(session["id"]), baseline) + interrupted = self.manager.storage.create_trial( + session["id"], + 1, + 0, + 1, + baseline["rewardConfig"], + None, + "trial-001-rung-0", + ) + self.manager.storage.update_trial(interrupted["id"], state="interrupted") + self.manager.storage.update_session(session["id"], state="interrupted") + self.manager.resume(session["id"]) + completed = self.wait_terminal(session["id"]) + self.assertEqual(completed["state"], "succeeded", completed["message"]) + self.assertNotIn(interrupted["id"], [trial["id"] for trial in completed["trials"]]) + + def test_create_validation_and_agent_capability(self): + self.assertTrue(self.manager.capability()["configured"]) + with self.assertRaisesRegex(Exception, "只支持"): + self.manager.parse_create({"taskId": "Other"}) + self.assertEqual(self.manager.test_agent()["model"], "fake") + + +if __name__ == "__main__": + unittest.main() diff --git a/training_server/tuning/__init__.py b/training_server/tuning/__init__.py new file mode 100644 index 00000000..d51a7928 --- /dev/null +++ b/training_server/tuning/__init__.py @@ -0,0 +1,26 @@ +"""Reward auto-tuning support for the local training service.""" + +from .schema import ( + BASE_REWARD_CONFIGURATION, + PARAMETER_SPECS, + WEIGHT_SPECS, + RewardConfigError, + apply_reward_configuration, + merge_proposal, + validate_configuration, + validate_proposal, +) +from .scoring import DEFAULT_OBJECTIVE_WEIGHTS, score_evaluation + +__all__ = [ + "BASE_REWARD_CONFIGURATION", + "DEFAULT_OBJECTIVE_WEIGHTS", + "PARAMETER_SPECS", + "WEIGHT_SPECS", + "RewardConfigError", + "apply_reward_configuration", + "merge_proposal", + "score_evaluation", + "validate_configuration", + "validate_proposal", +] diff --git a/training_server/tuning/advisor.py b/training_server/tuning/advisor.py new file mode 100644 index 00000000..afeffed5 --- /dev/null +++ b/training_server/tuning/advisor.py @@ -0,0 +1,134 @@ +"""PydanticAI adapter for the DeepSeek reward-tuning advisor.""" + +from __future__ import annotations + +import hashlib +import json +import os +from dataclasses import dataclass +from typing import Any + +from .schema import validate_proposal + +SYSTEM_PROMPT = """你是 Unitree Go2 强化学习奖励调参专家。 +只根据提供的数值配置、训练曲线摘要和固定评估结果提出下一轮稀疏修改。 +必须优先保持速度跟踪与跌倒安全门槛;每轮最多修改四个白名单标量,不得改变符号、函数、传感器或结构。 +不要建议 Python 代码、命令、文件路径或白名单外参数。输出必须符合 RewardProposal schema。 +""" + + +class AdvisorUnavailable(RuntimeError): + pass + + +@dataclass(frozen=True) +class AdvisorConfig: + api_key: str | None + base_url: str = "https://api.deepseek.com" + model: str = "deepseek-v4-flash" + + @classmethod + def from_environment(cls) -> AdvisorConfig: + return cls( + api_key=os.environ.get("DEEPSEEK_API_KEY"), + base_url=os.environ.get("MUJOCO_TUNING_AGENT_BASE_URL", "https://api.deepseek.com"), + model=os.environ.get("MUJOCO_TUNING_AGENT_MODEL", "deepseek-v4-flash"), + ) + + +class DeepSeekAdvisor: + def __init__(self, config: AdvisorConfig | None = None): + self.config = config or AdvisorConfig.from_environment() + self._cached_agent = None + + def capability(self) -> dict[str, Any]: + try: + import pydantic_ai # noqa: F401 + except ImportError: + installed = False + else: + installed = True + return { + "configured": bool(self.config.api_key) and installed, + "apiKeyConfigured": bool(self.config.api_key), + "frameworkInstalled": installed, + "model": self.config.model, + "baseUrl": self.config.base_url, + } + + def _agent(self): + if self._cached_agent is not None: + return self._cached_agent + if not self.config.api_key: + raise AdvisorUnavailable("未配置 DEEPSEEK_API_KEY") + try: + import httpx2 + from pydantic import BaseModel, Field + from pydantic_ai import Agent, PromptedOutput + from pydantic_ai.models.openai import OpenAIChatModel + from pydantic_ai.providers.openai import OpenAIProvider + except ImportError as error: + raise AdvisorUnavailable( + "缺少 PydanticAI,请安装 training_server/requirements.txt" + ) from error + + class RewardProposalOutput(BaseModel): + weights: dict[str, float] = Field(default_factory=dict) + params: dict[str, float] = Field(default_factory=dict) + rationale: str = Field(min_length=1, max_length=2000) + expected_impact: dict[str, str] = Field(default_factory=dict) + confidence: float = Field(ge=0.0, le=1.0) + + proxy = os.environ.get("HTTPS_PROXY") or os.environ.get("ALL_PROXY") + if proxy and proxy.startswith("socks://"): + proxy = "socks5://" + proxy.removeprefix("socks://") + http_client = httpx2.AsyncClient(proxy=proxy, trust_env=False, timeout=60.0) + provider = OpenAIProvider( + base_url=self.config.base_url, api_key=self.config.api_key, http_client=http_client + ) + model = OpenAIChatModel(self.config.model, provider=provider) # type: ignore[arg-type] + self._cached_agent = Agent( + model, + output_type=PromptedOutput(RewardProposalOutput), + system_prompt=SYSTEM_PROMPT, + retries=2, + model_settings={"temperature": 0.2}, + ) + return self._cached_agent + + def propose(self, context: dict[str, Any], previous: dict) -> dict[str, Any]: + prompt = json.dumps(context, ensure_ascii=False, separators=(",", ":"), allow_nan=False) + result = self._agent().run_sync(prompt) + output = result.output + patch = validate_proposal( + {"weights": dict(output.weights), "params": dict(output.params)}, previous + ) + try: + usage = result.usage() + usage_value = { + key: getattr(usage, key) + for key in ("requests", "input_tokens", "output_tokens", "total_tokens") + if getattr(usage, key, None) is not None + } + except (AttributeError, TypeError): + usage_value = {} + return { + "patch": patch, + "rationale": output.rationale, + "expectedImpact": dict(output.expected_impact), + "confidence": float(output.confidence), + "promptHash": hashlib.sha256(prompt.encode()).hexdigest(), + "usage": usage_value, + "model": self.config.model, + } + + def test_connection(self) -> dict[str, Any]: + base = { + "weights": {"track_linear_velocity": 1.0}, + "params": {}, + "instruction": "仅返回一个合法示例:把 track_linear_velocity 改为 1.1。", + } + # A minimal full previous config is supplied by callers for actual proposals; + # connectivity probing only verifies the provider and structured response path. + result = self._agent().run_sync(json.dumps(base, ensure_ascii=False)) + return {"ok": True, "model": self.config.model, "outputType": type(result.output).__name__} diff --git a/training_server/tuning/manager.py b/training_server/tuning/manager.py new file mode 100644 index 00000000..a63eeba7 --- /dev/null +++ b/training_server/tuning/manager.py @@ -0,0 +1,715 @@ +"""Persistent reward tuning session orchestrator.""" + +from __future__ import annotations + +import json +import os +import re +import shutil +import subprocess +import threading +from contextlib import suppress +from copy import deepcopy +from pathlib import Path +from typing import Any + +from .advisor import AdvisorUnavailable, DeepSeekAdvisor +from .process import GpuLease, ResourceBusyError, terminate_process +from .schema import ( + BASE_REWARD_CONFIGURATION, + merge_proposal, + validate_proposal, +) +from .scoring import DEFAULT_OBJECTIVE_WEIGHTS, score_evaluation, validate_objective_weights +from .storage import TuningStorage, now_iso +from .study import OptunaStudies +from .tensorboard import ingest_scalars + +ACTIVE_SESSION_STATES = { + "queued", + "running", + "evaluating", + "awaiting_approval", + "paused", + "interrupted", +} +RUNNING_STATES = {"queued", "running", "evaluating"} +RUN_NAME = re.compile(r"^[A-Za-z0-9_.-]{1,64}$") + + +class TuningError(RuntimeError): + pass + + +class TuningManager: + def __init__( + self, + trainer_root: Path, + python: str, + data_root: Path, + lease: GpuLease, + advisor: Any | None = None, + ): + self.trainer_root = trainer_root.expanduser().resolve() + self.python = python + self.data_root = data_root.expanduser().resolve() + self.data_root.mkdir(parents=True, exist_ok=True) + self.storage = TuningStorage(self.data_root / "tuning.sqlite3") + self.storage.recover_interrupted() + self.studies = OptunaStudies(self.data_root) + self.lease = lease + self.advisor = advisor or DeepSeekAdvisor() + self.lock = threading.RLock() + self.condition = threading.Condition(self.lock) + self.workers: dict[str, threading.Thread] = {} + self.processes: dict[str, subprocess.Popen[str]] = {} + self.cancel_events: dict[str, threading.Event] = {} + + def capability(self) -> dict[str, Any]: + capability = self.advisor.capability() + capability.update({"ready": (self.trainer_root / "scripts" / "evaluate.py").is_file()}) + return capability + + @staticmethod + def _integer(payload: dict, name: str, default: int, minimum: int, maximum: int) -> int: + value = payload.get(name, default) + if isinstance(value, bool) or not isinstance(value, int) or not minimum <= value <= maximum: + raise TuningError(f"{name} 必须在 {minimum}–{maximum} 之间") + return value + + def parse_create(self, payload: Any) -> tuple[str, dict, dict, bool]: + if not isinstance(payload, dict): + raise TuningError("请求体必须是 JSON 对象") + 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") + run_name = payload.get("runName", "auto-tune") + if not isinstance(run_name, str) or not RUN_NAME.fullmatch(run_name): + raise TuningError("runName 格式无效") + gpu_ids = payload.get("gpuIds", [0]) + if ( + not isinstance(gpu_ids, list) + or not gpu_ids + or any( + isinstance(value, bool) or not isinstance(value, int) or not 0 <= value <= 255 + for value in gpu_ids + ) + ): + raise TuningError("gpuIds 必须是非空非负整数数组") + trial_count = self._integer(payload, "trialCount", 12, 4, 20) + rung0 = self._integer(payload, "initialIterations", 300, 1, 1000000) + rung1 = self._integer(payload, "middleIterations", 900, rung0, 1000000) + rung2 = self._integer(payload, "finalIterations", 2000, rung1, 1000000) + config = { + "taskId": "Unitree-Go2-Flat", + "numEnvs": self._integer(payload, "numEnvs", 4096, 1, 16384), + "seed": self._integer(payload, "seed", 42, 0, 2147483647), + "runName": run_name, + "gpuIds": gpu_ids, + "trialCount": trial_count, + "rungs": [rung0, rung1, rung2], + "promote": [trial_count, min(4, trial_count), min(2, trial_count)], + "evalNumEnvs": self._integer(payload, "evalNumEnvs", 256, 1, 4096), + "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) + ) + fallback = payload.get("fallbackEnabled", False) + if not isinstance(fallback, bool): + raise TuningError("fallbackEnabled 必须是布尔值") + return mode, config, objective, fallback + + def create(self, payload: Any) -> dict: + mode, config, objective, fallback = self.parse_create(payload) + capability = self.capability() + if not capability["ready"]: + raise TuningError("评估入口未就绪") + if not capability["configured"] and not fallback: + raise TuningError("DeepSeek Agent 未配置;设置 DEEPSEEK_API_KEY 或显式启用 fallback") + with self.lock: + active = [ + session + for session in self.storage.list_sessions() + if session["state"] in ACTIVE_SESSION_STATES + ] + if active: + raise ResourceBusyError("已有调参 session 未结束") + session = self.storage.create_session(mode, config, objective, fallback) + self._start_worker(session["id"], resume=False) + return self.detail(session["id"]) + + def _start_worker(self, session_id: str, resume: bool) -> None: + cancel = threading.Event() + self.cancel_events[session_id] = cancel + worker = threading.Thread( + target=self._run_session, + args=(session_id, resume, cancel), + name=f"tuning-{session_id[:8]}", + daemon=True, + ) + self.workers[session_id] = worker + worker.start() + + def detail(self, session_id: str) -> dict: + session = self.storage.get_session(session_id) + session["trials"] = self.storage.list_trials(session_id) + session["proposals"] = self.storage.list_proposals(session_id) + session["audit"] = self.storage.audit_events(session_id) + return session + + def list(self) -> list[dict]: + return self.storage.list_sessions() + + def _session_root(self, session_id: str) -> Path: + root = (self.data_root / "sessions" / session_id).resolve() + if not root.is_relative_to(self.data_root): + raise TuningError("非法 session 路径") + root.mkdir(parents=True, exist_ok=True) + return root + + def _run_command( + self, + session_id: str, + command: list[str], + cwd: Path, + environment: dict[str, str], + log_path: Path, + ) -> int: + owner = f"tuning:{session_id}" + self.lease.acquire(owner) + try: + process = subprocess.Popen( + command, + cwd=cwd, + env=environment, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + encoding="utf-8", + errors="replace", + bufsize=1, + start_new_session=True, + ) + with self.lock: + self.processes[session_id] = process + assert process.stdout is not None + with log_path.open("a", encoding="utf-8") as log: + for line in process.stdout: + log.write(line) + log.flush() + if self.cancel_events[session_id].is_set(): + terminate_process(process) + break + return process.wait() + finally: + with self.lock: + self.processes.pop(session_id, None) + self.lease.release(owner) + + @staticmethod + def _latest_checkpoint(run_dir: Path) -> Path | None: + values = [] + for path in run_dir.glob("model_*.pt"): + match = re.fullmatch(r"model_(\d+)\.pt", path.name) + if match: + values.append((int(match.group(1)), path)) + return max(values, default=(0, None), key=lambda value: value[0])[1] + + def _execute_trial( + self, session: dict, trial: dict, resume_checkpoint: Path | None = None + ) -> dict: + session_id, trial_id = session["id"], trial["id"] + root = self._session_root(session_id) + run_dir = (root / trial["runDir"]).resolve() + if not run_dir.is_relative_to(root): + raise TuningError("trial 目录越界") + run_dir.mkdir(parents=True, exist_ok=True) + reward_path = run_dir / "reward_config.json" + reward_path.write_text( + json.dumps(trial["rewardConfig"], ensure_ascii=False, indent=2) + "\n", encoding="utf-8" + ) + config = session["config"] + command = [ + self.python, + "-u", + "scripts/train.py", + config["taskId"], + f"--env.scene.num-envs={config['numEnvs']}", + f"--agent.max-iterations={trial['targetIterations']}", + f"--agent.seed={config['seed']}", + f"--agent.run-name={config['runName']}-t{trial['number']}-r{trial['rung']}", + "--agent.logger=tensorboard", + "--agent.upload-model=False", + "--gpu-ids", + json.dumps(config["gpuIds"], separators=(",", ":")), + "--output-dir", + str(run_dir), + "--reward-config", + str(reward_path), + ] + if resume_checkpoint is not None: + command.extend(("--resume-checkpoint", str(resume_checkpoint))) + environment = os.environ.copy() + environment["WANDB_MODE"] = "disabled" + environment["WANDB_SILENT"] = "true" + self.storage.update_trial( + trial_id, state="training", started_at=now_iso(), message="正在训练" + ) + self.storage.update_session( + session_id, + state="running", + message=f"正在训练 trial {trial['number']} / rung {trial['rung']}", + ) + return_code = self._run_command( + session_id, command, self.trainer_root, environment, run_dir / "train.log" + ) + ingest_scalars(self.storage, trial_id, run_dir) + if self.cancel_events[session_id].is_set(): + raise TuningError("session 已取消") + checkpoint = self._latest_checkpoint(run_dir) + policy = run_dir / "policy.onnx" + if return_code != 0 or checkpoint is None or not policy.is_file(): + raise TuningError(f"训练失败(返回码 {return_code})或缺少 checkpoint/policy.onnx") + self._wait_if_paused(session_id, self.cancel_events[session_id]) + + eval_output = run_dir / "evaluation.json" + eval_command = [ + self.python, + "-u", + "scripts/evaluate.py", + config["taskId"], + "--checkpoint", + str(checkpoint), + "--output", + str(eval_output), + "--reward-config", + str(reward_path), + f"--num-envs={config['evalNumEnvs']}", + f"--steps-per-seed={config['evalSteps']}", + "--gpu-ids", + json.dumps(config["gpuIds"], separators=(",", ":")), + ] + self.storage.update_trial(trial_id, state="evaluating", message="正在固定协议评估") + self.storage.update_session( + session_id, state="evaluating", message=f"正在评估 trial {trial['number']}" + ) + return_code = self._run_command( + session_id, eval_command, self.trainer_root, environment, run_dir / "evaluate.log" + ) + ingest_scalars(self.storage, trial_id, run_dir / "evaluation-events") + if return_code != 0 or not eval_output.is_file(): + 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: + scored = { + "eligible": True, + "score": 0.0, + "components": {}, + "metrics": evaluation["metrics"], + } + else: + scored = score_evaluation( + baseline_trial["evaluation"]["metrics"], + evaluation["metrics"], + session["objectiveWeights"], + ) + evaluation["score"] = scored + rel_checkpoint = str(checkpoint.relative_to(root)) + rel_policy = str(policy.relative_to(root)) + self.storage.update_trial( + trial_id, + state="completed", + ended_at=now_iso(), + message="训练与评估完成", + checkpoint_path=rel_checkpoint, + policy_path=rel_policy, + evaluation=evaluation, + score=scored["score"], + eligible=scored["eligible"], + ) + try: + optuna_number = self.studies.record( + session_id, + trial["rewardConfig"], + scored["score"], + scored["eligible"], + trial["rung"], + ) + self.storage.audit( + session_id, "optuna_trial_recorded", {"trialId": trial_id, "number": optuna_number} + ) + except Exception as error: + self.storage.audit( + session_id, "optuna_record_failed", {"trialId": trial_id, "error": str(error)[:500]} + ) + return self.storage.get_trial(trial_id) + + def _best(self, session_id: str, rung: int | None = None) -> dict | None: + trials = [ + trial + for trial in self.storage.list_trials(session_id) + if trial["state"] == "completed" and trial["eligible"] + ] + if rung is not None: + trials = [trial for trial in trials if trial["rung"] == rung] + return max( + trials, + key=lambda trial: trial["score"] if trial["score"] is not None else -999, + default=None, + ) + + def _proposal_context(self, session: dict) -> dict: + trials = self.storage.list_trials(session["id"])[-12:] + rejected_feedback = [ + proposal["feedback"] + for proposal in self.storage.list_proposals(session["id"]) + if proposal["state"] == "rejected" and proposal["feedback"] + ][-4:] + return { + "task": session["config"]["taskId"], + "objectiveWeights": session["objectiveWeights"], + "allowlist": "服务端将验证固定 schema;最多四项修改", + "rejectedFeedback": rejected_feedback, + "trials": [ + { + "number": t["number"], + "rung": t["rung"], + "score": t["score"], + "eligible": t["eligible"], + "rewardConfig": t["rewardConfig"], + "evaluation": t["evaluation"] and t["evaluation"].get("metrics"), + } + for t in trials + ], + } + + def _fallback_patch(self, previous: dict, index: int) -> dict: + names = ("track_linear_velocity", "action_rate_l2", "body_orientation_l2", "foot_slip") + 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) + + def _request_proposal( + self, session: dict, previous: dict, base_trial_id: str, index: int + ) -> dict: + try: + result = self.advisor.propose(self._proposal_context(session), previous) + source = "agent" + except Exception as error: + if not session["fallbackEnabled"]: + raise AdvisorUnavailable(str(error)) from error + result = { + "patch": self._fallback_patch(previous, index), + "rationale": f"Agent 不可用,显式 fallback:{error}", + "expectedImpact": {}, + "confidence": 0.2, + "promptHash": None, + "usage": {}, + "model": "optuna-fallback", + } + source = "fallback" + proposal = self.storage.create_proposal( + session["id"], + base_trial_id, + result["patch"], + result["rationale"], + result.get("expectedImpact", {}), + result["confidence"], + source, + ) + self.storage.audit( + session["id"], + "proposal_created", + { + "proposalId": proposal["id"], + "source": source, + "promptHash": result.get("promptHash"), + "usage": result.get("usage", {}), + "model": result.get("model"), + }, + ) + return proposal + + def _wait_for_approval( + self, session_id: str, proposal_id: str, cancel: threading.Event + ) -> dict: + with self.condition: + while not cancel.is_set(): + proposal = self.storage.get_proposal(proposal_id) + if proposal["state"] != "pending": + return proposal + self.condition.wait(timeout=1.0) + raise TuningError("session 已取消") + + 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: + root = self._session_root(session_id) + for interrupted in [trial for trial in trials if trial["state"] == "interrupted"]: + run_dir = (root / interrupted["runDir"]).resolve() + if run_dir.is_relative_to(root): + shutil.rmtree(run_dir, ignore_errors=True) + self.storage.delete_trial(interrupted["id"]) + self.storage.audit( + session_id, + "session_resumed", + { + "discardedInterruptedTrials": [ + t["id"] for t in trials if t["state"] == "interrupted" + ] + }, + ) + trials = self.storage.list_trials(session_id) + if not trials: + baseline_dir = "trial-000-rung-0" + trial = self.storage.create_trial( + session_id, + 0, + 0, + session["config"]["rungs"][0], + deepcopy(BASE_REWARD_CONFIGURATION), + None, + baseline_dir, + ) + self._execute_trial(session, trial) + session = self.storage.get_session(session_id) + completed_rung0 = [ + t + for t in self.storage.list_trials(session_id) + if t["rung"] == 0 and t["state"] == "completed" + ] + next_number = len({t["number"] for t in completed_rung0}) + best_score = max((t["score"] or 0.0 for t in completed_rung0), default=0.0) + no_improve = 0 + while ( + next_number < session["config"]["trialCount"] + and no_improve < session["config"]["earlyStopPatience"] + ): + if cancel.is_set(): + raise TuningError("session 已取消") + self._wait_if_paused(session_id, cancel) + base = self._best(session_id, rung=0) or completed_rung0[0] + proposal = self._request_proposal( + session, base["rewardConfig"], base["id"], next_number + ) + if session["mode"] == "approval": + self.storage.update_session( + session_id, state="awaiting_approval", message="等待批准 Agent 建议" + ) + proposal = self._wait_for_approval(session_id, proposal["id"], cancel) + if proposal["state"] == "rejected": + self.storage.audit( + session_id, + "proposal_rejected", + {"proposalId": proposal["id"], "feedback": proposal["feedback"]}, + ) + continue + else: + self.storage.decide_proposal(proposal["id"], "approved", "自动模式") + proposal = self.storage.get_proposal(proposal["id"]) + self._wait_if_paused(session_id, cancel) + reward_config = merge_proposal(base["rewardConfig"], proposal["patch"]) + trial = self.storage.create_trial( + session_id, + next_number, + 0, + session["config"]["rungs"][0], + reward_config, + proposal["id"], + f"trial-{next_number:03d}-rung-0", + ) + result = self._execute_trial(session, trial) + if result["eligible"] and (result["score"] or -999) > best_score + 0.01: + best_score = result["score"] + no_improve = 0 + else: + no_improve += 1 + self.storage.update_session(session_id, consecutive_no_improve=no_improve) + next_number += 1 + + # Promote top configurations; each new rung resumes its own previous checkpoint. + for rung in (1, 2): + self._wait_if_paused(session_id, cancel) + previous = [ + t + for t in self.storage.list_trials(session_id) + if t["rung"] == rung - 1 and t["state"] == "completed" and t["eligible"] + ] + previous.sort(key=lambda t: t["score"] or -999, reverse=True) + promoted_numbers = { + trial["number"] + for trial in self.storage.list_trials(session_id) + if trial["rung"] == rung and trial["state"] == "completed" + } + for parent in previous[: session["config"]["promote"][rung]]: + if parent["number"] in promoted_numbers: + continue + if cancel.is_set(): + raise TuningError("session 已取消") + root = self._session_root(session_id) + checkpoint = root / parent["checkpointPath"] + trial = self.storage.create_trial( + session_id, + parent["number"], + rung, + session["config"]["rungs"][rung], + parent["rewardConfig"], + parent["proposalId"], + f"trial-{parent['number']:03d}-rung-{rung}", + ) + self._execute_trial(session, trial, checkpoint) + + best = ( + self._best(session_id, rung=2) + or self._best(session_id, rung=1) + or self._best(session_id, rung=0) + ) + if best is None: + raise TuningError("没有通过安全门槛的 trial") + preset_name = f"{session['config']['runName']}-{session_id[:8]}" + self.storage.save_preset(preset_name, session_id, best["id"], best["rewardConfig"]) + self.storage.update_session( + session_id, + state="succeeded", + best_trial_id=best["id"], + current_trial_id=None, + message="调参完成", + ) + self.storage.audit( + session_id, "session_completed", {"bestTrialId": best["id"], "preset": preset_name} + ) + except Exception as error: + state = self.storage.get_session(session_id)["state"] + if cancel.is_set() or state == "cancelled": + self.storage.update_session( + session_id, state="cancelled", message="调参已取消", current_trial_id=None + ) + else: + self.storage.update_session( + session_id, state="failed", message=str(error), current_trial_id=None + ) + self.storage.audit(session_id, "session_failed", {"error": str(error)[:1000]}) + finally: + with self.lock: + self.workers.pop(session_id, None) + self.processes.pop(session_id, None) + + def _wait_if_paused(self, session_id: str, cancel: threading.Event) -> None: + with self.condition: + while self.storage.get_session(session_id)["state"] == "paused" and not cancel.is_set(): + self.condition.wait(timeout=1.0) + + def approve(self, session_id: str, proposal_id: str, payload: Any) -> dict: + proposal = self.storage.get_proposal(proposal_id) + if proposal["sessionId"] != session_id: + raise TuningError("proposal 不属于该 session") + patch = proposal["patch"] + feedback = None + if isinstance(payload, dict): + feedback = payload.get("feedback") + if "patch" in payload: + base = self.storage.get_trial(proposal["baseTrialId"]) + patch = validate_proposal(payload["patch"], base["rewardConfig"]) + if not self.storage.decide_proposal(proposal_id, "approved", feedback, patch): + raise TuningError("proposal 已处理") + self.storage.audit( + session_id, + "proposal_approved", + {"proposalId": proposal_id, "modified": patch != proposal["patch"]}, + ) + with self.condition: + self.condition.notify_all() + return self.detail(session_id) + + def reject(self, session_id: str, proposal_id: str, payload: Any) -> dict: + feedback = payload.get("feedback", "") if isinstance(payload, dict) else "" + if not isinstance(feedback, str) or len(feedback) > 2000: + raise TuningError("feedback 无效") + proposal = self.storage.get_proposal(proposal_id) + if proposal["sessionId"] != session_id: + raise TuningError("proposal 不属于该 session") + if not self.storage.decide_proposal(proposal_id, "rejected", feedback): + raise TuningError("proposal 已处理") + with self.condition: + self.condition.notify_all() + return self.detail(session_id) + + def pause(self, session_id: str) -> dict: + session = self.storage.get_session(session_id) + if session["state"] not in RUNNING_STATES | {"awaiting_approval"}: + raise TuningError("当前状态不能暂停") + self.storage.update_session( + session_id, state="paused", message="已暂停后续调度;当前子进程将完成" + ) + return self.detail(session_id) + + def resume(self, session_id: str) -> dict: + session = self.storage.get_session(session_id) + if session["state"] == "paused": + self.storage.update_session(session_id, state="running", message="继续调参") + with self.condition: + self.condition.notify_all() + elif session["state"] == "interrupted": + self.storage.update_session(session_id, state="queued", message="从最近完整结果恢复") + self._start_worker(session_id, resume=True) + else: + raise TuningError("当前状态不能恢复") + return self.detail(session_id) + + def cancel(self, session_id: str) -> dict: + self.storage.get_session(session_id) + self.storage.update_session(session_id, state="cancelled", message="正在取消") + event = self.cancel_events.get(session_id) + if event: + event.set() + process = self.processes.get(session_id) + if process: + terminate_process(process) + with self.condition: + self.condition.notify_all() + return self.detail(session_id) + + def metrics( + self, session_id: str, trial_id: str, tags: list[str] | None, max_points: int + ) -> dict: + trial = self.storage.get_trial(trial_id) + if trial["sessionId"] != session_id: + raise TuningError("trial 不属于该 session") + return {"trialId": trial_id, "series": self.storage.metrics(trial_id, tags, max_points)} + + def best_artifact(self, session_id: str) -> Path: + session = self.storage.get_session(session_id) + if not session["bestTrialId"]: + raise TuningError("尚无最佳策略") + trial = self.storage.get_trial(session["bestTrialId"]) + if not trial["policyPath"]: + raise TuningError("最佳策略文件不存在") + root = self._session_root(session_id) + path = (root / trial["policyPath"]).resolve() + if not path.is_relative_to(root) or not path.is_file(): + raise TuningError("最佳策略文件不存在") + return path + + def preset_config(self, preset_id: str) -> dict: + return self.storage.get_preset(preset_id)["rewardConfig"] + + def test_agent(self) -> dict: + try: + return self.advisor.test_connection() + except Exception as error: + raise TuningError(f"Agent 连接测试失败:{error}") from error + + def shutdown(self) -> None: + for session_id in list(self.workers): + with suppress(KeyError, TuningError): + self.cancel(session_id) + for worker in list(self.workers.values()): + worker.join(timeout=7) diff --git a/training_server/tuning/process.py b/training_server/tuning/process.py new file mode 100644 index 00000000..0537e6f8 --- /dev/null +++ b/training_server/tuning/process.py @@ -0,0 +1,46 @@ +"""Shared GPU lease and process-group lifecycle helpers.""" + +from __future__ import annotations + +import os +import signal +import subprocess +import threading +from contextlib import suppress + + +class ResourceBusyError(RuntimeError): + pass + + +class GpuLease: + def __init__(self): + self.lock = threading.RLock() + self.owner: str | None = None + + def acquire(self, owner: str) -> None: + with self.lock: + if self.owner is not None and self.owner != owner: + raise ResourceBusyError(f"计算资源正由 {self.owner} 使用") + self.owner = owner + + def release(self, owner: str) -> None: + with self.lock: + if self.owner == owner: + self.owner = None + + def public(self) -> str | None: + with self.lock: + return self.owner + + +def terminate_process(process: subprocess.Popen[str], grace_seconds: float = 5.0) -> None: + if process.poll() is not None: + return + with suppress(ProcessLookupError): + os.killpg(process.pid, signal.SIGTERM) + try: + process.wait(timeout=grace_seconds) + except subprocess.TimeoutExpired: + with suppress(ProcessLookupError): + os.killpg(process.pid, signal.SIGKILL) diff --git a/training_server/tuning/schema.py b/training_server/tuning/schema.py new file mode 100644 index 00000000..66f4b09c --- /dev/null +++ b/training_server/tuning/schema.py @@ -0,0 +1,213 @@ +"""Pure-Python reward tuning schema shared by the service and trainer.""" + +from __future__ import annotations + +import math +from collections.abc import Mapping +from copy import deepcopy +from dataclasses import dataclass +from typing import Any + +MAX_PROPOSAL_CHANGES = 4 +MIN_CHANGE_RATIO = 0.5 +MAX_CHANGE_RATIO = 2.0 + + +class RewardConfigError(ValueError): + """A reward configuration or proposal violated the allowlist.""" + + +@dataclass(frozen=True) +class NumericSpec: + minimum: float + maximum: float + default: float + allow_zero: bool = True + + +WEIGHT_SPECS: dict[str, NumericSpec] = { + "track_linear_velocity": NumericSpec(0.5, 3.0, 1.0, False), + "track_angular_velocity": NumericSpec(0.25, 2.0, 1.0, False), + "body_orientation_l2": NumericSpec(-3.0, -0.1, -1.0, False), + "pose": NumericSpec(0.0, 2.5, 1.0), + "body_ang_vel": NumericSpec(-0.2, 0.0, -0.05), + "angular_momentum": NumericSpec(-0.1, 0.0, -0.025), + "is_terminated": NumericSpec(-400.0, -50.0, -200.0, False), + "joint_acc_l2": NumericSpec(-2.0e-6, 0.0, -2.5e-7), + "joint_pos_limits": NumericSpec(-30.0, -2.0, -10.0, False), + "action_rate_l2": NumericSpec(-0.2, -0.005, -0.05, False), + "foot_gait": NumericSpec(0.0, 1.5, 0.5), + "foot_clearance": NumericSpec(-3.0, 0.0, -1.0), + "foot_slip": NumericSpec(-1.0, 0.0, -0.25), + "soft_landing": NumericSpec(-5.0e-3, 0.0, -1.0e-3), + "stand_still": NumericSpec(-3.0, 0.0, -1.0), + "electrical_power": NumericSpec(-5.0e-3, 0.0, 0.0), +} + +PARAMETER_SPECS: dict[str, NumericSpec] = { + "track_linear_velocity.std": NumericSpec(0.25, 1.0, math.sqrt(0.25), False), + "track_angular_velocity.std": NumericSpec(0.35, 1.2, math.sqrt(0.5), False), + "pose.std_standing_scale": NumericSpec(0.5, 2.0, 1.0, False), + "pose.std_walking_scale": NumericSpec(0.5, 2.0, 1.0, False), + "pose.std_running_scale": NumericSpec(0.5, 2.0, 1.0, False), + "pose.walking_threshold": NumericSpec(0.05, 0.5, 0.1, False), + "pose.running_threshold": NumericSpec(1.0, 2.5, 1.5, False), + "foot_gait.period": NumericSpec(0.4, 0.8, 0.6, False), + "foot_gait.threshold": NumericSpec(0.45, 0.65, 0.56, False), + "foot_gait.command_threshold": NumericSpec(0.02, 0.3, 0.1, False), + "foot_clearance.target_height": NumericSpec(0.06, 0.16, 0.1, False), + "foot_clearance.command_threshold": NumericSpec(0.02, 0.3, 0.1, False), + "foot_slip.command_threshold": NumericSpec(0.02, 0.3, 0.1, False), + "soft_landing.command_threshold": NumericSpec(0.02, 0.3, 0.1, False), + "stand_still.command_threshold": NumericSpec(0.02, 0.3, 0.1, False), +} + +BASE_REWARD_CONFIGURATION: dict[str, dict[str, float]] = { + "weights": {name: spec.default for name, spec in WEIGHT_SPECS.items()}, + "params": {name: spec.default for name, spec in PARAMETER_SPECS.items()}, +} + + +def _number(name: str, value: Any, spec: NumericSpec) -> float: + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise RewardConfigError(f"{name} 必须是数值") + result = float(value) + if not math.isfinite(result): + raise RewardConfigError(f"{name} 必须是有限数值") + if result == 0.0 and not spec.allow_zero: + raise RewardConfigError(f"{name} 不允许关闭") + if result < spec.minimum or result > spec.maximum: + raise RewardConfigError(f"{name} 必须在 {spec.minimum}–{spec.maximum} 之间") + return result + + +def _mapping(value: Any, name: str) -> Mapping[str, Any]: + if not isinstance(value, Mapping): + raise RewardConfigError(f"{name} 必须是对象") + return value + + +def _cross_validate(config: Mapping[str, Mapping[str, float]]) -> None: + params = config["params"] + if params["pose.walking_threshold"] >= params["pose.running_threshold"]: + raise RewardConfigError("pose.walking_threshold 必须小于 pose.running_threshold") + + +def validate_configuration(value: Any) -> 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") + raw_weights = _mapping(root["weights"], "weights") + raw_params = _mapping(root["params"], "params") + if set(raw_weights) != set(WEIGHT_SPECS): + raise RewardConfigError("weights 必须完整且不能包含未知奖励项") + if set(raw_params) != set(PARAMETER_SPECS): + raise RewardConfigError("params 必须完整且不能包含未知参数") + config = { + "weights": { + name: _number(f"weights.{name}", raw_weights[name], spec) + for name, spec in WEIGHT_SPECS.items() + }, + "params": { + name: _number(f"params.{name}", raw_params[name], spec) + for name, spec in PARAMETER_SPECS.items() + }, + } + _cross_validate(config) + return config + + +def validate_proposal(value: Any, previous: Any) -> dict[str, dict[str, float]]: + """Validate a sparse Agent patch relative to a complete previous config.""" + current = validate_configuration(previous) + root = _mapping(value, "proposal") + if not set(root).issubset({"weights", "params"}): + raise RewardConfigError("proposal 只能包含 weights 和 params") + raw_weights = _mapping(root.get("weights", {}), "weights") + raw_params = _mapping(root.get("params", {}), "params") + if len(raw_weights) + len(raw_params) == 0: + 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) + if unknown_weights: + raise RewardConfigError(f"未知奖励项:{', '.join(sorted(unknown_weights))}") + if unknown_params: + raise RewardConfigError(f"未知奖励参数:{', '.join(sorted(unknown_params))}") + + 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]) + old = current["weights"][name] + if old != 0.0 and value_number != 0.0: + ratio = abs(value_number / old) + if ratio < MIN_CHANGE_RATIO or ratio > MAX_CHANGE_RATIO: + raise RewardConfigError( + f"weights.{name} 单轮变化必须在旧值幅度的 " + f"{MIN_CHANGE_RATIO}×–{MAX_CHANGE_RATIO}×" + ) + if value_number == old: + 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]) + old = current["params"][name] + ratio = abs(value_number / old) + if ratio < MIN_CHANGE_RATIO or ratio > MAX_CHANGE_RATIO: + raise RewardConfigError( + f"params.{name} 单轮变化必须在旧值的 {MIN_CHANGE_RATIO}×–{MAX_CHANGE_RATIO}×" + ) + if value_number == old: + raise RewardConfigError(f"params.{name} 没有发生变化") + patch["params"][name] = value_number + + candidate = deepcopy(current) + candidate["weights"].update(patch["weights"]) + candidate["params"].update(patch["params"]) + _cross_validate(candidate) + return patch + + +def merge_proposal(previous: Any, proposal: Any) -> dict[str, dict[str, float]]: + current = validate_configuration(previous) + patch = validate_proposal(proposal, current) + merged = deepcopy(current) + merged["weights"].update(patch["weights"]) + merged["params"].update(patch["params"]) + return validate_configuration(merged) + + +def apply_reward_configuration(env_cfg: Any, value: Any) -> None: + """Apply a validated full config to a fresh mjlab environment config.""" + config = validate_configuration(value) + for name, weight in config["weights"].items(): + if name not in env_cfg.rewards: + raise RewardConfigError(f"环境缺少奖励项:{name}") + env_cfg.rewards[name].weight = weight + + params = config["params"] + direct = { + "track_linear_velocity.std": ("track_linear_velocity", "std"), + "track_angular_velocity.std": ("track_angular_velocity", "std"), + "pose.walking_threshold": ("pose", "walking_threshold"), + "pose.running_threshold": ("pose", "running_threshold"), + "foot_gait.period": ("foot_gait", "period"), + "foot_gait.threshold": ("foot_gait", "threshold"), + "foot_gait.command_threshold": ("foot_gait", "command_threshold"), + "foot_clearance.target_height": ("foot_clearance", "target_height"), + "foot_clearance.command_threshold": ("foot_clearance", "command_threshold"), + "foot_slip.command_threshold": ("foot_slip", "command_threshold"), + "soft_landing.command_threshold": ("soft_landing", "command_threshold"), + "stand_still.command_threshold": ("stand_still", "command_threshold"), + } + for path, (term, parameter) in direct.items(): + env_cfg.rewards[term].params[parameter] = params[path] + for regime in ("standing", "walking", "running"): + key = f"std_{regime}" + scale = params[f"pose.{key}_scale"] + baseline = env_cfg.rewards["pose"].params[key] + env_cfg.rewards["pose"].params[key] = { + pattern: float(std) * scale for pattern, std in baseline.items() + } diff --git a/training_server/tuning/scoring.py b/training_server/tuning/scoring.py new file mode 100644 index 00000000..028effe5 --- /dev/null +++ b/training_server/tuning/scoring.py @@ -0,0 +1,118 @@ +"""Stable, reward-weight-independent evaluation scoring.""" + +from __future__ import annotations + +import math +from collections.abc import Mapping +from typing import Any + +DEFAULT_OBJECTIVE_WEIGHTS = { + "velocity_tracking": 0.35, + "action_smoothness": 0.20, + "posture_stability": 0.15, + "fall_avoidance": 0.15, + "foot_slip": 0.10, + "energy": 0.05, +} + +REQUIRED_METRICS = { + "linear_velocity_rmse", + "angular_velocity_rmse", + "mean_action_acc", + "orientation_error", + "fall_rate", + "slip_velocity", + "mechanical_power", +} + +PHYSICAL_FLOORS = { + "linear_velocity_rmse": 0.10, + "angular_velocity_rmse": 0.10, + "mean_action_acc": 0.01, + "orientation_error": 0.05, + "fall_rate": 0.02, + "slip_velocity": 0.05, + "mechanical_power": 10.0, +} + + +class EvaluationError(ValueError): + pass + + +def validate_objective_weights(value: Any) -> dict[str, float]: + if not isinstance(value, Mapping) or set(value) != set(DEFAULT_OBJECTIVE_WEIGHTS): + raise EvaluationError("objectiveWeights 必须完整包含六个目标") + result: dict[str, float] = {} + for key in DEFAULT_OBJECTIVE_WEIGHTS: + raw = value[key] + if isinstance(raw, bool) or not isinstance(raw, (int, float)): + raise EvaluationError(f"objectiveWeights.{key} 必须是数值") + number = float(raw) + if not math.isfinite(number) or number < 0.0 or number > 1.0: + raise EvaluationError(f"objectiveWeights.{key} 必须在 0–1 之间") + result[key] = number + if not math.isclose(sum(result.values()), 1.0, abs_tol=1.0e-6): + raise EvaluationError("objectiveWeights 总和必须为 1") + return result + + +def validate_metrics(value: Any) -> dict[str, float]: + if not isinstance(value, Mapping): + raise EvaluationError("metrics 必须是对象") + missing = REQUIRED_METRICS - set(value) + if missing: + raise EvaluationError(f"metrics 缺少:{', '.join(sorted(missing))}") + result: dict[str, float] = {} + for key in REQUIRED_METRICS: + raw = value[key] + if isinstance(raw, bool) or not isinstance(raw, (int, float)): + raise EvaluationError(f"metrics.{key} 必须是数值") + number = float(raw) + if not math.isfinite(number) or number < 0.0: + raise EvaluationError(f"metrics.{key} 必须是非负有限数值") + result[key] = number + if result["fall_rate"] > 1.0: + raise EvaluationError("metrics.fall_rate 必须在 0–1 之间") + return result + + +def _improvement(baseline: Mapping[str, float], current: Mapping[str, float], key: str) -> float: + scale = max(abs(baseline[key]), PHYSICAL_FLOORS[key]) + return max(-1.0, min(1.0, (baseline[key] - current[key]) / scale)) + + +def score_evaluation( + baseline_value: Any, + current_value: Any, + objective_weights: Any = DEFAULT_OBJECTIVE_WEIGHTS, +) -> dict[str, Any]: + baseline = validate_metrics(baseline_value) + current = validate_metrics(current_value) + weights = validate_objective_weights(objective_weights) + components = { + "velocity_tracking": 0.8 * _improvement(baseline, current, "linear_velocity_rmse") + + 0.2 * _improvement(baseline, current, "angular_velocity_rmse"), + "action_smoothness": _improvement(baseline, current, "mean_action_acc"), + "posture_stability": _improvement(baseline, current, "orientation_error"), + "fall_avoidance": _improvement(baseline, current, "fall_rate"), + "foot_slip": _improvement(baseline, current, "slip_velocity"), + "energy": _improvement(baseline, current, "mechanical_power"), + } + tracking_limit = max( + baseline["linear_velocity_rmse"] * 1.05, baseline["linear_velocity_rmse"] + 1.0e-6 + ) + eligible = ( + current["fall_rate"] <= baseline["fall_rate"] + 0.02 + and current["linear_velocity_rmse"] <= tracking_limit + ) + total = sum(weights[key] * components[key] for key in weights) + if not eligible: + total = min(total, -1.0) + return { + "eligible": eligible, + "score": total, + "components": components, + "metrics": current, + "baselineMetrics": baseline, + } diff --git a/training_server/tuning/storage.py b/training_server/tuning/storage.py new file mode 100644 index 00000000..7efae5cd --- /dev/null +++ b/training_server/tuning/storage.py @@ -0,0 +1,521 @@ +"""SQLite persistence for tuning sessions, trials, proposals and scalar data.""" + +from __future__ import annotations + +import json +import sqlite3 +import threading +import uuid +from collections.abc import Iterator +from contextlib import contextmanager +from datetime import UTC, datetime +from pathlib import Path +from typing import Any + +SCHEMA_VERSION = 1 + + +def now_iso() -> str: + return datetime.now(UTC).isoformat() + + +def _json(value: Any) -> str: + return json.dumps(value, ensure_ascii=False, separators=(",", ":"), allow_nan=False) + + +def _decode(value: str | None) -> Any: + return json.loads(value) if value else None + + +def _lttb(points: list[dict], threshold: int) -> list[dict]: + """Largest-Triangle-Three-Buckets downsampling preserving peaks and endpoints.""" + if threshold >= len(points) or threshold < 3: + return points[:threshold] + sampled = [points[0]] + bucket_width = (len(points) - 2) / (threshold - 2) + anchor_index = 0 + for bucket in range(threshold - 2): + average_start = int((bucket + 1) * bucket_width) + 1 + average_end = min(int((bucket + 2) * bucket_width) + 1, len(points)) + average_bucket = points[average_start:average_end] or [points[-1]] + average_x = sum(point["step"] for point in average_bucket) / len(average_bucket) + average_y = sum(point["value"] for point in average_bucket) / len(average_bucket) + range_start = int(bucket * bucket_width) + 1 + range_end = min(int((bucket + 1) * bucket_width) + 1, len(points) - 1) + anchor = points[anchor_index] + selected_index = range_start + maximum_area = -1.0 + for index in range(range_start, max(range_start + 1, range_end)): + point = points[index] + area = abs( + (anchor["step"] - average_x) * (point["value"] - anchor["value"]) + - (anchor["step"] - point["step"]) * (average_y - anchor["value"]) + ) + if area > maximum_area: + maximum_area = area + selected_index = index + sampled.append(points[selected_index]) + anchor_index = selected_index + sampled.append(points[-1]) + return sampled + + +class TuningStorage: + def __init__(self, path: Path): + self.path = path.expanduser().resolve() + self.path.parent.mkdir(parents=True, exist_ok=True) + self.local = threading.local() + self._migrate() + + def connection(self) -> sqlite3.Connection: + connection = getattr(self.local, "connection", None) + if connection is None: + connection = sqlite3.connect(self.path, timeout=10, isolation_level=None) + connection.row_factory = sqlite3.Row + connection.execute("PRAGMA foreign_keys=ON") + connection.execute("PRAGMA journal_mode=WAL") + connection.execute("PRAGMA busy_timeout=10000") + self.local.connection = connection + return connection + + @contextmanager + def transaction(self) -> Iterator[sqlite3.Connection]: + connection = self.connection() + connection.execute("BEGIN IMMEDIATE") + try: + yield connection + connection.execute("COMMIT") + except Exception: + connection.execute("ROLLBACK") + raise + + def _migrate(self) -> None: + connection = self.connection() + connection.executescript( + """ + CREATE TABLE IF NOT EXISTS schema_migrations(version INTEGER PRIMARY KEY); + CREATE TABLE IF NOT EXISTS sessions( + id TEXT PRIMARY KEY, state TEXT NOT NULL, mode TEXT NOT NULL, + created_at TEXT NOT NULL, updated_at TEXT NOT NULL, + config_json TEXT NOT NULL, objective_json TEXT NOT NULL, + message TEXT NOT NULL, current_trial_id TEXT, best_trial_id TEXT, + consecutive_no_improve INTEGER NOT NULL DEFAULT 0, + fallback_enabled INTEGER NOT NULL DEFAULT 0 + ); + CREATE TABLE IF NOT EXISTS trials( + id TEXT PRIMARY KEY, + session_id TEXT NOT NULL REFERENCES sessions(id) ON DELETE CASCADE, + number INTEGER NOT NULL, state TEXT NOT NULL, rung INTEGER NOT NULL, + target_iterations INTEGER NOT NULL, reward_config_json TEXT NOT NULL, + proposal_id TEXT, run_dir TEXT NOT NULL, checkpoint_path TEXT, + policy_path TEXT, evaluation_json TEXT, score REAL, eligible INTEGER, + created_at TEXT NOT NULL, started_at TEXT, ended_at TEXT, message TEXT NOT NULL, + UNIQUE(session_id, number, rung) + ); + CREATE TABLE IF NOT EXISTS proposals( + id TEXT PRIMARY KEY, + session_id TEXT NOT NULL REFERENCES sessions(id) ON DELETE CASCADE, + base_trial_id TEXT, state TEXT NOT NULL, source TEXT NOT NULL, + patch_json TEXT NOT NULL, rationale TEXT NOT NULL, + expected_json TEXT, confidence REAL NOT NULL, + created_at TEXT NOT NULL, decided_at TEXT, feedback TEXT + ); + CREATE TABLE IF NOT EXISTS metric_points( + trial_id TEXT NOT NULL REFERENCES trials(id) ON DELETE CASCADE, + tag TEXT NOT NULL, step INTEGER NOT NULL, + wall_time REAL NOT NULL, value REAL NOT NULL, + PRIMARY KEY(trial_id, tag, step) + ); + CREATE TABLE IF NOT EXISTS audit_events( + id INTEGER PRIMARY KEY AUTOINCREMENT, + session_id TEXT NOT NULL REFERENCES sessions(id) ON DELETE CASCADE, + event_type TEXT NOT NULL, payload_json TEXT NOT NULL, created_at TEXT NOT NULL + ); + CREATE TABLE IF NOT EXISTS presets( + id TEXT PRIMARY KEY, name TEXT NOT NULL UNIQUE, session_id TEXT NOT NULL, + trial_id TEXT NOT NULL, reward_config_json TEXT NOT NULL, created_at TEXT NOT NULL + ); + CREATE INDEX IF NOT EXISTS idx_trials_session ON trials(session_id, number, rung); + CREATE INDEX IF NOT EXISTS idx_proposals_session ON proposals(session_id, created_at); + CREATE INDEX IF NOT EXISTS idx_metrics_trial_tag ON metric_points(trial_id, tag, step); + """ + ) + connection.execute( + "INSERT OR IGNORE INTO schema_migrations(version) VALUES (?)", (SCHEMA_VERSION,) + ) + + def recover_interrupted(self) -> None: + at = now_iso() + with self.transaction() as connection: + connection.execute( + "UPDATE trials SET state='interrupted', ended_at=?, " + "message='服务重启中断,等待显式恢复' " + "WHERE state IN ('training','evaluating')", + (at,), + ) + connection.execute( + "UPDATE sessions SET state='interrupted', updated_at=?, " + "message='服务重启中断,可从完整 checkpoint 恢复' " + "WHERE state IN ('running','evaluating')", + (at,), + ) + + def create_session(self, mode: str, config: dict, objective: dict, fallback: bool) -> dict: + session_id, at = uuid.uuid4().hex, now_iso() + with self.transaction() as connection: + connection.execute( + "INSERT INTO sessions(" + "id,state,mode,created_at,updated_at,config_json,objective_json," + "message,fallback_enabled) " + "VALUES (?, 'queued', ?, ?, ?, ?, ?, '等待基线训练', ?)", + (session_id, mode, at, at, _json(config), _json(objective), int(fallback)), + ) + connection.execute( + "INSERT INTO audit_events(session_id,event_type,payload_json,created_at) " + "VALUES (?,?,?,?)", + (session_id, "session_created", _json({"mode": mode}), at), + ) + return self.get_session(session_id) + + def _session(self, row: sqlite3.Row) -> dict: + return { + "id": row["id"], + "state": row["state"], + "mode": row["mode"], + "createdAt": row["created_at"], + "updatedAt": row["updated_at"], + "config": _decode(row["config_json"]), + "objectiveWeights": _decode(row["objective_json"]), + "message": row["message"], + "currentTrialId": row["current_trial_id"], + "bestTrialId": row["best_trial_id"], + "consecutiveNoImprove": row["consecutive_no_improve"], + "fallbackEnabled": bool(row["fallback_enabled"]), + } + + def get_session(self, session_id: str) -> dict: + row = ( + self.connection().execute("SELECT * FROM sessions WHERE id=?", (session_id,)).fetchone() + ) + if row is None: + raise KeyError(session_id) + return self._session(row) + + def list_sessions(self, limit: int = 50) -> list[dict]: + rows = ( + self.connection() + .execute("SELECT * FROM sessions ORDER BY created_at DESC LIMIT ?", (limit,)) + .fetchall() + ) + return [self._session(row) for row in rows] + + def update_session(self, session_id: str, **changes: Any) -> bool: + columns = { + "state": "state", + "message": "message", + "current_trial_id": "current_trial_id", + "best_trial_id": "best_trial_id", + "consecutive_no_improve": "consecutive_no_improve", + } + values, assignments = [], [] + for key, value in changes.items(): + if key not in columns: + raise ValueError(key) + assignments.append(f"{columns[key]}=?") + values.append(value) + assignments.append("updated_at=?") + values.extend((now_iso(), session_id)) + cursor = self.connection().execute( + f"UPDATE sessions SET {', '.join(assignments)} WHERE id=?", values + ) + return cursor.rowcount == 1 + + def create_trial( + self, + session_id: str, + number: int, + rung: int, + target: int, + reward_config: dict, + proposal_id: str | None, + run_dir: str, + ) -> dict: + trial_id, at = uuid.uuid4().hex, now_iso() + with self.transaction() as connection: + connection.execute( + "INSERT INTO trials(" + "id,session_id,number,state,rung,target_iterations,reward_config_json," + "proposal_id,run_dir,created_at,message) " + "VALUES (?,?,?,'queued',?,?,?,?,?,?,'等待训练')", + ( + trial_id, + session_id, + number, + rung, + target, + _json(reward_config), + proposal_id, + run_dir, + at, + ), + ) + connection.execute( + "UPDATE sessions SET current_trial_id=?,updated_at=? WHERE id=?", + (trial_id, at, session_id), + ) + return self.get_trial(trial_id) + + def _trial(self, row: sqlite3.Row) -> dict: + return { + "id": row["id"], + "sessionId": row["session_id"], + "number": row["number"], + "state": row["state"], + "rung": row["rung"], + "targetIterations": row["target_iterations"], + "rewardConfig": _decode(row["reward_config_json"]), + "proposalId": row["proposal_id"], + "runDir": row["run_dir"], + "checkpointPath": row["checkpoint_path"], + "policyPath": row["policy_path"], + "evaluation": _decode(row["evaluation_json"]), + "score": row["score"], + "eligible": None if row["eligible"] is None else bool(row["eligible"]), + "createdAt": row["created_at"], + "startedAt": row["started_at"], + "endedAt": row["ended_at"], + "message": row["message"], + } + + def get_trial(self, trial_id: str) -> dict: + row = self.connection().execute("SELECT * FROM trials WHERE id=?", (trial_id,)).fetchone() + if row is None: + raise KeyError(trial_id) + return self._trial(row) + + def list_trials(self, session_id: str) -> list[dict]: + rows = ( + self.connection() + .execute("SELECT * FROM trials WHERE session_id=? ORDER BY number,rung", (session_id,)) + .fetchall() + ) + return [self._trial(row) for row in rows] + + def delete_trial(self, trial_id: str) -> bool: + cursor = self.connection().execute( + "DELETE FROM trials WHERE id=? AND state='interrupted'", (trial_id,) + ) + return cursor.rowcount == 1 + + def update_trial(self, trial_id: str, **changes: Any) -> bool: + columns = { + "state": "state", + "message": "message", + "checkpoint_path": "checkpoint_path", + "policy_path": "policy_path", + "score": "score", + "eligible": "eligible", + "started_at": "started_at", + "ended_at": "ended_at", + "evaluation": "evaluation_json", + } + values, assignments = [], [] + for key, value in changes.items(): + if key not in columns: + raise ValueError(key) + if key == "evaluation": + value = _json(value) + if key == "eligible": + value = int(value) + assignments.append(f"{columns[key]}=?") + values.append(value) + values.append(trial_id) + cursor = self.connection().execute( + f"UPDATE trials SET {', '.join(assignments)} WHERE id=?", values + ) + return cursor.rowcount == 1 + + def create_proposal( + self, + session_id: str, + base_trial_id: str | None, + patch: dict, + rationale: str, + expected: Any, + confidence: float, + source: str = "agent", + ) -> dict: + proposal_id, at = uuid.uuid4().hex, now_iso() + self.connection().execute( + "INSERT INTO proposals(" + "id,session_id,base_trial_id,state,source,patch_json,rationale," + "expected_json,confidence,created_at) " + "VALUES (?,?,?,'pending',?,?,?,?,?,?)", + ( + proposal_id, + session_id, + base_trial_id, + source, + _json(patch), + rationale, + _json(expected), + confidence, + at, + ), + ) + return self.get_proposal(proposal_id) + + def _proposal(self, row: sqlite3.Row) -> dict: + return { + "id": row["id"], + "sessionId": row["session_id"], + "baseTrialId": row["base_trial_id"], + "state": row["state"], + "source": row["source"], + "patch": _decode(row["patch_json"]), + "rationale": row["rationale"], + "expectedImpact": _decode(row["expected_json"]), + "confidence": row["confidence"], + "createdAt": row["created_at"], + "decidedAt": row["decided_at"], + "feedback": row["feedback"], + } + + def get_proposal(self, proposal_id: str) -> dict: + row = ( + self.connection() + .execute("SELECT * FROM proposals WHERE id=?", (proposal_id,)) + .fetchone() + ) + if row is None: + raise KeyError(proposal_id) + return self._proposal(row) + + def list_proposals(self, session_id: str) -> list[dict]: + rows = ( + self.connection() + .execute( + "SELECT * FROM proposals WHERE session_id=? ORDER BY created_at", (session_id,) + ) + .fetchall() + ) + return [self._proposal(row) for row in rows] + + def decide_proposal( + self, proposal_id: str, state: str, feedback: str | None, patch: dict | None = None + ) -> bool: + at = now_iso() + assignments, values = ["state=?", "feedback=?", "decided_at=?"], [state, feedback, at] + if patch is not None: + assignments.append("patch_json=?") + values.append(_json(patch)) + values.extend((proposal_id,)) + cursor = self.connection().execute( + f"UPDATE proposals SET {', '.join(assignments)} WHERE id=? AND state='pending'", values + ) + return cursor.rowcount == 1 + + def insert_metrics(self, trial_id: str, points: list[tuple[str, int, float, float]]) -> None: + self.connection().executemany( + "INSERT INTO metric_points(trial_id,tag,step,wall_time,value) " + "VALUES (?,?,?,?,?) ON CONFLICT(trial_id,tag,step) DO UPDATE SET " + "wall_time=excluded.wall_time,value=excluded.value", + [(trial_id, *point) for point in points], + ) + + def metrics( + self, trial_id: str, tags: list[str] | None = None, max_points: int = 1000 + ) -> list[dict]: + parameters: list[Any] = [trial_id] + clause = "trial_id=?" + if tags: + clause += f" AND tag IN ({','.join('?' for _ in tags)})" + parameters.extend(tags) + rows = ( + self.connection() + .execute( + f"SELECT tag,step,wall_time,value FROM metric_points " + f"WHERE {clause} ORDER BY tag,step", + parameters, + ) + .fetchall() + ) + grouped: dict[str, list[dict]] = {} + for row in rows: + grouped.setdefault(row["tag"], []).append( + {"step": row["step"], "wallTime": row["wall_time"], "value": row["value"]} + ) + series = [] + for tag, values in grouped.items(): + if len(values) > max_points: + values = _lttb(values, max_points) + series.append({"tag": tag, "points": values}) + return series + + def audit(self, session_id: str, event_type: str, payload: Any) -> None: + self.connection().execute( + "INSERT INTO audit_events(session_id,event_type,payload_json,created_at) " + "VALUES (?,?,?,?)", + (session_id, event_type, _json(payload), now_iso()), + ) + + def audit_events(self, session_id: str) -> list[dict]: + rows = ( + self.connection() + .execute("SELECT * FROM audit_events WHERE session_id=? ORDER BY id", (session_id,)) + .fetchall() + ) + return [ + { + "id": row["id"], + "type": row["event_type"], + "payload": _decode(row["payload_json"]), + "createdAt": row["created_at"], + } + for row in rows + ] + + def save_preset(self, name: str, session_id: str, trial_id: str, reward_config: dict) -> dict: + 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, + } + + 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"], + } + + 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 + ] diff --git a/training_server/tuning/study.py b/training_server/tuning/study.py new file mode 100644 index 00000000..1db8b639 --- /dev/null +++ b/training_server/tuning/study.py @@ -0,0 +1,55 @@ +"""Lazy Optuna study integration used for durable trial history and pruning metadata.""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any + + +class OptunaStudies: + def __init__(self, root: Path): + self.path = (root / "optuna.sqlite3").resolve() + self.url = f"sqlite:///{self.path}" + + def _study(self, session_id: str): + import optuna + + optuna.logging.set_verbosity(optuna.logging.WARNING) + return optuna.create_study( + study_name=f"reward-tuning-{session_id}", + storage=self.url, + direction="maximize", + load_if_exists=True, + pruner=optuna.pruners.SuccessiveHalvingPruner( + min_resource=300, reduction_factor=3, min_early_stopping_rate=0 + ), + ) + + def record( + self, + session_id: str, + reward_config: dict[str, Any], + score: float, + eligible: bool, + rung: int, + ) -> int: + """Record an externally proposed Agent config through Optuna ask/tell.""" + import optuna + + study = self._study(session_id) + trial = study.ask() + trial.set_user_attr("reward_config", reward_config) + trial.set_user_attr("eligible", eligible) + trial.set_user_attr("rung", rung) + state = optuna.trial.TrialState.COMPLETE if eligible else optuna.trial.TrialState.PRUNED + study.tell(trial, score if eligible else None, state=state) + return trial.number + + def summary(self, session_id: str) -> dict[str, Any]: + study = self._study(session_id) + completed = [trial for trial in study.trials if trial.value is not None] + return { + "studyName": study.study_name, + "trialCount": len(study.trials), + "bestValue": max((trial.value for trial in completed), default=None), + } diff --git a/training_server/tuning/tensorboard.py b/training_server/tuning/tensorboard.py new file mode 100644 index 00000000..7cca4876 --- /dev/null +++ b/training_server/tuning/tensorboard.py @@ -0,0 +1,35 @@ +"""TensorBoard scalar ingestion with optional dependency isolation.""" + +from __future__ import annotations + +from pathlib import Path + +from .storage import TuningStorage + + +class TensorboardUnavailable(RuntimeError): + pass + + +def ingest_scalars(storage: TuningStorage, trial_id: str, log_dir: Path) -> int: + """Reload all scalar events and idempotently upsert them into SQLite.""" + try: + from tensorboard.backend.event_processing.event_accumulator import EventAccumulator + except ImportError as error: + raise TensorboardUnavailable( + "缺少 tensorboard,请安装 training_server/requirements.txt" + ) from error + if not log_dir.is_dir(): + return 0 + accumulator = EventAccumulator(str(log_dir), size_guidance={"scalars": 0}) + try: + accumulator.Reload() + except (OSError, ValueError): + return 0 + points: list[tuple[str, int, float, float]] = [] + for tag in accumulator.Tags().get("scalars", []): + for event in accumulator.Scalars(tag): + points.append((tag, int(event.step), float(event.wall_time), float(event.value))) + if points: + storage.insert_metrics(trial_id, points) + return len(points) diff --git a/web_platform/ARCHITECTURE.md b/web_platform/ARCHITECTURE.md index 050c99c9..cd618cbd 100644 --- a/web_platform/ARCHITECTURE.md +++ b/web_platform/ARCHITECTURE.md @@ -14,7 +14,8 @@ src/ ├── rl/ ONNX 策略运行时、任务绑定、类型和面板 ├── simulation/ MuJoCo 会话、物理适配器和仿真控制组件 ├── telemetry/ 数据源抽象、记录器、导出和数据面板 -├── training/ 本地训练客户端、类型和面板 +├── training/ 本地训练/调参客户端、类型、共享连接和面板 +├── tuning/ 独立 tuning.html 的 Agent dashboard 与 scalar 图表 ├── viewer/ Three.js 场景、渲染和交互 ├── stores/ 跨域应用状态 └── test/ 全局测试初始化 @@ -38,4 +39,6 @@ src/ SimulationSession snapshot → app → viewer / 各业务面板 ``` +主工作台由 `index.html → src/main.tsx` 启动;自调参工作台由 Vite MPA 入口 `tuning.html → src/tuning/main.tsx` 启动,避免把 MuJoCo/Three.js 主应用依赖打入监控页面。两页仅通过训练 HTTP API和严格同源的短消息交接训练服务凭据/策略导入请求,不在 URL 中传 token。 + 测试文件使用 `*.test.ts(x)` 与被测模块共置;端到端测试统一保存在 `e2e/`。 diff --git a/web_platform/README.md b/web_platform/README.md index aa53bfe6..84e95d3b 100644 --- a/web_platform/README.md +++ b/web_platform/README.md @@ -15,7 +15,7 @@ - 内置平地、坡道、楼梯、可复现随机障碍物及 9 类系统参数化地形,可配置尺寸、摩擦、难度、种子与高度场采样精度 - 导入单文件 `.py` 控制器,通过本地 Pyodide 在 `mj_step` 前按仿真时间同步执行 - 导入 mjlab 导出的 `policy.onnx`,在浏览器本地执行 Go2-W 平衡/速度策略推理 -- 从图形界面向本机训练桥接服务发起 mjlab 强化学习训练、查看进度/日志、停止任务并导入训练生成的 ONNX +- 从图形界面向本机训练桥接服务发起 mjlab 强化学习训练、查看进度/日志、停止任务并导入训练生成的 ONNX;可在独立 TensorBoard 风格页面运行 DeepSeek 奖励函数自调参 - 可配置仿真遥测记录,实时查看速度、机身姿态、位置、驱动力等指标并导出 CSV/JSON - FPS、物理耗时和主线程步进预算提示 @@ -119,7 +119,9 @@ npm run training-server -- \ --trainer-python /path/to/training-env/bin/python ``` -服务启动时会在终端输出一个随机访问令牌;在界面中填写该令牌后连接。令牌仅保存在当前标签页的 `sessionStorage`。界面默认连接 `http://127.0.0.1:8765`,可选择服务端允许的任务、并行环境数、训练迭代、随机种子、CPU/GPU、GPU 编号和实验记录方式。W&B 默认为本地离线模式,无需登录或 API Key;也可完全禁用,只有明确选择在线模式时才会联网登录。训练期间页面轮询迭代进度与最近日志,可以停止任务;训练成功后点击“导入策略”,生成的 `policy.onnx` 会进入现有 ONNX 加载流程。 +服务启动时会在终端输出一个随机访问令牌;在界面中填写该令牌后连接。令牌仅保存在当前标签页的 `sessionStorage`。界面默认连接 `http://127.0.0.1:8765`,可选择服务端允许的任务、并行环境数、训练迭代、随机种子、CPU/GPU、GPU 编号和实验记录方式。W&B 默认为本地离线模式,无需登录或 API Key;也可完全禁用,只有明确选择在线模式时才会联网登录。训练期间页面轮询迭代进度与最近日志,可以停止任务;训练成功后点击“导入策略”,生成的 `policy.onnx` 会进入现有 ONNX 加载流程。普通训练还可以选择自调参产生的命名 reward preset,而不会改写仓库默认配置。 + +连接服务后点击“打开自调参 Agent 工作台”会打开独立 `tuning.html`。该页面提供自动/逐轮审批模式、目标权重与预算配置、TensorBoard scalar 筛选/平滑/缩放、trial/rung 排行、固定评估指标、Agent 决策时间线、参数 patch 修改审批、暂停/恢复/停止、最佳 preset JSON 和 ONNX 导出。新标签页 URL 不包含 token;同源 opener 会一次性交接凭据,直接打开页面时也可手工输入。DeepSeek key 始终由本地 Python 服务的 `DEEPSEEK_API_KEY` 环境变量读取,浏览器不会接触该 key。 桥接服务只监听本机回环地址,并检查 Host、Origin 和 Bearer Token;仅接受允许列表中的任务和经过范围校验的参数,不执行前端提供的 Shell 命令;一次只运行一个训练进程。默认任务使用仓库内置的 Go2 机器人资产与环境配置,**不会自动把浏览器中临时编辑的 MJCF/URDF 作为训练环境**。自定义浏览器模型训练需要在兼容的外部训练工程中注册 task,并通过服务的 `--trainer-root` 指定该工程。服务配置、接口和安全边界见 [`../training_server/README.md`](../training_server/README.md)。 diff --git a/web_platform/e2e/app.spec.ts b/web_platform/e2e/app.spec.ts index 3c067464..87551024 100644 --- a/web_platform/e2e/app.spec.ts +++ b/web_platform/e2e/app.spec.ts @@ -77,6 +77,14 @@ const LARGE_MODEL = ` `; +test('独立自调参工作台不需要加载 MuJoCo 主应用即可打开', async ({ page }) => { + await page.goto('/tuning.html'); + await expect(page.getByRole('heading', { name: 'Go2 奖励函数自调参 Agent' })).toBeVisible(); + await expect(page.getByText('新建 Unitree-Go2-Flat 调参 Session')).toBeVisible(); + await expect(page.getByLabel('访问令牌(仅当前标签页)')).toBeVisible(); + await expect(page.getByRole('button', { name: '启动自调参' })).toBeVisible(); +}); + test('显示中文平台骨架并加载单文件模型', async ({ page }) => { page.on('console', (message) => console.log(`[browser:${message.type()}] ${message.text()}`)); page.on('pageerror', (error) => console.log(`[browser:error] ${error.message}`)); diff --git a/web_platform/src/test/setup.ts b/web_platform/src/test/setup.ts index bb02c60c..447900aa 100644 --- a/web_platform/src/test/setup.ts +++ b/web_platform/src/test/setup.ts @@ -1 +1,22 @@ import '@testing-library/jest-dom/vitest'; + +Object.defineProperty(window, 'matchMedia', { + writable: true, + value: (query: string) => ({ + matches: false, + media: query, + onchange: null, + addListener: () => undefined, + removeListener: () => undefined, + addEventListener: () => undefined, + removeEventListener: () => undefined, + dispatchEvent: () => false, + }), +}); + +if (!globalThis.ResizeObserver) + globalThis.ResizeObserver = class ResizeObserver { + observe(): void {} + unobserve(): void {} + disconnect(): void {} + }; diff --git a/web_platform/src/training/LocalTrainingClient.test.ts b/web_platform/src/training/LocalTrainingClient.test.ts index 9126c140..ceb3d94a 100644 --- a/web_platform/src/training/LocalTrainingClient.test.ts +++ b/web_platform/src/training/LocalTrainingClient.test.ts @@ -51,6 +51,81 @@ describe('LocalTrainingClient', () => { ).rejects.toThrow('已有训练任务正在运行'); }); + it('调用调参 session、metrics 与审批接口且不泄露令牌到 URL', async () => { + const fetchMock = vi.fn().mockImplementation(() => + Promise.resolve( + new Response(JSON.stringify({ id: 'b'.repeat(32), state: 'running', series: [] }), { + status: 200, + headers: { 'Content-Type': 'application/json' }, + }), + ), + ); + vi.stubGlobal('fetch', fetchMock); + const client = new LocalTrainingClient('http://127.0.0.1:8765', 'deep-secret'); + await client.tuningMetrics('a'.repeat(32), 'b'.repeat(32), ['Evaluation/score'], 500); + await client.decideProposal('a'.repeat(32), 'c'.repeat(32), 'approve', { + feedback: 'ok', + patch: { weights: { pose: 1.1 }, params: {} }, + }); + expect(fetchMock.mock.calls[0][0]).toContain('/api/tuning/sessions/'); + expect(fetchMock.mock.calls[0][0]).toContain('maxPoints=500'); + expect(fetchMock.mock.calls[0][0]).not.toContain('deep-secret'); + const approval = fetchMock.mock.calls[1][1] as RequestInit; + expect(approval.method).toBe('POST'); + expect(new Headers(approval.headers).get('Authorization')).toBe('Bearer deep-secret'); + }); + + it('覆盖调参生命周期、preset 与最佳策略下载客户端方法', async () => { + const fetchMock = vi.fn().mockImplementation((input: string) => { + if (String(input).endsWith('/policy.onnx')) + return Promise.resolve(new Response(new Uint8Array([1, 2, 3]), { status: 200 })); + return Promise.resolve( + new Response( + JSON.stringify({ sessions: [], presets: [], configured: true, id: 'a'.repeat(32) }), + { + status: 200, + headers: { 'Content-Type': 'application/json' }, + }, + ), + ); + }); + vi.stubGlobal('fetch', fetchMock); + const client = new LocalTrainingClient('http://127.0.0.1:8765', 'secret-token'); + await client.tuningCapability(); + await client.testTuningAgent(); + await client.tuningSessions(); + await client.startTuning({ + taskId: 'Unitree-Go2-Flat', + mode: 'automatic', + runName: 'test', + numEnvs: 16, + seed: 42, + gpuIds: [0], + trialCount: 4, + initialIterations: 1, + middleIterations: 2, + finalIterations: 3, + evalNumEnvs: 8, + evalSteps: 10, + objectiveWeights: { + velocity_tracking: 0.35, + action_smoothness: 0.2, + posture_stability: 0.15, + fall_avoidance: 0.15, + foot_slip: 0.1, + energy: 0.05, + }, + fallbackEnabled: false, + }); + await client.tuningSession('a'.repeat(32)); + await client.tuningAction('a'.repeat(32), 'pause'); + await client.cancelTuning('a'.repeat(32)); + await client.presets(); + const policy = await client.downloadBestPolicy('a'.repeat(32)); + expect(policy.size).toBe(3); + expect(fetchMock).toHaveBeenCalledTimes(9); + }); + it('拒绝非 HTTP 地址和空访问令牌', () => { expect(() => new LocalTrainingClient('file:///tmp/socket', 'secret-token')).toThrow( 'http 或 https', diff --git a/web_platform/src/training/LocalTrainingClient.ts b/web_platform/src/training/LocalTrainingClient.ts index 65173846..3e4b38d0 100644 --- a/web_platform/src/training/LocalTrainingClient.ts +++ b/web_platform/src/training/LocalTrainingClient.ts @@ -1,4 +1,13 @@ -import type { TrainingJob, TrainingRequest, TrainingServerInfo } from './types'; +import type { + RewardPreset, + TuningCapability, + TuningCreateRequest, + TuningMetricsResponse, + TuningSession, + TrainingJob, + TrainingRequest, + TrainingServerInfo, +} from './types'; function normalizeEndpoint(value: string): string { const endpoint = value.trim().replace(/\/+$/, ''); @@ -61,12 +70,85 @@ export class LocalTrainingClient { return this.json(`/api/training/jobs/${encodeURIComponent(id)}`, { method: 'DELETE' }); } async downloadPolicy(id: string): Promise { - const response = await fetch( - `${this.endpoint}/api/training/jobs/${encodeURIComponent(id)}/artifacts/policy.onnx`, - this.requestInit(), + return this.download( + `/api/training/jobs/${encodeURIComponent(id)}/artifacts/policy.onnx`, + `policy-${id.slice(0, 8)}.onnx`, ); + } + tuningCapability(): Promise { + return this.json('/api/tuning/capabilities'); + } + testTuningAgent(): Promise<{ ok: boolean; model: string; outputType: string }> { + return this.json('/api/tuning/agent/test', { method: 'POST' }); + } + tuningSessions(): Promise { + return this.json<{ sessions: TuningSession[] }>('/api/tuning/sessions').then( + (value) => value.sessions, + ); + } + startTuning(request: TuningCreateRequest): Promise { + return this.json('/api/tuning/sessions', { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify(request), + }); + } + tuningSession(id: string): Promise { + return this.json(`/api/tuning/sessions/${encodeURIComponent(id)}`); + } + tuningAction(id: string, action: 'pause' | 'resume'): Promise { + return this.json(`/api/tuning/sessions/${encodeURIComponent(id)}/${action}`, { + method: 'POST', + }); + } + cancelTuning(id: string): Promise { + return this.json(`/api/tuning/sessions/${encodeURIComponent(id)}`, { method: 'DELETE' }); + } + decideProposal( + sessionId: string, + proposalId: string, + action: 'approve' | 'reject', + payload: { + feedback?: string; + patch?: { weights: Record; params: Record }; + }, + ): Promise { + return this.json( + `/api/tuning/sessions/${encodeURIComponent(sessionId)}/proposals/${encodeURIComponent(proposalId)}/${action}`, + { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify(payload), + }, + ); + } + tuningMetrics( + sessionId: string, + trialId: string, + tags: string[] = [], + maxPoints = 1000, + ): Promise { + const query = new URLSearchParams({ maxPoints: String(maxPoints) }); + if (tags.length) query.set('tags', tags.join(',')); + return this.json( + `/api/tuning/sessions/${encodeURIComponent(sessionId)}/trials/${encodeURIComponent(trialId)}/metrics?${query}`, + ); + } + presets(): Promise { + return this.json<{ presets: RewardPreset[] }>('/api/tuning/presets').then( + (value) => value.presets, + ); + } + downloadBestPolicy(sessionId: string): Promise { + return this.download( + `/api/tuning/sessions/${encodeURIComponent(sessionId)}/artifacts/best/policy.onnx`, + `best-policy-${sessionId.slice(0, 8)}.onnx`, + ); + } + private async download(path: string, name: string): Promise { + const response = await fetch(`${this.endpoint}${path}`, this.requestInit()); if (!response.ok) throw await responseError(response); const blob = await response.blob(); - return new File([blob], `policy-${id.slice(0, 8)}.onnx`, { type: 'application/octet-stream' }); + return new File([blob], name, { type: 'application/octet-stream' }); } } diff --git a/web_platform/src/training/LocalTrainingPanel.test.tsx b/web_platform/src/training/LocalTrainingPanel.test.tsx index 68076277..3c497a2f 100644 --- a/web_platform/src/training/LocalTrainingPanel.test.tsx +++ b/web_platform/src/training/LocalTrainingPanel.test.tsx @@ -37,6 +37,12 @@ describe('LocalTrainingPanel', () => { headers: { 'Content-Type': 'application/json' }, }), ) + .mockResolvedValueOnce( + new Response(JSON.stringify({ presets: [] }), { + status: 200, + headers: { 'Content-Type': 'application/json' }, + }), + ) .mockResolvedValueOnce( new Response(JSON.stringify(job), { status: 202, @@ -49,6 +55,12 @@ describe('LocalTrainingPanel', () => { headers: { 'Content-Type': 'application/json' }, }), ) + .mockResolvedValueOnce( + new Response(JSON.stringify({ presets: [] }), { + status: 200, + headers: { 'Content-Type': 'application/json' }, + }), + ) .mockResolvedValueOnce( new Response(JSON.stringify({ error: '训练任务不存在或服务已重启' }), { status: 404, @@ -62,10 +74,15 @@ describe('LocalTrainingPanel', () => { }); fireEvent.click(screen.getByRole('button', { name: '连接' })); expect(await screen.findByText('/opt/unitree_rl_mjlab')).toBeInTheDocument(); + const open = vi.spyOn(window, 'open').mockImplementation(() => null); + fireEvent.click(screen.getByRole('button', { name: '打开自调参 Agent 工作台' })); + expect(open).toHaveBeenCalledWith(expect.any(URL), '_blank'); + expect(String(open.mock.calls[0][0])).toContain('tuning.html'); + expect(String(open.mock.calls[0][0])).not.toContain('secret-token'); fireEvent.change(screen.getByLabelText('并行环境'), { target: { value: '32' } }); fireEvent.click(screen.getByRole('button', { name: '发起本地训练' })); - await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(2)); - const request = fetchMock.mock.calls[1][1] as RequestInit; + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(3)); + const request = fetchMock.mock.calls[2][1] as RequestInit; expect(JSON.parse(String(request.body))).toMatchObject({ taskId: 'Unitree-Go2-Flat', numEnvs: 32, @@ -80,7 +97,7 @@ describe('LocalTrainingPanel', () => { expect(tokenInput).toBeEnabled(); fireEvent.change(tokenInput, { target: { value: 'new-secret-token' } }); fireEvent.click(screen.getByRole('button', { name: '连接' })); - await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(4)); + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(6)); expect(await screen.findByRole('button', { name: '发起本地训练' })).toBeInTheDocument(); expect(sessionStorage.getItem('mujoco-local-training-token')).toBe('new-secret-token'); }); diff --git a/web_platform/src/training/LocalTrainingPanel.tsx b/web_platform/src/training/LocalTrainingPanel.tsx index 43766efc..653a489d 100644 --- a/web_platform/src/training/LocalTrainingPanel.tsx +++ b/web_platform/src/training/LocalTrainingPanel.tsx @@ -1,28 +1,25 @@ import { useEffect, useState, type ReactNode } from 'react'; -import { Download, Link, Play, Server, Square } from 'lucide-react'; +import { Download, ExternalLink, Link, Play, Server, Square } from 'lucide-react'; import { Badge, Button, ProgressBar, PropertyRow, Select } from '../components/ui'; import { LocalTrainingClient } from './LocalTrainingClient'; -import type { TrainingDevice, TrainingJob, TrainingServerInfo, WandbMode } from './types'; +import type { + RewardPreset, + TrainingDevice, + TrainingJob, + TrainingServerInfo, + WandbMode, +} from './types'; +import { + DEFAULT_TRAINING_ENDPOINT, + localStored, + rememberTrainingConnection, + sessionStored, + TRAINING_ENDPOINT_KEY, + TRAINING_JOB_KEY, + TRAINING_TOKEN_KEY, +} from './storage'; -const ENDPOINT_KEY = 'mujoco-local-training-endpoint', - JOB_KEY = 'mujoco-local-training-job', - TOKEN_KEY = 'mujoco-local-training-token'; -const DEFAULT_ENDPOINT = 'http://127.0.0.1:8765'; const ACTIVE_STATES = new Set(['queued', 'running']); -function stored(key: string, fallback = ''): string { - try { - return localStorage.getItem(key) ?? fallback; - } catch { - return fallback; - } -} -function sessionStored(key: string): string { - try { - return sessionStorage.getItem(key) ?? ''; - } catch { - return ''; - } -} function errorText(error: unknown): string { return error instanceof Error ? error.message : String(error); } @@ -37,10 +34,14 @@ function stateLabel(state: TrainingJob['state']): string { } export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File): void }) { - const [endpoint, setEndpoint] = useState(() => stored(ENDPOINT_KEY, DEFAULT_ENDPOINT)); - const [token, setToken] = useState(() => sessionStored(TOKEN_KEY)); + const [endpoint, setEndpoint] = useState(() => + localStored(TRAINING_ENDPOINT_KEY, DEFAULT_TRAINING_ENDPOINT), + ); + const [token, setToken] = useState(() => sessionStored(TRAINING_TOKEN_KEY)); const [server, setServer] = useState(); const [job, setJob] = useState(); + const [presets, setPresets] = useState([]); + const [rewardPresetId, setRewardPresetId] = useState(''); const [busy, setBusy] = useState(false), [error, setError] = useState(); const [taskId, setTaskId] = useState('Unitree-Go2-Flat'), @@ -60,26 +61,30 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File info = await client.health(); setServer(info); try { - localStorage.setItem(ENDPOINT_KEY, client.endpoint); - sessionStorage.setItem(TOKEN_KEY, client.token); + setPresets(await client.presets()); + } catch { + setPresets([]); + } + try { + rememberTrainingConnection(client.endpoint, client.token); } catch { /* 当前会话仍可连接 */ } if (info.tasks.length && !info.tasks.includes(taskId)) setTaskId(info.tasks[0]); - const remembered = info.activeJobId ?? stored(JOB_KEY); + const remembered = info.activeJobId ?? localStored(TRAINING_JOB_KEY); if (remembered) { try { const recovered = await client.job(remembered); setJob(recovered); try { - localStorage.setItem(JOB_KEY, recovered.id); + localStorage.setItem(TRAINING_JOB_KEY, recovered.id); } catch { /* ignore */ } } catch { setJob(undefined); try { - localStorage.removeItem(JOB_KEY); + localStorage.removeItem(TRAINING_JOB_KEY); } catch { /* ignore */ } @@ -116,6 +121,37 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File }; }, [endpoint, jobId, jobState, token]); + useEffect(() => { + const receive = (event: MessageEvent) => { + if ( + event.origin !== window.location.origin || + !event.source || + typeof event.data !== 'object' + ) + return; + const data = event.data as { type?: string; sessionId?: string }; + if (data.type === 'mujoco-tuning-ready') { + (event.source as Window).postMessage( + { type: 'mujoco-tuning-credentials', endpoint, token }, + event.origin, + ); + } + if (data.type === 'mujoco-tuning-import-policy' && data.sessionId) { + void new LocalTrainingClient(endpoint, token) + .downloadBestPolicy(data.sessionId) + .then(onPolicyReady) + .catch((value: unknown) => setError(errorText(value))); + } + }; + window.addEventListener('message', receive); + return () => window.removeEventListener('message', receive); + }, [endpoint, onPolicyReady, token]); + + const openTuningDashboard = () => { + rememberTrainingConnection(endpoint, token); + window.open(new URL('tuning.html', document.baseURI), '_blank'); + }; + const start = async () => { setBusy(true); setError(undefined); @@ -138,10 +174,11 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File device, gpuIds: ids, wandbMode, + rewardPresetId: rewardPresetId || undefined, }); setJob(next); try { - localStorage.setItem(JOB_KEY, next.id); + localStorage.setItem(TRAINING_JOB_KEY, next.id); } catch { /* ignore */ } @@ -217,6 +254,15 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File {server?.ready ? '可用' : '离线'} + {server?.ready && ( + + )} {server?.ready && !job && (
@@ -286,6 +332,21 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File />
+ + + setEndpoint(event.target.value)} + /> + + + setToken(event.target.value)} + /> + + + + + {!session ? ( + { + const next = await client().tuningSession(id); + setSession(next); + setSelectedTrialId(next.currentTrialId ?? next.bestTrialId ?? next.trials.at(-1)?.id); + }} + /> + ) : ( +
+ + +
+
+ + setTagFilter(event.target.value)} + placeholder="Episode_Reward / Evaluation" + /> + + + setSmoothing(Number(event.target.value))} + /> + + + {selectedTrial + ? `Trial ${selectedTrial.number} / rung ${selectedTrial.rung}` + : '请选择 trial'} + +
+ + {selectedTrial?.evaluation && } + +
+ + +
+ )} + {error && ( +
+ {error} +
+ )} + + ); +} + +function NewSessionForm(props: { + mode: TuningMode; + setMode(value: TuningMode): void; + runName: string; + setRunName(value: string): void; + numEnvs: number; + setNumEnvs(value: number): void; + trialCount: number; + setTrialCount(value: number): void; + gpuIds: string; + setGpuIds(value: string): void; + fallback: boolean; + setFallback(value: boolean): void; + objectives: ObjectiveWeights; + setObjectives(value: ObjectiveWeights): void; + start(): Promise; + busy: boolean; + sessions: TuningSession[]; + open(id: string): Promise; +}) { + return ( +
+
+

新建 Unitree-Go2-Flat 调参 Session

+
+ + props.setRunName(e.target.value)} + /> + + + + + + + + props.setGpuIds(e.target.value)} + /> + + +
+

目标权重

+
+ {(Object.keys(props.objectives) as (keyof ObjectiveWeights)[]).map((key) => ( + + ))} +
+

a + b, 0) - 1) < 1e-6 ? 'text-success' : 'text-danger'}`} + > + 总和:{Math.round(Object.values(props.objectives).reduce((a, b) => a + b, 0) * 100)}% +

+ +
+
+

历史 Sessions

+
+ {props.sessions.map((item) => ( + + ))} +
+
+
+ ); +} + +function EvaluationCard({ trial }: { trial: TuningTrial }) { + return ( +
+

固定协议评估

+
+ {Object.entries(trial.evaluation?.metrics ?? {}).map(([name, value]) => ( +
+

+ {name} +

+

{value.toFixed(5)}

+
+ ))} +
+
+ ); +} +function Leaderboard({ trials, select }: { trials: TuningTrial[]; select(id: string): void }) { + const ranked = [...trials] + .filter((trial) => trial.score !== undefined && trial.score !== null) + .sort((a, b) => (b.score ?? -999) - (a.score ?? -999)); + return ( +
+

排行榜

+
+ + + + + + + + + + + + {ranked.map((trial, index) => ( + select(trial.id)} + > + + + + + + + ))} + +
排名TrialRung分数安全门槛
{index + 1}T{trial.number}{trial.rung}{trial.score?.toFixed(5)}{trial.eligible ? '通过' : '未通过'}
+
+
+ ); +} +function ApprovalCard(props: { + proposal: TuningProposal; + patchText: string; + setPatchText(value: string): void; + feedback: string; + setFeedback(value: string): void; + decide(proposal: TuningProposal, action: 'approve' | 'reject'): Promise; + busy: boolean; +}) { + return ( +
+
+

等待审批

+ 置信度 {Math.round(props.proposal.confidence * 100)}% +
+

{props.proposal.rationale}

+