Files
Mujoco_WASM/web_platform/src/rl/runtime/OnnxPolicyRuntime.ts
T
chenlin 438e56bcc8
web-platform-ci / TypeScript, lint, unit, build (push) Has been cancelled
web-platform-ci / Playwright E2E (push) Has been cancelled
feat(training): release V0.9.1 避障训练与基础策略迁移
2026-09-08 10:50:13 +08:00

234 lines
8.5 KiB
TypeScript

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<void>;
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<OnnxPolicyRuntime> {
const task =
deployment?.taskId === GO2_OBSTACLE_AVOIDANCE_TASK.id
? { ...GO2_OBSTACLE_AVOIDANCE_TASK, observationSize: deployment.observationSize }
: GO2W_VELOCITY_TASK;
const session = await ort.InferenceSession.create(model.slice(), {
executionProviders: ['wasm'],
graphOptimizationLevel: 'all',
});
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)),
);
}
}