260 lines
9.9 KiB
TypeScript
260 lines
9.9 KiB
TypeScript
import { readFileSync } from 'node:fs';
|
|
import { beforeEach, describe, expect, it, vi } from 'vitest';
|
|
import type { ProjectManifest } from '../project/types';
|
|
import fixture from '../rl/fixtures/obstacleDeployment.json';
|
|
import { validatePolicyDeployment } from '../rl/deployment';
|
|
const ort = vi.hoisted(() => ({ create: vi.fn() }));
|
|
vi.mock('onnxruntime-web/wasm', () => ({
|
|
env: { wasm: {} },
|
|
InferenceSession: { create: ort.create },
|
|
}));
|
|
vi.mock('@mujoco/mujoco', async (original) => {
|
|
const actual = await original<typeof import('@mujoco/mujoco')>();
|
|
return {
|
|
...actual,
|
|
default: () =>
|
|
actual.default({ wasmBinary: readFileSync('node_modules/@mujoco/mujoco/mujoco.wasm') }),
|
|
};
|
|
});
|
|
import { MainThreadPhysicsAdapter } from './PhysicsAdapter';
|
|
const deployment = validatePolicyDeployment(fixture);
|
|
function graph(size = 47) {
|
|
return {
|
|
inputNames: ['obs'],
|
|
outputNames: ['actions'],
|
|
inputMetadata: [{ isTensor: true, type: 'float32', shape: [1, size] }],
|
|
outputMetadata: [{ isTensor: true, type: 'float32', shape: [1, 12] }],
|
|
release: vi.fn().mockResolvedValue(undefined),
|
|
};
|
|
}
|
|
function project(): ProjectManifest {
|
|
const doc = new DOMParser().parseFromString(
|
|
readFileSync('training_server/rl/src/assets/robots/unitree_go2/xmls/go2.xml', 'utf8'),
|
|
'application/xml',
|
|
);
|
|
doc.querySelectorAll('mesh, geom[mesh]').forEach((e) => e.remove());
|
|
const actuators = doc.createElement('actuator');
|
|
for (const name of deployment.jointNames) {
|
|
const motor = doc.createElement('motor');
|
|
motor.setAttribute('joint', name);
|
|
motor.setAttribute('name', `${name}_motor`);
|
|
actuators.append(motor);
|
|
}
|
|
doc.documentElement.append(actuators);
|
|
const data = new TextEncoder().encode(new XMLSerializer().serializeToString(doc));
|
|
return {
|
|
id: 'transaction',
|
|
name: 'go2',
|
|
files: [{ path: 'go2.xml', data, size: data.length, source: 'file', mimeType: 'text/xml' }],
|
|
entries: [{ path: 'go2.xml', format: 'mjcf', label: 'go2' }],
|
|
maps: [],
|
|
selectedEntry: 'go2.xml',
|
|
totalBytes: data.length,
|
|
};
|
|
}
|
|
async function existing() {
|
|
const adapter = new MainThreadPhysicsAdapter(),
|
|
source = project(),
|
|
oldGraph = graph();
|
|
await adapter.load(source, 'go2.xml');
|
|
ort.create.mockResolvedValueOnce(oldGraph);
|
|
await adapter.loadRLPolicy(new Uint8Array([1]), 'old-flat.onnx');
|
|
adapter.setRLPolicyEnabled(true);
|
|
const old = adapter.session!;
|
|
old.data.time = 7;
|
|
old.data.qvel[0] = 0.123;
|
|
old.data.qpos[0] = 1.25;
|
|
old.data.ctrl[0] = 0.3;
|
|
adapter.setPaused(false);
|
|
return { adapter, source, old, oldGraph, before: adapter.snapshot()! };
|
|
}
|
|
beforeEach(() => {
|
|
vi.clearAllMocks();
|
|
ort.create.mockReset();
|
|
});
|
|
describe('PhysicsAdapter training transaction', () => {
|
|
it.each(['wrong graph', 'ORT initialization', 'binding'] as const)(
|
|
'%s失败保留原session/workspace/策略及物理状态',
|
|
async (failure) => {
|
|
const { adapter, source, old, oldGraph, before } = await existing();
|
|
const workspace = adapter.workspace,
|
|
dispose = vi.spyOn(old, 'dispose');
|
|
const candidate = graph(47);
|
|
if (failure === 'ORT initialization')
|
|
ort.create.mockRejectedValueOnce(new Error('ORT init failed'));
|
|
else ort.create.mockResolvedValueOnce(candidate);
|
|
if (failure === 'binding') {
|
|
const text = new TextDecoder()
|
|
.decode(source.files[0].data)
|
|
.replace('name="FL_hip_joint_motor"', 'name="incompatible_name"');
|
|
source.files[0].data = new TextEncoder().encode(text);
|
|
}
|
|
try {
|
|
await expect(
|
|
adapter.load(source, 'go2.xml', {
|
|
trainingDeployment: deployment,
|
|
trainingPolicy: { data: new Uint8Array([2]), path: 'candidate.onnx' },
|
|
}),
|
|
).rejects.toThrow(
|
|
failure === 'wrong graph' ? /维度/ : failure === 'binding' ? /驱动器/ : /ORT init/,
|
|
);
|
|
expect(adapter.session).toBe(old);
|
|
expect(adapter.workspace).toBe(workspace);
|
|
expect(adapter.snapshot()).toEqual(before);
|
|
expect(dispose).not.toHaveBeenCalled();
|
|
expect(oldGraph.release).not.toHaveBeenCalled();
|
|
if (failure === 'wrong graph') expect(candidate.release).toHaveBeenCalledOnce();
|
|
} finally {
|
|
adapter.dispose();
|
|
}
|
|
},
|
|
);
|
|
it('完成ORT初始化前不发布候选,成功后仍保留旧资源以供viewer失败回滚', async () => {
|
|
const { adapter, source, old, oldGraph, before } = await existing();
|
|
let finish!: (value: ReturnType<typeof graph>) => void;
|
|
let started!: () => void;
|
|
const initialized = new Promise<void>((resolve) => {
|
|
started = resolve;
|
|
});
|
|
ort.create.mockImplementationOnce(() => {
|
|
started();
|
|
return new Promise((resolve) => {
|
|
finish = resolve;
|
|
});
|
|
});
|
|
const candidate = graph(81);
|
|
try {
|
|
const loading = adapter.load(source, 'go2.xml', {
|
|
trainingDeployment: deployment,
|
|
trainingPolicy: { data: new Uint8Array([2]), path: 'new.onnx' },
|
|
});
|
|
await initialized;
|
|
expect(adapter.session).toBe(old);
|
|
expect(oldGraph.release).not.toHaveBeenCalled();
|
|
finish(candidate);
|
|
await loading;
|
|
expect(adapter.session).not.toBe(old);
|
|
expect(adapter.snapshot()?.rlPolicy).toMatchObject({
|
|
path: 'new.onnx',
|
|
observationSize: 81,
|
|
enabled: true,
|
|
});
|
|
expect(oldGraph.release).not.toHaveBeenCalled();
|
|
adapter.rollbackRetired();
|
|
expect(adapter.session).toBe(old);
|
|
expect(adapter.snapshot()).toEqual(before);
|
|
for (let i = 0; i < 10; i++) await Promise.resolve();
|
|
expect(candidate.release).toHaveBeenCalledOnce();
|
|
expect(oldGraph.release).not.toHaveBeenCalled();
|
|
} finally {
|
|
adapter.dispose();
|
|
}
|
|
});
|
|
});
|
|
|
|
it('候选ORT等待期间另一场景完成加载,迟到候选不得覆盖最新场景', async () => {
|
|
const { adapter, source } = await existing();
|
|
let finish!: (value: ReturnType<typeof graph>) => void;
|
|
let started!: () => void;
|
|
const initialized = new Promise<void>((resolve) => {
|
|
started = resolve;
|
|
});
|
|
ort.create.mockImplementationOnce(() => {
|
|
started();
|
|
return new Promise((resolve) => {
|
|
finish = resolve;
|
|
});
|
|
});
|
|
const candidate = graph(81);
|
|
try {
|
|
const stale = adapter.load(source, 'go2.xml', {
|
|
trainingDeployment: deployment,
|
|
trainingPolicy: { data: new Uint8Array([2]), path: 'stale.onnx' },
|
|
});
|
|
await initialized;
|
|
await adapter.load(source, 'go2.xml');
|
|
const latest = adapter.session,
|
|
before = adapter.snapshot();
|
|
const rejected = expect(stale).rejects.toThrow(/取消/);
|
|
finish(candidate);
|
|
await rejected;
|
|
expect(adapter.session).toBe(latest);
|
|
expect(adapter.snapshot()).toEqual(before);
|
|
for (let i = 0; i < 10; i++) await Promise.resolve();
|
|
expect(candidate.release).toHaveBeenCalledOnce();
|
|
} finally {
|
|
adapter.dispose();
|
|
}
|
|
});
|
|
|
|
describe('PhysicsAdapter custom_boxes applied scene compiler', () => {
|
|
it('真实WASM读取工程路径/嵌套body变换/旋转多实例;与共享CPU布局halfsize和原点一致', async () => {
|
|
const { default: shared } =
|
|
await import('../../../training_server/tests/fixtures/custom-boxes.json');
|
|
const source = project();
|
|
const descriptor = {
|
|
schemaVersion: 1,
|
|
id: 'warehouse',
|
|
name: '仓库',
|
|
coordinateSystem: { units: 'm', up: 'Z', forward: '+X' },
|
|
physics: { source: 'physics/world.xml' },
|
|
spawnPoints: [],
|
|
};
|
|
for (const [path, text] of [
|
|
['maps/warehouse/map.json', JSON.stringify(descriptor)],
|
|
[
|
|
'maps/warehouse/physics/world.xml',
|
|
'<mujoco><worldbody><body name="nested" pos="0 1 .5"><geom name="wall" type="box" size=".4 .3 .5" friction=".8 .005 .0001"/><geom name="decoration" type="sphere" size="10" contype="0" conaffinity="0"/></body></worldbody></mujoco>',
|
|
],
|
|
]) {
|
|
const data = new TextEncoder().encode(text);
|
|
source.files.push({ path, data, size: data.length, source: 'file', mimeType: 'text/xml' });
|
|
}
|
|
const assets = [
|
|
{
|
|
id: 'one',
|
|
name: 'one',
|
|
selection: {
|
|
kind: 'project' as const,
|
|
descriptorPath: 'maps/warehouse/map.json',
|
|
positionX: 1,
|
|
positionY: 1,
|
|
},
|
|
},
|
|
];
|
|
const adapter = new MainThreadPhysicsAdapter();
|
|
try {
|
|
await adapter.load(source, 'go2.xml', { mapAssets: assets });
|
|
const initial = adapter.exportTrainingTerrain(assets, { spawn: [-2, -1], target: [2, -1] });
|
|
// Geometry from actual WASM matches the JSON consumed by CPU TerrainGenerator tests.
|
|
expect(initial.boxes[1]).toEqual(shared.boxes[1]);
|
|
expect(initial.actualObstacleCount).toBe(1);
|
|
adapter.session!.data.qpos[0] = 8;
|
|
expect(adapter.exportTrainingTerrain(assets).spawn[0]).not.toBe(8);
|
|
expect(() => adapter.exportTrainingTerrain([])).toThrow(/过时/);
|
|
const multi = [
|
|
...assets,
|
|
{
|
|
id: 'two',
|
|
name: 'two',
|
|
selection: { ...assets[0].selection, positionX: -2, positionY: 1, yawDeg: 45 },
|
|
},
|
|
];
|
|
await adapter.load(source, 'go2.xml', { mapAssets: multi });
|
|
const layout = adapter.exportTrainingTerrain(multi, { spawn: [-2, -1], target: [2, -1] });
|
|
expect(layout.boxes).toHaveLength(3);
|
|
expect(layout.boxes[2].pos[0]).toBeCloseTo(-2 - Math.SQRT1_2, 12);
|
|
expect(layout.boxes[2].pos[1]).toBeCloseTo(1 + Math.SQRT1_2, 12);
|
|
expect(layout.boxes[2].size[0]).toBeCloseTo(0.7 * Math.SQRT1_2, 12);
|
|
expect(layout.boxes[2].size[1]).toBeCloseTo(0.7 * Math.SQRT1_2, 12);
|
|
adapter.rollbackRetired();
|
|
expect(adapter.exportTrainingTerrain(assets).actualObstacleCount).toBe(1);
|
|
// A deployment replacement cannot pretend to be the original applied editor map.
|
|
await adapter.load(source, 'go2.xml', { mapAssets: assets, trainingDeployment: deployment });
|
|
expect(() => adapter.exportTrainingTerrain(assets)).toThrow(/已应用/);
|
|
} finally {
|
|
adapter.dispose();
|
|
}
|
|
});
|
|
});
|