From 6ac255e85926cd2014c1434dec2d116ba6a834da Mon Sep 17 00:00:00 2001 From: cen617-code <1057290604@qq.com> Date: Wed, 26 Aug 2026 13:36:30 +0800 Subject: [PATCH] feat(web-platform): release V0.4.2 RL Module --- wasm/package-lock.json | 129 +++++++++++++++++- wasm/package.json | 1 + wasm/web_platform/README.md | 26 ++++ wasm/web_platform/src/app/App.tsx | 23 ++-- .../src/app/components/RLPolicyPanel.tsx | 30 ++++ .../src/app/components/SidebarPanel.tsx | 6 +- .../src/rl/runtime/Go2wPolicyBindings.ts | 79 +++++++++++ .../src/rl/runtime/OnnxPolicyRuntime.ts | 83 +++++++++++ .../src/rl/tasks/go2wVelocity.test.ts | 28 ++++ .../web_platform/src/rl/tasks/go2wVelocity.ts | 57 ++++++++ wasm/web_platform/src/rl/types.ts | 30 ++++ .../src/simulation/PhysicsAdapter.ts | 9 ++ .../src/simulation/SimulationSession.ts | 34 ++++- 13 files changed, 518 insertions(+), 17 deletions(-) create mode 100644 wasm/web_platform/src/app/components/RLPolicyPanel.tsx create mode 100644 wasm/web_platform/src/rl/runtime/Go2wPolicyBindings.ts create mode 100644 wasm/web_platform/src/rl/runtime/OnnxPolicyRuntime.ts create mode 100644 wasm/web_platform/src/rl/tasks/go2wVelocity.test.ts create mode 100644 wasm/web_platform/src/rl/tasks/go2wVelocity.ts create mode 100644 wasm/web_platform/src/rl/types.ts diff --git a/wasm/package-lock.json b/wasm/package-lock.json index f9ec0335..e05aceee 100644 --- a/wasm/package-lock.json +++ b/wasm/package-lock.json @@ -14,6 +14,7 @@ "fflate": "^0.8.3", "lucide-react": "^0.555.0", "monaco-editor": "^0.55.1", + "onnxruntime-web": "^1.29.0", "pyodide": "^0.29.4", "react": "^19.2.8", "react-dom": "^19.2.8", @@ -1056,6 +1057,63 @@ "node": ">=20" } }, + "node_modules/@protobufjs/aspromise": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/@protobufjs/aspromise/-/aspromise-1.1.2.tgz", + "integrity": "sha512-j+gKExEuLmKwvz3OgROXtrJ2UG2x8Ch2YZUxahh+s1F2HZ+wAceUNLkvy6zKCPVRkU++ZWQrdxsUeQXmcg4uoQ==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/base64": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/@protobufjs/base64/-/base64-1.1.2.tgz", + "integrity": "sha512-AZkcAA5vnN/v4PDqKyMR5lx7hZttPDgClv83E//FMNhR2TMcLUhfRUBHCmSl0oi9zMgDDqRUJkSxO3wm85+XLg==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/codegen": { + "version": "2.0.5", + "resolved": "https://registry.npmjs.org/@protobufjs/codegen/-/codegen-2.0.5.tgz", + "integrity": "sha512-zgXFLzW3Ap33e6d0Wlj4MGIm6Ce8O89n/apUaGNB/jx+hw+ruWEp7EwGUshdLKVRCxZW12fp9r40E1mQrf/34g==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/eventemitter": { + "version": "1.1.1", + "resolved": "https://registry.npmjs.org/@protobufjs/eventemitter/-/eventemitter-1.1.1.tgz", + "integrity": "sha512-vW1GmwMZNnL+gMRaovlh9yZX74kc+TTU3FObkkurpMaRtBfLP3ldjS9KQWlwZgraRE0+dheEEoAxdzcJQ8eXZg==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/fetch": { + "version": "1.1.1", + "resolved": "https://registry.npmjs.org/@protobufjs/fetch/-/fetch-1.1.1.tgz", + "integrity": "sha512-GpptLrs57adMSuHi3VNj0mAF8dwh36LMaYF6XyJ6JMWlVsc+t42tm1HSEDmOs3A8fC9yyeisgLhsTVQokOZ0zw==", + "license": "BSD-3-Clause", + "dependencies": { + "@protobufjs/aspromise": "^1.1.1" + } + }, + "node_modules/@protobufjs/float": { + "version": "1.0.2", + "resolved": "https://registry.npmjs.org/@protobufjs/float/-/float-1.0.2.tgz", + "integrity": "sha512-Ddb+kVXlXst9d+R9PfTIxh1EdNkgoRe5tOX6t01f1lYWOvJnSPDBlG241QLzcyPdoNTsblLUdujGSE4RzrTZGQ==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/path": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/@protobufjs/path/-/path-1.1.2.tgz", + "integrity": "sha512-6JOcJ5Tm08dOHAbdR3GrvP+yUUfkjG5ePsHYczMFLq3ZmMkAD98cDgcT2iA1lJ9NVwFd4tH/iSSoe44YWkltEA==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/pool": { + "version": "1.1.0", + "resolved": "https://registry.npmjs.org/@protobufjs/pool/-/pool-1.1.0.tgz", + "integrity": "sha512-0kELaGSIDBKvcgS4zkjz1PeddatrjYcmMWOlAuAPwAeccUrPHdUqo/J6LiymHHEiJT5NrF1UVwxY14f+fy4WQw==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/utf8": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/@protobufjs/utf8/-/utf8-1.1.2.tgz", + "integrity": "sha512-b1UQwcEZ4yCnMCD8DAL1VlbvBJE9/IX4FTIp7BG1xYpf29SLazLSrqUkj4w7Y5y7cCVP6E5tcqqcI0xemPkHug==", + "license": "BSD-3-Clause" + }, "node_modules/@rolldown/binding-android-arm64": { "version": "1.0.3", "resolved": "https://registry.npmjs.org/@rolldown/binding-android-arm64/-/binding-android-arm64-1.0.3.tgz", @@ -1535,7 +1593,6 @@ "version": "24.9.2", "resolved": "https://registry.npmjs.org/@types/node/-/node-24.9.2.tgz", "integrity": "sha512-uWN8YqxXxqFMX2RqGOrumsKeti4LlmIMIyV0lgut4jx7KQBcBiW6vkDtIBvHnHIquwNfJhk8v2OtmO8zXWHfPA==", - "dev": true, "dependencies": { "undici-types": "~7.16.0" } @@ -3048,6 +3105,12 @@ "node": ">=16" } }, + "node_modules/flatbuffers": { + "version": "25.9.23", + "resolved": "https://registry.npmjs.org/flatbuffers/-/flatbuffers-25.9.23.tgz", + "integrity": "sha512-MI1qs7Lo4Syw0EOzUl0xjs2lsoeqFku44KpngfIduHBYvzm8h2+7K8YMQh1JtVVVrUvhLpNwqVi4DERegUJhPQ==", + "license": "Apache-2.0" + }, "node_modules/flatted": { "version": "3.4.4", "resolved": "https://registry.npmjs.org/flatted/-/flatted-3.4.4.tgz", @@ -3168,6 +3231,12 @@ "url": "https://github.com/sponsors/sindresorhus" } }, + "node_modules/guid-typescript": { + "version": "1.0.9", + "resolved": "https://registry.npmjs.org/guid-typescript/-/guid-typescript-1.0.9.tgz", + "integrity": "sha512-Y8T4vYhEfwJOTbouREvG+3XDsjr8E3kIr7uf+JZ0BYloFsttiHU0WfvANVsR7TxNUJa/WpCnw/Ino/p+DeBhBQ==", + "license": "ISC" + }, "node_modules/hasown": { "version": "2.0.4", "resolved": "https://registry.npmjs.org/hasown/-/hasown-2.0.4.tgz", @@ -3807,6 +3876,12 @@ "url": "https://github.com/sponsors/sindresorhus" } }, + "node_modules/long": { + "version": "5.3.2", + "resolved": "https://registry.npmjs.org/long/-/long-5.3.2.tgz", + "integrity": "sha512-mNAgZ1GmyNhD7AuqnTG3/VQ26o760+ZYBPKjPvugO8+nLbYfX6TVpJPseBvopbdY+qpZ/lKUnmEc1LeZYS3QAA==", + "license": "Apache-2.0" + }, "node_modules/lru-cache": { "version": "10.4.3", "resolved": "https://registry.npmjs.org/lru-cache/-/lru-cache-10.4.3.tgz", @@ -4056,6 +4131,26 @@ "node": ">=12.20.0" } }, + "node_modules/onnxruntime-common": { + "version": "1.29.0", + "resolved": "https://registry.npmjs.org/onnxruntime-common/-/onnxruntime-common-1.29.0.tgz", + "integrity": "sha512-/F63/e2VJoaVXGGNu6S5QH7jivBThGO95OzAVXXQ8hTta/b1QxI8udHa6cI3+3mAb5WWIIaMMwfZw01oivjJ1g==", + "license": "MIT" + }, + "node_modules/onnxruntime-web": { + "version": "1.29.0", + "resolved": "https://registry.npmjs.org/onnxruntime-web/-/onnxruntime-web-1.29.0.tgz", + "integrity": "sha512-LuQlpX6MFLJZu756erwUeb1mNfoJGbs1kzDwJGNlf5RvfYMdqhcY3vNpDPK40CUV2HoWTkIj+uS0o36GFHjeYw==", + "license": "MIT", + "dependencies": { + "flatbuffers": "^25.1.24", + "guid-typescript": "^1.0.9", + "long": "^5.2.3", + "onnxruntime-common": "1.29.0", + "platform": "^1.3.6", + "protobufjs": "^7.2.4" + } + }, "node_modules/optionator": { "version": "0.9.4", "resolved": "https://registry.npmjs.org/optionator/-/optionator-0.9.4.tgz", @@ -4204,6 +4299,12 @@ "node": ">= 6" } }, + "node_modules/platform": { + "version": "1.3.6", + "resolved": "https://registry.npmjs.org/platform/-/platform-1.3.6.tgz", + "integrity": "sha512-fnWVljUchTro6RiCFvCXBbNhJc2NijN7oIQxbwsyL0buWJPG85v81ehlHI9fXrJsMNgTofEoWIQeClKpgxFLrg==", + "license": "MIT" + }, "node_modules/playwright": { "version": "1.62.1", "resolved": "https://registry.npmjs.org/playwright/-/playwright-1.62.1.tgz", @@ -4458,6 +4559,29 @@ "url": "https://github.com/chalk/ansi-styles?sponsor=1" } }, + "node_modules/protobufjs": { + "version": "7.6.5", + "resolved": "https://registry.npmjs.org/protobufjs/-/protobufjs-7.6.5.tgz", + "integrity": "sha512-/FPD0nUc9jH6rfFjji9IBqOz4pcSE3CsT1m7Ep6Mdb0LxSUMj8hgl6GomOvZzpNpAqqGaXA0P3VSrZLFzIhQrw==", + "hasInstallScript": true, + "license": "BSD-3-Clause", + "dependencies": { + "@protobufjs/aspromise": "^1.1.2", + "@protobufjs/base64": "^1.1.2", + "@protobufjs/codegen": "^2.0.5", + "@protobufjs/eventemitter": "^1.1.1", + "@protobufjs/fetch": "^1.1.1", + "@protobufjs/float": "^1.0.2", + "@protobufjs/path": "^1.1.2", + "@protobufjs/pool": "^1.1.0", + "@protobufjs/utf8": "^1.1.1", + "@types/node": ">=13.7.0", + "long": "^5.3.2" + }, + "engines": { + "node": ">=12.0.0" + } + }, "node_modules/punycode": { "version": "2.3.1", "resolved": "https://registry.npmjs.org/punycode/-/punycode-2.3.1.tgz", @@ -5241,8 +5365,7 @@ "node_modules/undici-types": { "version": "7.16.0", "resolved": "https://registry.npmjs.org/undici-types/-/undici-types-7.16.0.tgz", - "integrity": "sha512-Zz+aZWSj8LE6zoxD+xrjh4VfkIG8Ya6LvYkZqtUQGJPZjYl53ypCaUwWqo7eI0x66KBGeRo+mlBEkMSeSZ38Nw==", - "dev": true + "integrity": "sha512-Zz+aZWSj8LE6zoxD+xrjh4VfkIG8Ya6LvYkZqtUQGJPZjYl53ypCaUwWqo7eI0x66KBGeRo+mlBEkMSeSZ38Nw==" }, "node_modules/update-browserslist-db": { "version": "1.3.1", diff --git a/wasm/package.json b/wasm/package.json index 6222b96a..f7ecee29 100644 --- a/wasm/package.json +++ b/wasm/package.json @@ -61,6 +61,7 @@ "fflate": "^0.8.3", "lucide-react": "^0.555.0", "monaco-editor": "^0.55.1", + "onnxruntime-web": "^1.29.0", "pyodide": "^0.29.4", "react": "^19.2.8", "react-dom": "^19.2.8", diff --git a/wasm/web_platform/README.md b/wasm/web_platform/README.md index 4e3935dd..51b2ff10 100644 --- a/wasm/web_platform/README.md +++ b/wasm/web_platform/README.md @@ -13,6 +13,7 @@ - 播放、暂停、单步、重置、0.25×–4× 速度 - actuator 滑杆、hinge/slide 关节拖动、动态 body 外力拖拽 - 导入单文件 `.py` 控制器,通过本地 Pyodide 在 `mj_step` 前按仿真时间同步执行 +- 导入 mjlab 导出的 `policy.onnx`,在浏览器本地执行 Go2-W 平衡/速度策略推理 - FPS、物理耗时和主线程步进预算提示 ## 开发 @@ -62,6 +63,30 @@ Python 控制器是可信的单文件脚本,必须同步定义 `step(ctx, stat 当前 Python 与 MuJoCo 都运行在主线程,以保证闭环调用严格位于 `mj_step` 前。仅运行可信脚本;死循环仍可能阻塞页面。Pyodide 及 Python 标准库由 npm 包随生产构建离线发布,不从 CDN 下载;暂不支持第三方 Python 包、`pip` 或多文件 import。 +## ONNX 强化学习策略 + +平台只负责策略推理,训练仍在 Python/mjlab 中完成。当前内置任务兼容 `unitree_rl_mjlab` Go2 velocity 的部署观测顺序: + +```text +base_ang_vel(3) + projected_gravity(3) + velocity_command(3) ++ gait_phase(2) + joint_pos_rel(12) + joint_vel_rel(12) + last_action(12) += 47 维观测 +``` + +策略必须具有一个 `float32` 输入和至少一个 `float32` 输出,输入末维为 47、输出末维为 12。动作按 `FL、FR、RL、RR` 的 hip/thigh/calf 顺序解释,转换为 `default_joint_pos + 0.25 * action` 的关节目标。对于 motor 模型,平台使用与部署配置一致的 kp/kd 执行位置 PD;对于 position actuator,直接写入目标位置。Go2-W 的四个轮电机在此首版腿式策略中保持零力矩。 + +使用步骤: + +1. 导入浮动基座 Go2-W MJCF/URDF,确保腿部关节与 actuator 使用 Unitree 标准命名; +2. 打开右侧“控制 → ONNX 强化学习策略”; +3. 从工程中选择或单独导入 `policy.onnx`; +4. 加载策略,设置前向、侧向和偏航速度,启用策略后播放仿真。 + +ONNX Runtime Web 的推理接口是异步的。物理循环会在每个 `mj_step` 前持续施加最近一次已完成的动作,并以 50 Hz 提交新观测;界面会显示推理耗时和次数。ONNX 与 Python 控制器互斥,启用其中一个会停止另一个。 + +> [!IMPORTANT] +> 当前内置契约是参考 mjlab Go2 的 12 腿关节策略,不是包含四个轮电机动作的 16 自由度 Go2-W 专用策略。若训练 Go2-W 轮式策略,需要后续同时扩展训练端部署配置和浏览器任务清单,确保观测、动作及归一化完全一致。 + ## 示例 `fixtures/` 包含(用于测试和手工验收,不会打进生产构建): @@ -78,6 +103,7 @@ Python 控制器是可信的单文件脚本,必须同步定义 `step(ctx, stat - 仅面向桌面版 Chrome、Edge、Firefox;未适配手机和平板。 - 物理运行在主线程、单线程 WASM。超出每帧预算时限制追帧并提示。 +- ONNX Runtime Web 当前使用单线程 WASM;策略必须将观测归一化包含在导出的 ONNX 图内,平台不会额外加载训练 checkpoint 的运行均值。 - 不支持 Xacro、账号或云端保存;Python 控制器暂不支持第三方包和不可信代码隔离。 - 关节拖动只支持 hinge/slide;ball/free joint 只读。 - MuJoCo WASM 本身不支持 DAE mesh。平台会移除 DAE visual,并以 collision 几何显示;DAE collision 会替换为半径 0.05 m 的占位球体并在界面警告。高精度仿真应先将 DAE 转为 OBJ/STL 或改为 URDF primitive。 diff --git a/wasm/web_platform/src/app/App.tsx b/wasm/web_platform/src/app/App.tsx index 653ccb32..68a620f8 100644 --- a/wasm/web_platform/src/app/App.tsx +++ b/wasm/web_platform/src/app/App.tsx @@ -2,11 +2,12 @@ /* eslint-disable react-hooks/exhaustive-deps */ import {useCallback,useEffect,useRef,useState,type ChangeEvent,type DragEvent} from 'react'; import {Camera,ChevronLeft,ChevronRight,CircleHelp,Code2,Crosshair,Download,Hand,Maximize,MousePointer2,PanelsTopLeft,Pause,Play,RotateCcw,Settings as SettingsIcon,SunMoon} from 'lucide-react'; -import type {ProjectManifest} from '../project/types'; +import {DEFAULT_IMPORT_LIMITS,type ProjectManifest} from '../project/types'; import {filesFromDrop,importBrowserFiles,normalizeProjectPath,ProjectImportError} from '../project/importer'; import {MainThreadPhysicsAdapter,type UrdfBaseMode,type UrdfEnhancementOptions,type UrdfLoadMode} from '../simulation/PhysicsAdapter'; import type {ActuatorParameters} from '../simulation/SimulationSession'; import type {ControllerCommand,ControllerStatus} from '../controller/types'; +import type {RLCommand,RLPolicyStatus} from '../rl/types'; import {MuJoCoViewer,type InteractionMode,type ViewerTheme} from '../viewer/MuJoCoViewer'; import {useAppStore,type AppDiagnostic} from '../stores/useAppStore'; import {WorkbenchHeader} from './components/WorkbenchHeader'; @@ -37,24 +38,24 @@ function urdfLinkNames(project:ProjectManifest|null,path:string|undefined):strin export function App(){ const state=useAppStore(); const manifest=useRef(null),notificationId=useRef(0),loadInFlight=useRef(false),importInFlight=useRef(false),adapter=useRef(new MainThreadPhysicsAdapter()),root=useRef(null),viewerHost=useRef(null),viewer=useRef(null),urdfEnhancementsRef=useRef({addActuators:true,addSensors:true,sensorType:'camera'}); - const [forceScale,setForceScale]=useState(50),[leftOpen,setLeftOpen]=useState(true),[rightOpen,setRightOpen]=useState(true),[helpOpen,setHelpOpen]=useState(false),[commandOpen,setCommandOpen]=useState(false),[sourceOpen,setSourceOpen]=useState(false),[generatedMjcf,setGeneratedMjcf]=useState(),[generatedMjcfPath,setGeneratedMjcfPath]=useState(),[pendingUrdfPath,setPendingUrdfPath]=useState(),[pendingUrdfMounts,setPendingUrdfMounts]=useState([]),[removeConfirmOpen,setRemoveConfirmOpen]=useState(false),[fullscreen,setFullscreen]=useState(false),[settingsOpen,setSettingsOpen]=useState(false),[layoutOpen,setLayoutOpen]=useState(false),[diagnosticsOpen,setDiagnosticsOpen]=useState(false),[importProgress,setImportProgress]=useState(),[notifications,setNotifications]=useState([]),[toast,setToast]=useState(),[selectedControllerPath,setSelectedControllerPath]=useState(),[controllerStatus,setControllerStatus]=useState(); + const [forceScale,setForceScale]=useState(50),[leftOpen,setLeftOpen]=useState(true),[rightOpen,setRightOpen]=useState(true),[helpOpen,setHelpOpen]=useState(false),[commandOpen,setCommandOpen]=useState(false),[sourceOpen,setSourceOpen]=useState(false),[generatedMjcf,setGeneratedMjcf]=useState(),[generatedMjcfPath,setGeneratedMjcfPath]=useState(),[pendingUrdfPath,setPendingUrdfPath]=useState(),[pendingUrdfMounts,setPendingUrdfMounts]=useState([]),[removeConfirmOpen,setRemoveConfirmOpen]=useState(false),[fullscreen,setFullscreen]=useState(false),[settingsOpen,setSettingsOpen]=useState(false),[layoutOpen,setLayoutOpen]=useState(false),[diagnosticsOpen,setDiagnosticsOpen]=useState(false),[importProgress,setImportProgress]=useState(),[notifications,setNotifications]=useState([]),[toast,setToast]=useState(),[selectedControllerPath,setSelectedControllerPath]=useState(),[controllerStatus,setControllerStatus]=useState(),[selectedPolicyPath,setSelectedPolicyPath]=useState(),[policyStatus,setPolicyStatus]=useState(); const [urdfMode,setUrdfMode]=useState('mjcf'),urdfModeRef=useRef('mjcf'); const [baseMode,setBaseMode]=useState('floating'),baseModeRef=useRef('floating'); const [showCollision,setShowCollision]=useState(false),[showSensorCamera,setShowSensorCamera]=useState(true),[theme,setTheme]=useState(initialTheme),[jointAdvanced,setJointAdvanced]=useState(false),[ignoreJointLimits,setIgnoreJointLimits]=useState(false),[angleUnit,setAngleUnit]=useState<'rad'|'deg'>('rad'); - useEffect(()=>{if(!viewerHost.current)return;viewer.current=new MuJoCoViewer(viewerHost.current,{onSelection:state.setSelection,onFrame:(frame,fps,snapshot)=>{const memory=(performance as Performance&{memory?:{usedJSHeapSize:number}}).memory?.usedJSHeapSize;state.setMetrics(fps,frame.stepMs,memory===undefined?undefined:memory/1048576,frame.overBudget);if(snapshot){state.setSnapshot(snapshot);setControllerStatus(snapshot.controller);if(snapshot.controller?.error)state.setPaused(true);}},onError:error=>state.setDiagnostic(diagnostic(error.message.includes('控制器')?'仿真':'渲染',error))});return()=>{viewer.current?.dispose();viewer.current=null;adapter.current.dispose();};},[]); + useEffect(()=>{if(!viewerHost.current)return;viewer.current=new MuJoCoViewer(viewerHost.current,{onSelection:state.setSelection,onFrame:(frame,fps,snapshot)=>{const memory=(performance as Performance&{memory?:{usedJSHeapSize:number}}).memory?.usedJSHeapSize;state.setMetrics(fps,frame.stepMs,memory===undefined?undefined:memory/1048576,frame.overBudget);if(snapshot){state.setSnapshot(snapshot);setControllerStatus(snapshot.controller);setPolicyStatus(snapshot.rlPolicy);if(snapshot.controller?.error||snapshot.rlPolicy?.error){adapter.current.setPaused(true);state.setPaused(true);}}},onError:error=>state.setDiagnostic(diagnostic(error.message.includes('控制器')?'仿真':'渲染',error))});return()=>{viewer.current?.dispose();viewer.current=null;adapter.current.dispose();};},[]); useEffect(()=>{viewer.current?.setMode(state.mode);},[state.mode]); useEffect(()=>{if(viewer.current)viewer.current.forceScale=forceScale;},[forceScale]); useEffect(()=>{viewer.current?.setShowCollision(showCollision);},[showCollision]); useEffect(()=>{viewer.current?.setShowSensorCamera(showSensorCamera);},[showSensorCamera]); useEffect(()=>{viewer.current?.setTheme(theme);document.documentElement.style.colorScheme=theme;try{localStorage.setItem('mujoco-platform-theme',theme);}catch{/* 当前会话仍可切换 */}},[theme]); useEffect(()=>{const change=()=>setFullscreen(document.fullscreenElement===root.current);document.addEventListener('fullscreenchange',change);return()=>document.removeEventListener('fullscreenchange',change);},[]); - const loadEntry=useCallback(async(path:string,requestedMode?:UrdfLoadMode)=>{if(!manifest.current||loadInFlight.current)return;loadInFlight.current=true;setIgnoreJointLimits(false);setControllerStatus(undefined);state.setEntry(path);state.setLoading(true);setImportProgress({label:'初始化 WASM 与编译模型',value:.65});state.setDiagnostic(undefined);setGeneratedMjcf(undefined);setGeneratedMjcfPath(undefined);viewer.current?.attach(null);state.setSnapshot(undefined);state.setSelection(null);try{const snapshot=await adapter.current.load(manifest.current,path,requestedMode??urdfModeRef.current,baseModeRef.current,urdfEnhancementsRef.current);const supportFiles=adapter.current.cachedSupportFiles();if(supportFiles.length&&manifest.current){manifest.current=mergeCachedFiles(manifest.current,supportFiles);state.setProject(manifest.current.name,manifest.current.files.map(file=>({path:file.path,size:file.size})),manifest.current.entries,path);}setImportProgress({label:'创建视口场景',value:.92});adapter.current.setSpeed(useAppStore.getState().speed);state.setSnapshot(snapshot);state.setPaused(true);viewer.current?.attach(adapter.current.session);try{setGeneratedMjcf(new TextDecoder().decode(adapter.current.exportMjcf()));setGeneratedMjcfPath(convertedCachePath(path));}catch(error){console.warn('[MuJoCo] 无法生成源码预览',error);}const notice:WorkbenchNotification={id:++notificationId.current,title:snapshot.warnings.length?'URDF 兼容处理':'模型加载完成',detail:snapshot.warnings.length?snapshot.warnings.join('\n'):path,tone:snapshot.warnings.length?'warning':'success',at:Date.now()};setNotifications(items=>[notice,...items].slice(0,20));setToast(notice);}catch(error){state.setDiagnostic(diagnostic('模型编译',error,path));const notice:WorkbenchNotification={id:++notificationId.current,title:'模型编译失败',detail:error instanceof Error?error.message:String(error),tone:'danger',at:Date.now()};setNotifications(items=>[notice,...items].slice(0,20));setToast(notice);}finally{loadInFlight.current=false;setImportProgress(undefined);state.setLoading(false);}},[]); + const loadEntry=useCallback(async(path:string,requestedMode?:UrdfLoadMode)=>{if(!manifest.current||loadInFlight.current)return;loadInFlight.current=true;setIgnoreJointLimits(false);setControllerStatus(undefined);setPolicyStatus(undefined);state.setEntry(path);state.setLoading(true);setImportProgress({label:'初始化 WASM 与编译模型',value:.65});state.setDiagnostic(undefined);setGeneratedMjcf(undefined);setGeneratedMjcfPath(undefined);viewer.current?.attach(null);state.setSnapshot(undefined);state.setSelection(null);try{const snapshot=await adapter.current.load(manifest.current,path,requestedMode??urdfModeRef.current,baseModeRef.current,urdfEnhancementsRef.current);const supportFiles=adapter.current.cachedSupportFiles();if(supportFiles.length&&manifest.current){manifest.current=mergeCachedFiles(manifest.current,supportFiles);state.setProject(manifest.current.name,manifest.current.files.map(file=>({path:file.path,size:file.size})),manifest.current.entries,path);}setImportProgress({label:'创建视口场景',value:.92});adapter.current.setSpeed(useAppStore.getState().speed);state.setSnapshot(snapshot);state.setPaused(true);viewer.current?.attach(adapter.current.session);try{setGeneratedMjcf(new TextDecoder().decode(adapter.current.exportMjcf()));setGeneratedMjcfPath(convertedCachePath(path));}catch(error){console.warn('[MuJoCo] 无法生成源码预览',error);}const notice:WorkbenchNotification={id:++notificationId.current,title:snapshot.warnings.length?'URDF 兼容处理':'模型加载完成',detail:snapshot.warnings.length?snapshot.warnings.join('\n'):path,tone:snapshot.warnings.length?'warning':'success',at:Date.now()};setNotifications(items=>[notice,...items].slice(0,20));setToast(notice);}catch(error){state.setDiagnostic(diagnostic('模型编译',error,path));const notice:WorkbenchNotification={id:++notificationId.current,title:'模型编译失败',detail:error instanceof Error?error.message:String(error),tone:'danger',at:Date.now()};setNotifications(items=>[notice,...items].slice(0,20));setToast(notice);}finally{loadInFlight.current=false;setImportProgress(undefined);state.setLoading(false);}},[]); const requestLoadEntry=useCallback(async(path:string)=>{const entry=manifest.current?.entries.find(candidate=>candidate.path===path);if(entry?.format==='urdf'&&urdfModeRef.current==='mjcf'){setPendingUrdfMounts(urdfLinkNames(manifest.current,path));setPendingUrdfPath(path);return;}await loadEntry(path);},[loadEntry]); const confirmUrdfOptions=(options:UrdfEnhancementOptions)=>{const path=pendingUrdfPath;if(!path)return;urdfEnhancementsRef.current=options;setPendingUrdfPath(undefined);setPendingUrdfMounts([]);void loadEntry(path);}; const skipUrdfOptions=()=>confirmUrdfOptions({addActuators:false,addSensors:false,sensorType:'camera'}); - const ingest=useCallback(async(files:File[],lockOwned=false)=>{if(importInFlight.current&&!lockOwned)return;importInFlight.current=true;state.setLoading(true);setImportProgress({label:'读取工程文件',value:.12});try{const next=await importBrowserFiles(files);setImportProgress({label:'处理模型资源与入口',value:.38});manifest.current=next;setSelectedControllerPath(next.files.find(file=>/\.py$/i.test(file.path))?.path);state.setProject(next.name,next.files.map(({path,size})=>({path,size})),next.entries,next.selectedEntry);if(next.selectedEntry)await requestLoadEntry(next.selectedEntry);}catch(error){state.setDiagnostic(diagnostic(error instanceof ProjectImportError&&/ZIP/.test(error.message)?'ZIP':'导入',error,error instanceof ProjectImportError?error.path:undefined));const notice:WorkbenchNotification={id:++notificationId.current,title:'工程导入失败',detail:error instanceof Error?error.message:String(error),tone:'danger',at:Date.now()};setNotifications(items=>[notice,...items].slice(0,20));setToast(notice);}finally{importInFlight.current=false;setImportProgress(undefined);state.setLoading(false);}},[requestLoadEntry]); + const ingest=useCallback(async(files:File[],lockOwned=false)=>{if(importInFlight.current&&!lockOwned)return;importInFlight.current=true;state.setLoading(true);setImportProgress({label:'读取工程文件',value:.12});try{const next=await importBrowserFiles(files);setImportProgress({label:'处理模型资源与入口',value:.38});manifest.current=next;setSelectedControllerPath(next.files.find(file=>/\.py$/i.test(file.path))?.path);setSelectedPolicyPath(next.files.find(file=>/\.onnx$/i.test(file.path))?.path);state.setProject(next.name,next.files.map(({path,size})=>({path,size})),next.entries,next.selectedEntry);if(next.selectedEntry)await requestLoadEntry(next.selectedEntry);}catch(error){state.setDiagnostic(diagnostic(error instanceof ProjectImportError&&/ZIP/.test(error.message)?'ZIP':'导入',error,error instanceof ProjectImportError?error.path:undefined));const notice:WorkbenchNotification={id:++notificationId.current,title:'工程导入失败',detail:error instanceof Error?error.message:String(error),tone:'danger',at:Date.now()};setNotifications(items=>[notice,...items].slice(0,20));setToast(notice);}finally{importInFlight.current=false;setImportProgress(undefined);state.setLoading(false);}},[requestLoadEntry]); const removeProject=()=>{if(state.projectName)setRemoveConfirmOpen(true);}; - const confirmRemoveProject=()=>{viewer.current?.attach(null);adapter.current.dispose();manifest.current=null;setGeneratedMjcf(undefined);setGeneratedMjcfPath(undefined);setPendingUrdfPath(undefined);setPendingUrdfMounts([]);setSelectedControllerPath(undefined);setControllerStatus(undefined);state.clearProject();setRemoveConfirmOpen(false);}; + const confirmRemoveProject=()=>{viewer.current?.attach(null);adapter.current.dispose();manifest.current=null;setGeneratedMjcf(undefined);setGeneratedMjcfPath(undefined);setPendingUrdfPath(undefined);setPendingUrdfMounts([]);setSelectedControllerPath(undefined);setControllerStatus(undefined);setSelectedPolicyPath(undefined);setPolicyStatus(undefined);state.clearProject();setRemoveConfirmOpen(false);}; const changeUrdfMode=(value:UrdfLoadMode)=>{setUrdfMode(value);urdfModeRef.current=value;const entry=state.entries.find(candidate=>candidate.path===state.selectedEntry);if(entry?.format!=='urdf')return;if(value==='mjcf'){setPendingUrdfMounts(urdfLinkNames(manifest.current,entry.path));setPendingUrdfPath(entry.path);}else void loadEntry(entry.path,value);}; const changeBaseMode=(value:UrdfBaseMode)=>{setBaseMode(value);baseModeRef.current=value;const entry=state.entries.find(candidate=>candidate.path===state.selectedEntry);if(entry?.format==='urdf'&&urdfModeRef.current==='mjcf')void loadEntry(entry.path,'mjcf');}; const changeFiles=(event:ChangeEvent)=>{void ingest(Array.from(event.target.files??[]));event.target.value='';}; @@ -72,9 +73,15 @@ export function App(){ const loadControllerSource=async(source:string,path:string)=>{state.setLoading(true);setImportProgress({label:'初始化 Python 运行时并加载控制器',value:.5});state.setDiagnostic(undefined);try{const status=await adapter.current.loadPythonController(source,path);setControllerStatus(status);state.setSnapshot(adapter.current.snapshot()??undefined);notify('Python 控制器已加载',`${status.name} · ${status.controlHz} Hz`);}catch(error){state.setDiagnostic(diagnostic('仿真',error,path));}finally{setImportProgress(undefined);state.setLoading(false);}}; const loadControllerPath=(path:string)=>{const file=manifest.current?.files.find(candidate=>candidate.path===path);if(!file){state.setDiagnostic(diagnostic('仿真',new Error('工程中找不到控制脚本'),path));return;}setSelectedControllerPath(path);void loadControllerSource(new TextDecoder().decode(file.data),path);}; const importController=(file:File)=>{void (async()=>{try{if(!/\.py$/i.test(file.name))throw new Error('请选择 .py 文件');if(file.size>1024*1024)throw new Error('Python 控制脚本不能超过 1 MiB');const path=normalizeProjectPath(file.name),data=new Uint8Array(await file.arrayBuffer());if(manifest.current){const index=manifest.current.files.findIndex(candidate=>candidate.path===path),files=manifest.current.files.slice(),entry={path,data,size:data.byteLength,source:'file' as const,mimeType:file.type||'text/x-python'};if(index>=0)files[index]=entry;else files.push(entry);manifest.current={...manifest.current,files,totalBytes:files.reduce((total,item)=>total+item.size,0)};state.setProject(manifest.current.name,files.map(({path:filePath,size})=>({path:filePath,size})),manifest.current.entries,manifest.current.selectedEntry);state.setSnapshot(adapter.current.snapshot()??undefined);}setSelectedControllerPath(path);await loadControllerSource(new TextDecoder().decode(data),path);}catch(error){state.setDiagnostic(diagnostic('仿真',error,file.name));}})();}; - const toggleController=(enabled:boolean)=>{adapter.current.setControllerEnabled(enabled);const snapshot=adapter.current.snapshot()??undefined;setControllerStatus(snapshot?.controller);state.setSnapshot(snapshot);}; + const toggleController=(enabled:boolean)=>{adapter.current.setControllerEnabled(enabled);const snapshot=adapter.current.snapshot()??undefined;setControllerStatus(snapshot?.controller);setPolicyStatus(snapshot?.rlPolicy);state.setSnapshot(snapshot);}; const sendControllerCommand=(command:ControllerCommand)=>{try{adapter.current.sendControllerCommand(command);const snapshot=adapter.current.snapshot()??undefined;setControllerStatus(snapshot?.controller);state.setSnapshot(snapshot);}catch(error){state.setDiagnostic(diagnostic('仿真',error,selectedControllerPath));}}; const removeController=()=>{adapter.current.removeController();setControllerStatus(undefined);state.setSnapshot(adapter.current.snapshot()??undefined);}; + const loadPolicyBytes=async(data:Uint8Array,path:string)=>{state.setLoading(true);setImportProgress({label:'初始化 ONNX Runtime 并加载策略',value:.55});state.setDiagnostic(undefined);try{const status=await adapter.current.loadRLPolicy(data,path);setPolicyStatus(status);state.setSnapshot(adapter.current.snapshot()??undefined);notify('ONNX 策略已加载',`${status.taskName} · ${status.observationSize} → ${status.actionSize}`);}catch(error){state.setDiagnostic(diagnostic('仿真',error,path));}finally{setImportProgress(undefined);state.setLoading(false);}}; + const loadPolicyPath=(path:string)=>{const file=manifest.current?.files.find(candidate=>candidate.path===path);if(!file){state.setDiagnostic(diagnostic('仿真',new Error('工程中找不到 ONNX 策略'),path));return;}setSelectedPolicyPath(path);void loadPolicyBytes(file.data,path);}; + const importPolicy=(file:File)=>{void (async()=>{try{if(!/\.onnx$/i.test(file.name))throw new Error('请选择 .onnx 文件');if(file.size>64*1024*1024)throw new Error('ONNX 策略不能超过 64 MiB');const path=normalizeProjectPath(file.name),data=new Uint8Array(await file.arrayBuffer());if(manifest.current){const index=manifest.current.files.findIndex(candidate=>candidate.path===path),files=manifest.current.files.slice(),entry={path,data,size:data.byteLength,source:'file' as const,mimeType:file.type||'application/octet-stream'};if(index>=0)files[index]=entry;else files.push(entry);const totalBytes=files.reduce((total,item)=>total+item.size,0);if(totalBytes>DEFAULT_IMPORT_LIMITS.maxTotalBytes)throw new Error('加入 ONNX 后工程总大小超过 512 MiB');manifest.current={...manifest.current,files,totalBytes};state.setProject(manifest.current.name,files.map(({path:filePath,size})=>({path:filePath,size})),manifest.current.entries,manifest.current.selectedEntry);state.setSnapshot(adapter.current.snapshot()??undefined);}setSelectedPolicyPath(path);await loadPolicyBytes(data,path);}catch(error){state.setDiagnostic(diagnostic('仿真',error,file.name));}})();}; + const togglePolicy=(enabled:boolean)=>{adapter.current.setRLPolicyEnabled(enabled);const snapshot=adapter.current.snapshot()??undefined;setPolicyStatus(snapshot?.rlPolicy);setControllerStatus(snapshot?.controller);state.setSnapshot(snapshot);}; + const setPolicyCommand=(command:RLCommand)=>{adapter.current.setRLCommand(command);const snapshot=adapter.current.snapshot()??undefined;setPolicyStatus(snapshot?.rlPolicy);state.setSnapshot(snapshot);}; + const removePolicy=()=>{adapter.current.removeRLPolicy();setPolicyStatus(undefined);state.setSnapshot(adapter.current.snapshot()??undefined);}; const notify=(title:string,detail:string,tone:WorkbenchNotification['tone']='success')=>{const notice:WorkbenchNotification={id:++notificationId.current,title,detail,tone,at:Date.now()};setNotifications(items=>[notice,...items].slice(0,20));setToast(notice);}; const saveCachedSource=async(path:string,text:string)=>{if(!manifest.current)return;manifest.current=upsertCachedMjcf(manifest.current,path,text);state.setProject(manifest.current.name,manifest.current.files.map(file=>({path:file.path,size:file.size})),manifest.current.entries,path);notify('转换后的 MJCF 已保存到缓存',path);await loadEntry(path);}; const exportUrdf=()=>{if(!manifest.current||selectedFormat!=='urdf'||!state.selectedEntry)return;const text=readCachedText(manifest.current,state.selectedEntry);downloadBytes(new TextEncoder().encode(text),exportedFileName(manifest.current.name,'urdf'));notify('URDF 已导出',state.selectedEntry);}; @@ -101,7 +108,7 @@ export function App(){ ]; return
event.preventDefault()} onDrop={drop}> setSourceOpen(true)} onTogglePause={togglePause} onStep={singleStep} onReset={reset} onSpeed={changeSpeed} onToggleLeft={()=>setLeftOpen(value=>!value)} onToggleRight={()=>setRightOpen(value=>!value)} onToggleTheme={()=>setTheme(value=>value==='dark'?'light':'dark')} onHelp={()=>setHelpOpen(true)} endActions={<>setNotifications(items=>items.filter(item=>item.id!==id))} onClear={()=>setNotifications([])} onOpenLog={()=>setDiagnosticsOpen(true)}/>setLayoutOpen(true)}>setSettingsOpen(true)}>} compactMenu={setCommandOpen(true)} onLayout={()=>setLayoutOpen(true)} onSettings={()=>setSettingsOpen(true)} onFullscreen={toggleFullscreen} onHelp={()=>setHelpOpen(true)} onTheme={()=>setTheme(value=>value==='dark'?'light':'dark')}/>} onCommands={()=>setCommandOpen(true)} onToggleFullscreen={toggleFullscreen} center={viewer.current?.resetCamera()}/>}/> -
viewer.current?.highlightJoint(jointId)}/>
setToast(undefined)}/>{Boolean(state.snapshot?.model.ncam)&&(showSensorCamera?
摄像头
:)}{state.entries.length>1&&!state.selectedEntry&&!pendingUrdfPath&&} {state.diagnostic&&state.setDiagnostic(undefined)} onRetry={state.diagnostic.category==='模型编译'&&state.diagnostic.path?()=>void loadEntry(state.diagnostic!.path!):undefined} onOpenProject={()=>{setLeftOpen(true);state.setDiagnostic(undefined);}}/>}
/\.py$/i.test(file.path)).map(file=>file.path)} selectedControllerPath={selectedControllerPath} controllerStatus={controllerStatus} onUrdfMode={changeUrdfMode} onBaseMode={changeBaseMode} onShowCollision={setShowCollision} onResetJoints={resetJoints} onToggleJointLimits={toggleJointLimits} onToggleAdvanced={()=>setJointAdvanced(value=>!value)} onToggleAngleUnit={()=>setAngleUnit(value=>value==='rad'?'deg':'rad')} onActuator={setActuator} onActuatorParameters={setActuatorParameters} onJoint={setJoint} onForceScale={setForceScale} onSelectControllerPath={setSelectedControllerPath} onLoadControllerPath={loadControllerPath} onImportController={importController} onToggleController={toggleController} onControllerCommand={sendControllerCommand} onRemoveController={removeController}/>
+
viewer.current?.highlightJoint(jointId)}/>
setToast(undefined)}/>{Boolean(state.snapshot?.model.ncam)&&(showSensorCamera?
摄像头
:)}{state.entries.length>1&&!state.selectedEntry&&!pendingUrdfPath&&} {state.diagnostic&&state.setDiagnostic(undefined)} onRetry={state.diagnostic.category==='模型编译'&&state.diagnostic.path?()=>void loadEntry(state.diagnostic!.path!):undefined} onOpenProject={()=>{setLeftOpen(true);state.setDiagnostic(undefined);}}/>}
/\.py$/i.test(file.path)).map(file=>file.path)} selectedControllerPath={selectedControllerPath} controllerStatus={controllerStatus} policyPaths={state.files.filter(file=>/\.onnx$/i.test(file.path)).map(file=>file.path)} selectedPolicyPath={selectedPolicyPath} policyStatus={policyStatus} onUrdfMode={changeUrdfMode} onBaseMode={changeBaseMode} onShowCollision={setShowCollision} onResetJoints={resetJoints} onToggleJointLimits={toggleJointLimits} onToggleAdvanced={()=>setJointAdvanced(value=>!value)} onToggleAngleUnit={()=>setAngleUnit(value=>value==='rad'?'deg':'rad')} onActuator={setActuator} onActuatorParameters={setActuatorParameters} onJoint={setJoint} onForceScale={setForceScale} onSelectControllerPath={setSelectedControllerPath} onLoadControllerPath={loadControllerPath} onImportController={importController} onToggleController={toggleController} onControllerCommand={sendControllerCommand} onRemoveController={removeController} onSelectPolicyPath={setSelectedPolicyPath} onLoadPolicyPath={loadPolicyPath} onImportPolicy={importPolicy} onTogglePolicy={togglePolicy} onPolicyCommand={setPolicyCommand} onRemovePolicy={removePolicy}/>
{pendingUrdfPath&&}{sourceOpen&&generatedMjcf&&generatedMjcfPath&&setSourceOpen(false)} onSave={saveCachedSource}/>}setHelpOpen(false)}/>setDiagnosticsOpen(false)} onClear={()=>setNotifications([])}/>setSettingsOpen(false)} theme={theme} angleUnit={angleUnit} showCollision={showCollision} jointAdvanced={jointAdvanced} forceScale={forceScale} onTheme={setTheme} onAngleUnit={setAngleUnit} onShowCollision={setShowCollision} onJointAdvanced={setJointAdvanced} onForceScale={setForceScale}/>setLayoutOpen(false)} leftOpen={leftOpen} rightOpen={rightOpen} onLeftOpen={setLeftOpen} onRightOpen={setRightOpen} onPreset={applyLayoutPreset} onReset={()=>applyLayoutPreset('default')}/>setCommandOpen(false)} commands={commands}/>setRemoveConfirmOpen(false)}>

