import * as ort from 'onnxruntime-web/wasm'; import { GO2W_VELOCITY_TASK, clampGo2wCommand } from '../tasks/go2wVelocity'; import type { PolicyDeployment } from '../deployment'; import { GO2_OBSTACLE_AVOIDANCE_TASK } from '../tasks/go2ObstacleAvoidance'; import type { NavigationStatus, RLCommand, RLPolicyStatus } from '../types'; ort.env.wasm.numThreads = 1; ort.env.wasm.proxy = false; export interface PolicyRuntimeBindings { observe(time: number, lastAction: Float32Array, command: RLCommand): Float32Array; apply(action: Float32Array): void; clear(): void; reset?(time: number): void; navigationStatus?(): NavigationStatus; setNavigationTarget?(target: [number, number]): void; resetNavigationTarget?(): void; } function message(error: unknown): string { return error instanceof Error ? error.message : String(error); } /** * ONNX Runtime Web 的 run() 是异步 API。物理循环会在每个 mj_step 前持续施加最近一次 * 完成的动作,并按控制频率启动下一次推理,避免阻塞 MuJoCo 的同步步进循环。 */ export class OnnxPolicyRuntime { private enabled = false; private disposed = false; private inFlight = false; private nextInferenceTime = 0; private action = new Float32Array(GO2W_VELOCITY_TASK.actionSize); private commandValue: RLCommand = { linearX: 0, linearY: 0, angularZ: 0 }; private inferenceCount = 0; private lastInferenceMs = 0; private error?: string; private epoch = 0; private runPromise?: Promise; private constructor( private readonly session: ort.InferenceSession, private readonly bindings: PolicyRuntimeBindings, private readonly path: string, private readonly inputName: string, private readonly outputName: string, private readonly task: { id: string; name: string; observationSize: number; actionSize: number; controlHz: number; }, ) {} static async load( model: Uint8Array, path: string, bindings: PolicyRuntimeBindings, deployment?: PolicyDeployment, ): Promise { const task = deployment?.taskId === GO2_OBSTACLE_AVOIDANCE_TASK.id ? { ...GO2_OBSTACLE_AVOIDANCE_TASK, observationSize: deployment.observationSize } : GO2W_VELOCITY_TASK; const session = await ort.InferenceSession.create(model.slice(), { executionProviders: ['wasm'], graphOptimizationLevel: 'all', }); 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 !== task.observationSize) throw new Error(`策略观测维度不匹配:模型 ${fixedInput},任务 ${task.observationSize}`); if (typeof fixedOutput === 'number' && fixedOutput > 0 && fixedOutput !== task.actionSize) throw new Error(`策略动作维度不匹配:模型 ${fixedOutput},任务 ${task.actionSize}`); if (deployment && (fixedInput !== task.observationSize || fixedOutput !== task.actionSize)) throw new Error('部署策略必须声明固定的观测/动作特征维度'); return new OnnxPolicyRuntime( session, bindings, path, session.inputNames[0], session.outputNames[0], task, ); } catch (error) { await session.release(); throw error; } } status(): RLPolicyStatus { return { taskId: this.task.id, taskName: this.task.name, path: this.path, loaded: !this.disposed, enabled: this.enabled, controlHz: this.task.controlHz, observationSize: this.task.observationSize, actionSize: this.task.actionSize, inputName: this.inputName, outputName: this.outputName, command: { ...this.commandValue }, inferenceCount: this.inferenceCount, lastInferenceMs: this.lastInferenceMs, navigation: this.navigationStatus(), error: this.error, }; } navigationStatus(): NavigationStatus | undefined { return this.disposed ? undefined : this.bindings.navigationStatus?.(); } setNavigationTarget(target: [number, number]): void { if (!this.disposed) this.bindings.setNavigationTarget?.(target); } resetNavigationTarget(): void { if (!this.disposed) this.bindings.resetNavigationTarget?.(); } setCommand(command: RLCommand): void { this.commandValue = clampGo2wCommand(command); } 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(); this.bindings.reset?.(time); } step(time: number): void { if (!this.enabled || this.disposed) return; this.bindings.apply(this.action); if (this.inFlight || time + 1e-9 < this.nextInferenceTime) return; let observation: Float32Array; try { observation = this.bindings.observe(time, this.action, this.commandValue); if (observation.length !== this.task.observationSize || !observation.every(Number.isFinite)) throw new Error('策略观测维度错误或包含非有限数'); } catch (error) { this.fail(error); return; } this.inFlight = true; this.nextInferenceTime = time + 1 / this.task.controlHz; const started = performance.now(), epoch = this.epoch; const input = new ort.Tensor('float32', observation, [1, observation.length]); this.runPromise = this.session .run({ [this.inputName]: input }) .then((outputs) => { try { const output = outputs[this.outputName]; if (!output || output.type !== 'float32') throw new Error(`找不到 float32 输出:${this.outputName}`); if (output.data.length !== this.task.actionSize) throw new Error( `策略动作维度错误:期望 ${this.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)), ); } }