Files
Mujoco_WASM/web_platform/src/rl/runtime/OnnxPolicyRuntime.test.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

218 lines
7.6 KiB
TypeScript
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import { beforeEach, describe, expect, it, vi } from 'vitest';
import fixture from '../fixtures/obstacleDeployment.json';
import { validatePolicyDeployment } from '../deployment';
const mock = vi.hoisted(() => ({
create: vi.fn(),
tensors: [] as { dispose: ReturnType<typeof vi.fn> }[],
}));
vi.mock('onnxruntime-web/wasm', () => ({
env: { wasm: {} },
InferenceSession: { create: mock.create },
Tensor: class {
dispose = vi.fn();
constructor(
public type: string,
public data: Float32Array,
public dims: number[],
) {
mock.tensors.push(this);
}
},
}));
import { OnnxPolicyRuntime } from './OnnxPolicyRuntime';
const deployment = validatePolicyDeployment(fixture);
const flush = async () => {
for (let i = 0; i < 10; i++) await Promise.resolve();
};
function setup(size = 81) {
const session = {
inputNames: ['obs'],
outputNames: ['actions'],
inputMetadata: [{ isTensor: true, type: 'float32', shape: [1, size] }],
outputMetadata: [{ isTensor: true, type: 'float32', shape: [1, 12] }],
run: vi.fn(),
release: vi.fn().mockResolvedValue(undefined),
};
const bindings = {
observe: vi.fn(() => new Float32Array(size)),
apply: vi.fn(),
clear: vi.fn(),
reset: vi.fn(),
};
mock.create.mockResolvedValue(session);
return { session, bindings };
}
beforeEach(() => {
vi.clearAllMocks();
mock.tensors.length = 0;
});
describe('OnnxPolicyRuntime held-action', () => {
it('动态目标转发不重建ORT、不启用策略或打断in-flight,status实时读取', async () => {
const { session, bindings } = setup();
let finish!: (outputs: Record<string, unknown>) => void;
session.run.mockImplementation(
() =>
new Promise((resolve) => {
finish = resolve;
}),
);
const navigation = {
target: [1, 2] as [number, number],
defaultTarget: [5, 0] as [number, number],
distance: 3,
targetHeight: 0,
};
const targetBindings = {
...bindings,
navigationStatus: () => navigation,
setNavigationTarget: vi.fn(),
resetNavigationTarget: vi.fn(),
};
const runtime = await OnnxPolicyRuntime.load(
new Uint8Array(),
'goal.onnx',
targetBindings,
deployment,
);
runtime.setNavigationTarget([2, 3]);
expect(targetBindings.setNavigationTarget).toHaveBeenCalledWith([2, 3]);
expect(runtime.status().enabled).toBe(false);
expect(runtime.status().navigation).toBe(navigation);
runtime.setEnabled(true, 0);
runtime.step(0);
runtime.setNavigationTarget([3, 4]);
runtime.step(0.02);
expect(session.run).toHaveBeenCalledOnce();
expect(mock.create).toHaveBeenCalledOnce();
runtime.resetNavigationTarget();
expect(targetBindings.resetNavigationTarget).toHaveBeenCalledOnce();
finish({ actions: { type: 'float32', data: new Float32Array(12), dispose: vi.fn() } });
await flush();
runtime.dispose();
expect(runtime.status().navigation).toBeUndefined();
});
it('81维契约单in-flight,持有最新动作,不阻塞物理步;reset拒绝旧promise', async () => {
const { session, bindings } = setup();
let finish!: (outputs: Record<string, unknown>) => void;
session.run.mockImplementation(
() =>
new Promise((resolve) => {
finish = resolve;
}),
);
const runtime = await OnnxPolicyRuntime.load(
new Uint8Array([1]),
'policy.onnx',
bindings,
deployment,
);
runtime.setEnabled(true, 0);
runtime.step(0);
runtime.step(0.02);
runtime.step(0.04);
expect(session.run).toHaveBeenCalledOnce();
expect(bindings.apply).toHaveBeenCalledTimes(3);
const output = { type: 'float32', data: new Float32Array(12).fill(0.3), dispose: vi.fn() };
finish({ actions: output });
await flush();
runtime.step(0.06);
expect(bindings.apply.mock.lastCall?.[0][0]).toBeCloseTo(0.3);
expect(session.run).toHaveBeenCalledTimes(2);
expect(output.dispose).toHaveBeenCalledOnce();
runtime.reset(0);
finish({ actions: { ...output, data: new Float32Array(12).fill(0.9) } });
await flush();
runtime.step(0);
expect(bindings.apply.mock.lastCall?.[0][0]).toBe(0);
expect(bindings.reset).toHaveBeenCalledWith(0);
runtime.dispose();
expect(session.release).not.toHaveBeenCalled();
finish({ actions: output });
await flush();
expect(session.release).toHaveBeenCalledOnce();
expect(mock.tensors.every((t) => t.dispose.mock.calls.length === 1)).toBe(true);
});
it('81策略不能默认为47、Rough维度不能加载,失败释放session', async () => {
const { session, bindings } = setup();
await expect(OnnxPolicyRuntime.load(new Uint8Array(), 'wrong.onnx', bindings)).rejects.toThrow(
/维度/,
);
expect(session.release).toHaveBeenCalledOnce();
session.inputMetadata[0].shape = [1, 234];
await expect(
OnnxPolicyRuntime.load(new Uint8Array(), 'rough.onnx', bindings, deployment),
).rejects.toThrow(/维度/);
});
it('原47维Flat保持向后兼容,非法观测/输出失败关闭并清动作', async () => {
const { session, bindings } = setup(47);
session.run.mockResolvedValue({
actions: { type: 'float32', data: new Float32Array(12).fill(NaN), dispose: vi.fn() },
});
const runtime = await OnnxPolicyRuntime.load(new Uint8Array(), 'flat.onnx', bindings);
expect(runtime.status().observationSize).toBe(47);
runtime.setEnabled(true, 0);
runtime.step(0);
await flush();
expect(runtime.status().enabled).toBe(false);
expect(runtime.status().error).toMatch(/非有限/);
runtime.reset(0);
bindings.observe.mockReturnValue(new Float32Array(81));
runtime.setEnabled(true, 0);
runtime.step(0);
expect(runtime.status().error).toMatch(/维度/);
runtime.dispose();
await flush();
});
});
it('默认Flat作业的兼容契约仍检查真实graph,拒绝81维及动态特征', async () => {
const flat = validatePolicyDeployment({
...fixture,
taskId: 'Unitree-Go2-Flat',
observationSize: 47,
observationTerms: fixture.observationTerms.slice(0, 7),
terrain: undefined,
terrainPreset: undefined,
terrainParams: undefined,
sensorCfg: undefined,
navigation: undefined,
});
const { session, bindings } = setup(47);
const runtime = await OnnxPolicyRuntime.load(new Uint8Array(), 'legacy.onnx', bindings, flat);
expect(runtime.status().observationSize).toBe(47);
runtime.dispose();
await flush();
session.inputMetadata[0].shape = [1, 81];
await expect(
OnnxPolicyRuntime.load(new Uint8Array(), 'legacy-wrong.onnx', bindings, flat),
).rejects.toThrow(/维度/);
session.inputMetadata[0].shape = [1, -1];
await expect(
OnnxPolicyRuntime.load(new Uint8Array(), 'legacy-dynamic.onnx', bindings, flat),
).rejects.toThrow(/固定/);
});
it('97 metadata不能冒充81 graph;真实97 shape沿用single-flight生命周期', async () => {
const multi = validatePolicyDeployment({
...fixture,
observationSize: 97,
sensorCfg: { ...fixture.sensorCfg, sensorMode: 'multi_ring_raycast', rayCount: 48 },
});
const { session, bindings } = setup(81);
await expect(
OnnxPolicyRuntime.load(new Uint8Array([1]), 'multi.onnx', bindings, multi),
).rejects.toThrow(/维度/);
expect(session.release).toHaveBeenCalledOnce();
const next = setup(97);
const runtime = await OnnxPolicyRuntime.load(
new Uint8Array([1]),
'multi.onnx',
next.bindings,
multi,
);
expect(runtime.status().observationSize).toBe(97);
runtime.dispose();
await flush();
expect(next.session.release).toHaveBeenCalledOnce();
});