确定从当前会话中移除“{state.projectName}”吗?

该操作不会删除本地文件。

; } diff --git a/wasm/web_platform/src/app/components/RLPolicyPanel.tsx b/wasm/web_platform/src/app/components/RLPolicyPanel.tsx new file mode 100644 index 00000000..391f6899 --- /dev/null +++ b/wasm/web_platform/src/app/components/RLPolicyPanel.tsx @@ -0,0 +1,30 @@ +import {useRef,type ChangeEvent} from 'react'; +import {BrainCircuit,FileUp,Power,RotateCw,Trash2} from 'lucide-react'; +import type {RLCommand,RLPolicyStatus} from '../../rl/types'; +import {Badge,Button,PropertyRow,Select} from '../../components/ui'; + +export interface RLPolicyPanelProps { + paths:string[];selectedPath?:string;status?:RLPolicyStatus;loading:boolean; + onSelectPath(path:string):void;onLoadPath(path:string):void;onImport(file:File):void; + onToggle(enabled:boolean):void;onCommand(command:RLCommand):void;onRemove():void; +} + +export function RLPolicyPanel({paths,selectedPath,status,loading,onSelectPath,onLoadPath,onImport,onToggle,onCommand,onRemove}:RLPolicyPanelProps){ + const input=useRef(null); + const importFile=(event:ChangeEvent)=>{const file=event.target.files?.[0];if(file)onImport(file);event.target.value='';}; + const command=status?.command??{linearX:0,linearY:0,angularZ:0}; + return
+ + {paths.length>0&&} +
+ {status?
+
{status.taskName}{status.enabled?'推理中':'已停止'}
+ +

速度指令(机身坐标系)

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

{status.error}

} +
+
:

加载 mjlab 导出的单输入、单动作输出 policy.onnx。首个内置任务使用 47 维 Go2 actor 观测和 12 维腿部关节位置动作;Go2-W 轮电机保持零力矩。

} +
; +} + +function CommandInput({label,value,min,max,onChange}:{label:string;value:number;min:number;max:number;onChange(value:number):void}){return ;} diff --git a/wasm/web_platform/src/app/components/SidebarPanel.tsx b/wasm/web_platform/src/app/components/SidebarPanel.tsx index 1cc74aa0..67959bc5 100644 --- a/wasm/web_platform/src/app/components/SidebarPanel.tsx +++ b/wasm/web_platform/src/app/components/SidebarPanel.tsx @@ -7,10 +7,12 @@ import type {ActuatorInfo,ActuatorParameters,SimulationSnapshot} from '../../sim import type {UrdfBaseMode,UrdfLoadMode} from '../../simulation/PhysicsAdapter'; import type {ViewerSelection} from '../../viewer/MuJoCoViewer'; import type {ControllerCommand,ControllerStatus} from '../../controller/types'; +import type {RLCommand,RLPolicyStatus} from '../../rl/types'; import {Badge,Button,CollapsibleSection,CopyButton,PropertyRow,ResizablePanel,Select,Tabs} from '../../components/ui'; import {TreeSearchField} from './TreeSearchField'; import {ProjectBreadcrumb} from './ProjectBreadcrumb'; import {PythonControllerPanel} from './PythonControllerPanel'; +import {RLPolicyPanel} from './RLPolicyPanel'; export function SidebarPanel({title,side,children,visible=true}:{title:string;side:'left'|'right';children:ReactNode;visible?:boolean}){return ;} @@ -20,16 +22,18 @@ interface ModelControlsProps{ snapshot?:SimulationSnapshot;selection:ViewerSelection|null;selectedFormat?:ModelEntry['format'];loading:boolean;visible?:boolean; urdfMode:UrdfLoadMode;baseMode:UrdfBaseMode;showCollision:boolean;ignoreJointLimits:boolean;jointAdvanced:boolean;angleUnit:'rad'|'deg';forceScale:number; controllerPaths:string[];selectedControllerPath?:string;controllerStatus?:ControllerStatus; + policyPaths:string[];selectedPolicyPath?:string;policyStatus?:RLPolicyStatus; onUrdfMode:(value:UrdfLoadMode)=>void;onBaseMode:(value:UrdfBaseMode)=>void;onShowCollision:(value:boolean)=>void; onResetJoints:()=>void;onToggleJointLimits:()=>void;onToggleAdvanced:()=>void;onToggleAngleUnit:()=>void; onActuator:(id:number,value:number)=>void;onActuatorParameters:(id:number,parameters:ActuatorParameters)=>void;onJoint:(id:number,value:number)=>void;onForceScale:(value:number)=>void; onSelectControllerPath:(path:string)=>void;onLoadControllerPath:(path:string)=>void;onImportController:(file:File)=>void;onToggleController:(enabled:boolean)=>void;onControllerCommand:(command:ControllerCommand)=>void;onRemoveController:()=>void; + onSelectPolicyPath:(path:string)=>void;onLoadPolicyPath:(path:string)=>void;onImportPolicy:(file:File)=>void;onTogglePolicy:(enabled:boolean)=>void;onPolicyCommand:(command:RLCommand)=>void;onRemovePolicy:()=>void; } export function ModelControlsSidebar(props:ModelControlsProps){const [tab,setTab]=useState<'properties'|'controls'>('properties'),s=props.snapshot;if(!s)return
导入模型后显示属性
; const properties=<>{s.model.nbody} Body}>
{props.selectedFormat==='urdf'&&

MJCF 模式保留 visual mesh、添加物理地面,并将模型最低点对齐到 z=0。

} {props.selection?
}/>}/>value.toFixed(3)).join(', ')} action={}/>
:

