234 lines
8.5 KiB
TypeScript
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)),
|
|
);
|
|
}
|
|
}
|