Files
Mujoco_WASM/web_platform/src/simulation/PhysicsAdapter.trainingTransaction.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

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();
}
});
});