在视口中单击物体

}
; - const controls=<>{s.controller.enabled?'运行':'停止'}:undefined}>{s.actuators.length}}>{s.actuators.length?s.actuators.map(actuator=>props.onActuator(actuator.id,value)} onParameters={parameters=>props.onActuatorParameters(actuator.id,parameters)}/>):

模型没有驱动器

}
+ const controls=<>{s.rlPolicy.enabled?'推理':'停止'}:undefined}>{s.controller.enabled?'运行':'停止'}:undefined}>{s.actuators.length}}>{s.actuators.length?s.actuators.map(actuator=>props.onActuator(actuator.id,value)} onParameters={parameters=>props.onActuatorParameters(actuator.id,parameters)}/>):

模型没有驱动器

}
{s.joints.length}}>
{s.joints.map(joint=>{const scale=joint.type===3&&props.angleUnit==='deg'?180/Math.PI:1,unit=joint.type===3?(props.angleUnit==='deg'?'°':' rad'):joint.type===2?' m':'';return props.onJoint(joint.id,value/scale)}/>;})}

选择“外力施加”,在动态物体上按住拖动,松开即清零。

; return ,content:properties},{value:'controls',label:'控制',icon:,content:controls}]}/>; diff --git a/wasm/web_platform/src/rl/runtime/Go2wPolicyBindings.ts b/wasm/web_platform/src/rl/runtime/Go2wPolicyBindings.ts new file mode 100644 index 00000000..a5d6a734 --- /dev/null +++ b/wasm/web_platform/src/rl/runtime/Go2wPolicyBindings.ts @@ -0,0 +1,79 @@ +import type {MjData,MjModel} from '@mujoco/mujoco'; +import {buildGo2wObservation,GO2W_VELOCITY_TASK} from '../tasks/go2wVelocity'; +import type {JointBinding,RLCommand} from '../types'; +import type {PolicyRuntimeBindings} from './OnnxPolicyRuntime'; + +interface BoundJoint extends JointBinding {positionActuator:boolean;controlScale:number;} + +function rotateInverse(quaternion:readonly number[],vector:readonly number[]):[number,number,number]{ + const [w,x,y,z]=quaternion,[vx,vy,vz]=vector; + const tx=2*(y*vz-z*vy),ty=2*(z*vx-x*vz),tz=2*(x*vy-y*vx); + return [vx-w*tx+(y*tz-z*ty),vy-w*ty+(z*tx-x*tz),vz-w*tz+(x*ty-y*tx)]; +} + +/** 将 mjlab Go2 velocity 的 47 维 actor 观测和 12 维关节位置动作映射到 MuJoCo。 */ +export class Go2wPolicyBindings implements PolicyRuntimeBindings { + private readonly joints:BoundJoint[]; + private readonly baseBodyId:number; + private readonly baseFreeJointId:number; + private readonly gyroSensorId?:number; + private readonly wheelActuatorIds:number[]; + + constructor(private readonly model:MjModel,private readonly data:MjData,private readonly setActuator:(id:number,value:number)=>void){ + const jointIds=new Map(),actuatorIds=new Map(),sensorIds=new Map(),bodyIds=new Map(); + for(let id=0;id{ + const jointId=jointIds.get(name);if(jointId===undefined)throw new Error(`Go2-W 策略找不到关节:${name}`); + const short=name.replace(/_joint$/,''); + const actuatorId=actuatorIds.get(short)??actuatorIds.get(`${name}_motor`); + if(actuatorId===undefined)throw new Error(`Go2-W 策略找不到驱动器:${short} 或 ${name}_motor`); + const joint=model.jnt(jointId),actuator=model.actuator(actuatorId); + try{ + const address=Number(model.actuator_ctrladr[actuatorId]??actuatorId),nextAddress=actuatorId+11e-5||Math.abs(gain-GO2W_VELOCITY_TASK.stiffness[index])>1e-4||Math.abs(Number(actuator.biasprm[2])+GO2W_VELOCITY_TASK.damping[index])>1e-4))throw new Error(`position 驱动器 ${actuator.name||actuatorId} 的 gear/kp/kd 与 mjlab deploy 配置不一致`); + const controlScale=gear*gain; + if(!Number.isFinite(controlScale)||Math.abs(controlScale)<1e-9)throw new Error(`驱动器 ${actuator.name||actuatorId} 的 gear × gain 无效`); + return {name,jointId,qposAddress:Number(joint.qposadr),qvelAddress:Number(joint.dofadr),actuatorId,positionActuator,controlScale}; + }finally{actuator.delete();joint.delete();} + }); + this.wheelActuatorIds=['FL','FR','RL','RR'].flatMap(prefix=>{ + const id=actuatorIds.get(`${prefix}_wheel`)??actuatorIds.get(`${prefix}_wheel_joint_motor`)??actuatorIds.get(`${prefix}_foot_joint_motor`); + return id===undefined?[]:[id]; + }); + } + + observe(time:number,lastAction:Float32Array,command:RLCommand):Float32Array{ + const quaternion=Array.from(this.data.xquat.subarray(this.baseBodyId*4,this.baseBodyId*4+4),Number); + const projectedGravity=rotateInverse(quaternion,[0,0,-1]); + let angularVelocity:[number,number,number]; + if(this.gyroSensorId!==undefined){const address=Number(this.model.sensor_adr[this.gyroSensorId]);angularVelocity=[Number(this.data.sensordata[address]),Number(this.data.sensordata[address+1]),Number(this.data.sensordata[address+2])];} + else {const joint=this.model.jnt(this.baseFreeJointId);try{const address=Number(joint.dofadr)+3;angularVelocity=[Number(this.data.qvel[address]),Number(this.data.qvel[address+1]),Number(this.data.qvel[address+2])];}finally{joint.delete();}} + return buildGo2wObservation({angularVelocity,projectedGravity,command,time,jointPosition:this.joints.map(item=>Number(this.data.qpos[item.qposAddress])),jointVelocity:this.joints.map(item=>Number(this.data.qvel[item.qvelAddress])),lastAction:Array.from(lastAction)}); + } + + apply(action:Float32Array):void{ + for(let index=0;index=this.model.nsite||Number(this.model.site_bodyid[siteId])!==this.baseBodyId)return false;const offset=siteId*4;return Math.abs(Number(this.model.site_quat[offset])-1)<1e-5&&Math.abs(Number(this.model.site_quat[offset+1]))<1e-5&&Math.abs(Number(this.model.site_quat[offset+2]))<1e-5&&Math.abs(Number(this.model.site_quat[offset+3]))<1e-5;} + private findFloatingBaseBody():number{for(let jointId=0;jointId; + + private constructor(private readonly session:ort.InferenceSession,private readonly bindings:PolicyRuntimeBindings,private readonly path:string,private readonly inputName:string,private readonly outputName:string){} + + static async load(model:Uint8Array,path:string,bindings:PolicyRuntimeBindings):Promise{ + const session=await ort.InferenceSession.create(model.slice(),{executionProviders:['wasm'],graphOptimizationLevel:'all'}); + try{ + if(session.inputNames.length!==1)throw new Error(`当前仅支持单输入策略,模型包含 ${session.inputNames.length} 个输入`); + if(session.outputNames.length<1)throw new Error('ONNX 策略没有输出'); + const input=session.inputMetadata[0],output=session.outputMetadata[0]; + if(!input?.isTensor||input.type!=='float32')throw new Error('策略输入必须是 float32 Tensor'); + if(!output?.isTensor||output.type!=='float32')throw new Error('策略输出必须是 float32 Tensor'); + if(input.shape.length!==2||output.shape.length!==2)throw new Error(`策略输入/输出必须是二维 [batch, features],实际为 [${input.shape}] / [${output.shape}]`); + const inputBatch=input.shape[0],outputBatch=output.shape[0],fixedInput=input.shape[1],fixedOutput=output.shape[1]; + if(typeof inputBatch==='number'&&inputBatch!==-1&&inputBatch!==1)throw new Error(`策略输入 batch 必须为 1 或动态维度,实际为 ${inputBatch}`); + if(typeof outputBatch==='number'&&outputBatch!==-1&&outputBatch!==1)throw new Error(`策略输出 batch 必须为 1 或动态维度,实际为 ${outputBatch}`); + if(typeof fixedInput==='number'&&fixedInput>0&&fixedInput!==GO2W_VELOCITY_TASK.observationSize)throw new Error(`策略观测维度不匹配:模型 ${fixedInput},任务 ${GO2W_VELOCITY_TASK.observationSize}`); + if(typeof fixedOutput==='number'&&fixedOutput>0&&fixedOutput!==GO2W_VELOCITY_TASK.actionSize)throw new Error(`策略动作维度不匹配:模型 ${fixedOutput},任务 ${GO2W_VELOCITY_TASK.actionSize}`); + return new OnnxPolicyRuntime(session,bindings,path,session.inputNames[0],session.outputNames[0]); + }catch(error){await session.release();throw error;} + } + + status():RLPolicyStatus{return {taskId:GO2W_VELOCITY_TASK.id,taskName:GO2W_VELOCITY_TASK.name,path:this.path,loaded:!this.disposed,enabled:this.enabled,controlHz:GO2W_VELOCITY_TASK.controlHz,observationSize:GO2W_VELOCITY_TASK.observationSize,actionSize:GO2W_VELOCITY_TASK.actionSize,inputName:this.inputName,outputName:this.outputName,command:{...this.commandValue},inferenceCount:this.inferenceCount,lastInferenceMs:this.lastInferenceMs,error:this.error};} + setCommand(command:RLCommand):void{this.commandValue=clampGo2wCommand(command);} + setEnabled(enabled:boolean,time:number):void{if(this.disposed)return;this.epoch+=1;this.enabled=enabled;this.error=undefined;this.nextInferenceTime=time;if(!enabled){this.action.fill(0);this.bindings.clear();}} + reset(time:number):void{this.epoch+=1;this.action.fill(0);this.nextInferenceTime=time;this.error=undefined;this.bindings.clear();} + + step(time:number):void{ + if(!this.enabled||this.disposed)return; + this.bindings.apply(this.action); + if(this.inFlight||time+1e-9{ + try{ + const output=outputs[this.outputName]; + if(!output||output.type!=='float32')throw new Error(`找不到 float32 输出:${this.outputName}`); + if(output.data.length!==GO2W_VELOCITY_TASK.actionSize)throw new Error(`策略动作维度错误:期望 ${GO2W_VELOCITY_TASK.actionSize},实际 ${output.data.length}`); + const next=Float32Array.from(output.data as Float32Array,Number); + for(const value of next)if(!Number.isFinite(value))throw new Error('策略输出包含非有限数'); + if(!this.disposed&&this.enabled&&epoch===this.epoch){this.action=next;this.inferenceCount+=1;this.lastInferenceMs=performance.now()-started;} + }finally{for(const value of Object.values(outputs))value.dispose();} + }).catch(error=>{if(epoch===this.epoch)this.fail(error);}).finally(()=>{input.dispose();this.inFlight=false;this.runPromise=undefined;}); + } + + private fail(error:unknown):void{if(this.disposed)return;this.error=message(error);this.enabled=false;this.bindings.clear();} + dispose():void{if(this.disposed)return;this.disposed=true;this.enabled=false;this.epoch+=1;this.bindings.clear();const pending=this.runPromise??Promise.resolve();void pending.catch(()=>{}).finally(()=>this.session.release().catch(error=>console.warn('[ONNX] 释放推理会话失败',error)));} +} diff --git a/wasm/web_platform/src/rl/tasks/go2wVelocity.test.ts b/wasm/web_platform/src/rl/tasks/go2wVelocity.test.ts new file mode 100644 index 00000000..6207eb67 --- /dev/null +++ b/wasm/web_platform/src/rl/tasks/go2wVelocity.test.ts @@ -0,0 +1,28 @@ +import {describe,expect,it} from 'vitest'; +import {buildGo2wObservation,clampGo2wCommand,go2wGaitPhase,GO2W_VELOCITY_TASK} from './go2wVelocity'; + +describe('Go2-W velocity task',()=>{ + it('按 mjlab deploy 顺序构造 47 维 actor 观测',()=>{ + const jointPosition=GO2W_VELOCITY_TASK.defaultJointPosition.map(value=>value+0.1); + const observation=buildGo2wObservation({angularVelocity:[1,2,3],projectedGravity:[0,0,-1],command:{linearX:0.5,linearY:-0.25,angularZ:0.2},time:0,jointPosition,jointVelocity:Array(12).fill(0.3),lastAction:Array(12).fill(-0.4)}); + expect(observation).toHaveLength(47); + [1,2,3,0,0,-1,0.5,-0.25,0.2,0,1].forEach((value,index)=>expect(observation[index]).toBeCloseTo(value)); + for(const value of observation.slice(11,23))expect(value).toBeCloseTo(0.1); + for(const value of observation.slice(23,35))expect(value).toBeCloseTo(0.3); + for(const value of observation.slice(35,47))expect(value).toBeCloseTo(-0.4); + }); + + it('静止时关闭步态相位,并限制速度命令范围',()=>{ + expect(go2wGaitPhase(0.15,{linearX:0,linearY:0,angularZ:0})).toEqual([0,0]); + const moving=go2wGaitPhase(0.15,{linearX:1,linearY:0,angularZ:0}); + expect(moving[0]).toBeCloseTo(1); + expect(moving[1]).toBeCloseTo(0); + expect(clampGo2wCommand({linearX:4,linearY:-4,angularZ:3})).toEqual({linearX:1,linearY:-0.5,angularZ:1}); + }); + + it('拒绝维度错误或非有限观测',()=>{ + const valid={angularVelocity:[0,0,0],projectedGravity:[0,0,-1],command:{linearX:0,linearY:0,angularZ:0},time:0,jointPosition:Array(12).fill(0),jointVelocity:Array(12).fill(0),lastAction:Array(12).fill(0)}; + expect(()=>buildGo2wObservation({...valid,lastAction:[0]})).toThrow(/观测维度/); + expect(()=>buildGo2wObservation({...valid,angularVelocity:[Number.NaN,0,0]})).toThrow(/非有限数/); + }); +}); diff --git a/wasm/web_platform/src/rl/tasks/go2wVelocity.ts b/wasm/web_platform/src/rl/tasks/go2wVelocity.ts new file mode 100644 index 00000000..54eef7fb --- /dev/null +++ b/wasm/web_platform/src/rl/tasks/go2wVelocity.ts @@ -0,0 +1,57 @@ +import type {RLCommand} from '../types'; + +export const GO2W_VELOCITY_TASK={ + id:'unitree-go2w-velocity' as const, + name:'Unitree Go2-W 平衡/速度控制', + controlHz:50, + gaitPeriod:0.6, + observationSize:47, + actionSize:12, + commandLimits:{linearX:[-0.5,1] as const,linearY:[-0.5,0.5] as const,angularZ:[-1,1] as const}, + jointNames:[ + 'FL_hip_joint','FL_thigh_joint','FL_calf_joint', + 'FR_hip_joint','FR_thigh_joint','FR_calf_joint', + 'RL_hip_joint','RL_thigh_joint','RL_calf_joint', + 'RR_hip_joint','RR_thigh_joint','RR_calf_joint', + ] as const, + defaultJointPosition:[-0.1,0.9,-1.8,0.1,0.9,-1.8,-0.1,0.9,-1.8,0.1,0.9,-1.8] as const, + actionScale:[0.25,0.25,0.25,0.25,0.25,0.25,0.25,0.25,0.25,0.25,0.25,0.25] as const, + stiffness:[20,20,40,20,20,40,20,20,40,20,20,40] as const, + damping:[1,1,2,1,1,2,1,1,2,1,1,2] as const, +}; + +export function clampGo2wCommand(command:RLCommand):RLCommand { + const limits=GO2W_VELOCITY_TASK.commandLimits; + const clamp=(value:number,range:readonly[number,number])=>Math.min(range[1],Math.max(range[0],Number.isFinite(value)?value:0)); + return {linearX:clamp(command.linearX,limits.linearX),linearY:clamp(command.linearY,limits.linearY),angularZ:clamp(command.angularZ,limits.angularZ)}; +} + +export function go2wGaitPhase(time:number,command:RLCommand):[number,number] { + if(Math.hypot(command.linearX,command.linearY,command.angularZ)<0.1)return [0,0]; + const phase=((time/GO2W_VELOCITY_TASK.gaitPeriod)%1+1)%1; + return [Math.sin(phase*2*Math.PI),Math.cos(phase*2*Math.PI)]; +} + +export function buildGo2wObservation(values:{ + angularVelocity:readonly number[]; + projectedGravity:readonly number[]; + command:RLCommand; + time:number; + jointPosition:readonly number[]; + jointVelocity:readonly number[]; + lastAction:readonly number[]; +}):Float32Array { + const phase=go2wGaitPhase(values.time,values.command); + const observation=new Float32Array([ + ...values.angularVelocity.slice(0,3), + ...values.projectedGravity.slice(0,3), + values.command.linearX,values.command.linearY,values.command.angularZ, + ...phase, + ...values.jointPosition.map((value,index)=>value-GO2W_VELOCITY_TASK.defaultJointPosition[index]), + ...values.jointVelocity, + ...values.lastAction, + ]); + if(observation.length!==GO2W_VELOCITY_TASK.observationSize)throw new Error(`Go2-W 观测维度错误:期望 ${GO2W_VELOCITY_TASK.observationSize},实际 ${observation.length}`); + for(const value of observation)if(!Number.isFinite(value))throw new Error('Go2-W 观测包含非有限数'); + return observation; +} diff --git a/wasm/web_platform/src/rl/types.ts b/wasm/web_platform/src/rl/types.ts new file mode 100644 index 00000000..d38dfeb7 --- /dev/null +++ b/wasm/web_platform/src/rl/types.ts @@ -0,0 +1,30 @@ +export interface RLCommand { + linearX:number; + linearY:number; + angularZ:number; +} + +export interface RLPolicyStatus { + taskId:'unitree-go2w-velocity'; + taskName:string; + path:string; + loaded:boolean; + enabled:boolean; + controlHz:number; + observationSize:number; + actionSize:number; + inputName:string; + outputName:string; + command:RLCommand; + inferenceCount:number; + lastInferenceMs:number; + error?:string; +} + +export interface JointBinding { + name:string; + jointId:number; + qposAddress:number; + qvelAddress:number; + actuatorId:number; +} diff --git a/wasm/web_platform/src/simulation/PhysicsAdapter.ts b/wasm/web_platform/src/simulation/PhysicsAdapter.ts index 22426a4d..8ab2a90f 100644 --- a/wasm/web_platform/src/simulation/PhysicsAdapter.ts +++ b/wasm/web_platform/src/simulation/PhysicsAdapter.ts @@ -5,6 +5,7 @@ import {enhanceConvertedMjcf,groundConvertedMjcf,type UrdfBaseMode,type UrdfEnha import {MemfsWorkspace} from '../project/workspace'; import {SimulationSession,type ActuatorParameters,type FrameResult,type SimulationSnapshot} from './SimulationSession'; import type {ControllerCommand,ControllerStatus} from '../controller/types'; +import type {RLCommand,RLPolicyStatus} from '../rl/types'; export type UrdfLoadMode='mjcf'|'native'; export type {UrdfBaseMode,UrdfEnhancementOptions}; @@ -28,6 +29,10 @@ export interface PhysicsAdapter { setControllerEnabled(enabled:boolean):void; sendControllerCommand(command:ControllerCommand):void; removeController():void; + loadRLPolicy(model:Uint8Array,path:string):Promise; + setRLPolicyEnabled(enabled:boolean):void; + setRLCommand(command:RLCommand):void; + removeRLPolicy():void; cachedSupportFiles():ProjectFile[]; exportMjcf(): Uint8Array; dispose(): void; @@ -94,6 +99,10 @@ export class MainThreadPhysicsAdapter implements PhysicsAdapter { setControllerEnabled(enabled:boolean):void{this.session?.setControllerEnabled(enabled);} sendControllerCommand(command:ControllerCommand):void{this.session?.sendControllerCommand(command);} removeController():void{this.session?.removeController();} + async loadRLPolicy(model:Uint8Array,path:string):Promise{if(!this.session)throw new Error('请先加载模型');return this.session.loadRLPolicy(model,path);} + setRLPolicyEnabled(enabled:boolean):void{this.session?.setRLPolicyEnabled(enabled);} + setRLCommand(command:RLCommand):void{this.session?.setRLCommand(command);} + removeRLPolicy():void{this.session?.removeRLPolicy();} cachedSupportFiles():ProjectFile[]{return this.supportFiles.map(file=>({...file,data:file.data.slice()}));} exportMjcf():Uint8Array{ if(!this.session||!this.workspace)throw new Error('尚未加载可导出的模型'); diff --git a/wasm/web_platform/src/simulation/SimulationSession.ts b/wasm/web_platform/src/simulation/SimulationSession.ts index cf23b0fc..17923c07 100644 --- a/wasm/web_platform/src/simulation/SimulationSession.ts +++ b/wasm/web_platform/src/simulation/SimulationSession.ts @@ -2,12 +2,15 @@ import type {MainModule, MjData, MjModel, MjvPerturb, MjvScene} from '@mujoco/mu import {meshIdFromSceneDataId} from './geometry'; import {PythonControllerRuntime} from '../controller/PythonControllerRuntime'; import type {ControllerBindings,ControllerCommand,ControllerStatus} from '../controller/types'; +import {Go2wPolicyBindings} from '../rl/runtime/Go2wPolicyBindings'; +import {OnnxPolicyRuntime} from '../rl/runtime/OnnxPolicyRuntime'; +import type {RLCommand,RLPolicyStatus} from '../rl/types'; export interface ActuatorParameters {gear:number;gain:number;kp:number;kv:number;ctrlLimited:boolean;ctrlMin:number;ctrlMax:number;forceLimited:boolean;forceMin:number;forceMax:number;} export interface ActuatorInfo extends ActuatorParameters {id:number;name:string;value:number;min:number;max:number;limited:boolean;jointId?:number;jointName?:string;jointType?:number;unit:string;kind:'motor'|'position'|'velocity'|'other';controlCount:number;} export interface JointInfo {id:number;name:string;type:number;value:number;min:number;max:number;limitMin:number;limitMax:number;limited:boolean;limitsIgnored:boolean;editable:boolean;bodyId:number;axis:[number,number,number];} export interface BodyInfo {id:number;name:string;parentId:number;} -export interface SimulationSnapshot {time: number; qpos: number[]; qvel: number[]; ctrl: number[]; actuators: ActuatorInfo[]; joints: JointInfo[]; bodies: BodyInfo[]; warnings: string[]; controller?:ControllerStatus; model:{nbody:number;njnt:number;ngeom:number;ncam:number;nactuator:number;nu:number;nq:number;nv:number};} +export interface SimulationSnapshot {time: number; qpos: number[]; qvel: number[]; ctrl: number[]; actuators: ActuatorInfo[]; joints: JointInfo[]; bodies: BodyInfo[]; warnings: string[]; controller?:ControllerStatus; rlPolicy?:RLPolicyStatus; model:{nbody:number;njnt:number;ngeom:number;ncam:number;nactuator:number;nu:number;nq:number;nv:number};} export interface FrameResult {steps: number; stepMs: number; overBudget: boolean;} export class SimulationSession { @@ -27,6 +30,8 @@ export class SimulationSession { private jointLimits:{limited:boolean;min:number;max:number;type:number}[]=[]; private pythonController?:PythonControllerRuntime; private controllerLoadGeneration=0; + private rlPolicy?:OnnxPolicyRuntime; + private rlPolicyLoadGeneration=0; constructor(readonly module: MainModule, modelPath: string, readonly warnings: string[] = []) { let model: MjModel | undefined; let data: MjData | undefined; let perturb: MjvPerturb | undefined; @@ -43,7 +48,7 @@ export class SimulationSession { setPaused(paused: boolean): void {this.paused = paused; this.accumulator = 0; this.lastNow = undefined;} setSpeed(speed: number): void {this.speed = Math.min(4, Math.max(0.1, speed));} - reset(): void {this.setPaused(true);this.module.mj_resetData(this.model,this.data);this.module.mj_forward(this.model,this.data);this.clearExternalForce();this.data.ctrl.fill(0);this.pythonController?.reset(Number(this.data.time));} + reset(): void {this.setPaused(true);this.module.mj_resetData(this.model,this.data);this.module.mj_forward(this.model,this.data);this.clearExternalForce();this.data.ctrl.fill(0);this.pythonController?.reset(Number(this.data.time));this.rlPolicy?.reset(Number(this.data.time));} singleStep(): void {this.runController();this.applyForce();this.module.mj_step(this.model,this.data);} advance(now: number): FrameResult { @@ -69,16 +74,35 @@ export class SimulationSession { } setControllerEnabled(enabled:boolean):void { + if(enabled&&this.pythonController){this.data.ctrl.fill(0);this.rlPolicy?.setEnabled(false,Number(this.data.time));} this.pythonController?.setEnabled(enabled,Number(this.data.time)); if(!enabled)this.data.ctrl.fill(0); } + async loadRLPolicy(model:Uint8Array,path:string):Promise{ + const generation=++this.rlPolicyLoadGeneration; + const bindings=new Go2wPolicyBindings(this.model,this.data,(id,value)=>this.setActuator(id,value)); + const runtime=await OnnxPolicyRuntime.load(model,path,bindings); + if(this.disposed||generation!==this.rlPolicyLoadGeneration){runtime.dispose();throw new Error('模型已切换,ONNX 策略加载已取消');} + this.data.ctrl.fill(0);this.rlPolicy?.dispose();this.rlPolicy=runtime; + return runtime.status(); + } + + setRLPolicyEnabled(enabled:boolean):void { + if(enabled&&this.rlPolicy){this.data.ctrl.fill(0);this.pythonController?.setEnabled(false,Number(this.data.time));} + this.rlPolicy?.setEnabled(enabled,Number(this.data.time)); + if(!enabled)this.data.ctrl.fill(0); + } + + setRLCommand(command:RLCommand):void {this.rlPolicy?.setCommand(command);} + removeRLPolicy():void {this.rlPolicyLoadGeneration+=1;this.rlPolicy?.dispose();this.rlPolicy=undefined;this.data.ctrl.fill(0);} + sendControllerCommand(command:ControllerCommand):void {this.pythonController?.command(command);} removeController():void {this.controllerLoadGeneration+=1;this.pythonController?.dispose();this.pythonController=undefined;this.data.ctrl.fill(0);} private runController():void { - try{this.pythonController?.stepIfDue(Number(this.data.time));} + try{this.pythonController?.stepIfDue(Number(this.data.time));this.rlPolicy?.step(Number(this.data.time));} catch(error){this.setPaused(true);this.data.ctrl.fill(0);throw error;} } @@ -224,7 +248,7 @@ export class SimulationSession { try {return {id,name:body.name||`body_${id}`,parentId:Number(this.model.body_parentid[id])};} finally { body.delete(); } }); - return {time:Number(this.data.time),qpos:Array.from(this.data.qpos),qvel:Array.from(this.data.qvel),ctrl:Array.from(this.data.ctrl),actuators,joints,bodies,warnings:this.warnings,controller:this.pythonController?.status(),model:{nbody:this.model.nbody,njnt:this.model.njnt,ngeom:this.model.ngeom,ncam:this.model.ncam,nactuator:this.model.nactuator,nu:this.model.nu,nq:this.model.nq,nv:this.model.nv}}; + return {time:Number(this.data.time),qpos:Array.from(this.data.qpos),qvel:Array.from(this.data.qvel),ctrl:Array.from(this.data.ctrl),actuators,joints,bodies,warnings:this.warnings,controller:this.pythonController?.status(),rlPolicy:this.rlPolicy?.status(),model:{nbody:this.model.nbody,njnt:this.model.njnt,ngeom:this.model.ngeom,ncam:this.model.ncam,nactuator:this.model.nactuator,nu:this.model.nu,nq:this.model.nq,nv:this.model.nv}}; } - dispose(): void {if(this.disposed)return;this.disposed=true;this.removeController();this.clearExternalForce();this.perturb.delete();this.data.delete();this.model.delete();} + dispose(): void {if(this.disposed)return;this.disposed=true;this.removeController();this.removeRLPolicy();this.clearExternalForce();this.perturb.delete();this.data.delete();this.model.delete();} }