218 lines
7.6 KiB
TypeScript
218 lines
7.6 KiB
TypeScript
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();
|
||
});
|