feat(training): release V0.9.1 避障训练与基础策略迁移
This commit is contained in:
@@ -1,3 +1,10 @@
|
||||
import {
|
||||
composeTrainingMap,
|
||||
trainingTerrainFromCompiledScene,
|
||||
type CompiledMapGeometry,
|
||||
type TrainingSceneCoordinates,
|
||||
} from '../map/trainingMap';
|
||||
import type { PolicyDeployment, TrainingTerrain } from '../rl/deployment';
|
||||
import type { MainModule } from '@mujoco/mujoco';
|
||||
import type { ProjectFile, ProjectManifest } from '../project/types';
|
||||
import { prepareProjectForMujoco } from '../project/importer';
|
||||
@@ -36,6 +43,9 @@ export interface PhysicsLoadProgress {
|
||||
}
|
||||
|
||||
export interface PhysicsLoadOptions {
|
||||
trainingDeployment?: PolicyDeployment;
|
||||
/** Candidate policy is initialized/validated before replacing the active session. */
|
||||
trainingPolicy?: { data: Uint8Array; path: string };
|
||||
urdfMode?: UrdfLoadMode;
|
||||
baseMode?: UrdfBaseMode;
|
||||
enhancements?: UrdfEnhancementOptions;
|
||||
@@ -68,9 +78,15 @@ export interface PhysicsAdapter {
|
||||
setControllerEnabled(enabled: boolean): void;
|
||||
sendControllerCommand(command: ControllerCommand): void;
|
||||
removeController(): void;
|
||||
loadRLPolicy(model: Uint8Array, path: string): Promise<RLPolicyStatus>;
|
||||
loadRLPolicy(
|
||||
model: Uint8Array,
|
||||
path: string,
|
||||
deployment?: PolicyDeployment,
|
||||
): Promise<RLPolicyStatus>;
|
||||
setRLPolicyEnabled(enabled: boolean): void;
|
||||
setRLCommand(command: RLCommand): void;
|
||||
setNavigationTarget(target: [number, number]): void;
|
||||
resetNavigationTarget(): void;
|
||||
removeRLPolicy(): void;
|
||||
configureDataRecorder(config: Partial<DataRecorderConfig>): DataRecorderStatus | undefined;
|
||||
startDataRecording(): DataRecorderStatus | undefined;
|
||||
@@ -82,6 +98,10 @@ export interface PhysicsAdapter {
|
||||
releaseRetired(): void;
|
||||
rollbackRetired(): void;
|
||||
exportMjcf(): Uint8Array;
|
||||
exportTrainingTerrain(
|
||||
assets: readonly PlacedMapAsset[],
|
||||
coordinates?: TrainingSceneCoordinates,
|
||||
): TrainingTerrain;
|
||||
dispose(): void;
|
||||
}
|
||||
|
||||
@@ -108,6 +128,15 @@ export class MainThreadPhysicsAdapter implements PhysicsAdapter {
|
||||
session: SimulationSession | null = null;
|
||||
workspace: MemfsWorkspace | null = null;
|
||||
private supportFiles: ProjectFile[] = [];
|
||||
private trainingScenes = new WeakMap<
|
||||
SimulationSession,
|
||||
{
|
||||
assets: string;
|
||||
geometries: CompiledMapGeometry[];
|
||||
initialPose: number[];
|
||||
extent: number;
|
||||
}
|
||||
>();
|
||||
private loadGeneration = 0;
|
||||
private disposed = false;
|
||||
private retiredSession: SimulationSession | null = null;
|
||||
@@ -196,7 +225,7 @@ export class MainThreadPhysicsAdapter implements PhysicsAdapter {
|
||||
intermediate.dispose();
|
||||
}
|
||||
}
|
||||
if (hasMaps) {
|
||||
if (hasMaps && !options.trainingDeployment?.terrain) {
|
||||
report(0.72, '组合机器人与物理地图');
|
||||
if (entry?.format === 'urdf' && urdfMode === 'native')
|
||||
throw new Error('原生 URDF 模式暂不支持地图,请切换为转换模式');
|
||||
@@ -232,13 +261,102 @@ export class MainThreadPhysicsAdapter implements PhysicsAdapter {
|
||||
modelRelativePath = mapPath;
|
||||
modelPath = workspace.path(modelRelativePath);
|
||||
}
|
||||
if (options.trainingDeployment?.terrain) {
|
||||
if (entry?.format === 'urdf' && urdfMode === 'native')
|
||||
throw new Error('训练地图需要URDF转换模式');
|
||||
const intermediate = new SimulationSession(module, modelPath);
|
||||
try {
|
||||
const flatPath = `${modelRelativePath}.training.xml`;
|
||||
if (!module.mj_saveLastXML(workspace.path(flatPath), intermediate.model))
|
||||
throw new Error('无法展开训练场景');
|
||||
workspace.writeGenerated(
|
||||
flatPath,
|
||||
composeTrainingMap(
|
||||
new TextEncoder().encode(workspace.readText(flatPath)),
|
||||
options.trainingDeployment,
|
||||
),
|
||||
);
|
||||
modelPath = workspace.path(flatPath);
|
||||
} finally {
|
||||
intermediate.dispose();
|
||||
}
|
||||
warnings.push('已用策略配套训练布局替换场景地形;编辑器地图未更改。Go2-W动力学不等同Go2。');
|
||||
}
|
||||
report(0.84, '编译模型与物理数据');
|
||||
console.info('[MuJoCo] 编译模型', modelPath);
|
||||
nextSession = new SimulationSession(module, modelPath, warnings);
|
||||
if (options.trainingDeployment) nextSession.configureDeployment(options.trainingDeployment);
|
||||
if (entry?.format === 'urdf' && urdfMode === 'native') {
|
||||
const offset = nextSession.alignLowestPointToGround();
|
||||
warnings.push(`原生 URDF 已整体平移 ${offset.toFixed(4)} m,使最低点位于 z=0`);
|
||||
}
|
||||
if (placedMapAssets?.length && !options.trainingDeployment) {
|
||||
const sanitize = (id: string) => id.replace(/[^a-zA-Z0-9_-]/g, '_');
|
||||
const prefixes = placedMapAssets.map((asset) =>
|
||||
asset.selection.kind === 'builtin'
|
||||
? `__platform_map_${sanitize(asset.id)}__`
|
||||
: `__platform_map_${sanitize(`${resolveProjectMap(prepared.manifest, asset.selection.descriptorPath).definition.id}_${asset.id}`)}_`,
|
||||
);
|
||||
const model = nextSession.model,
|
||||
data = nextSession.data;
|
||||
const geometries: CompiledMapGeometry[] = [];
|
||||
for (let id = 0; id < model.ngeom; id++) {
|
||||
const geom = model.geom(id);
|
||||
try {
|
||||
const name = geom.name;
|
||||
geometries.push({
|
||||
name,
|
||||
type: Number(model.geom_type[id]),
|
||||
position: Array.from(data.geom_xpos.slice(id * 3, id * 3 + 3)),
|
||||
rotation: Array.from(data.geom_xmat.slice(id * 9, id * 9 + 9)),
|
||||
size: Array.from(model.geom_size.slice(id * 3, id * 3 + 3)),
|
||||
friction: Array.from(model.geom_friction.slice(id * 3, id * 3 + 3)),
|
||||
collision: Boolean(model.geom_contype[id] || model.geom_conaffinity[id]),
|
||||
static: Number(model.body_weldid[Number(model.geom_bodyid[id])]) === 0,
|
||||
map:
|
||||
prefixes.some((prefix) => name.startsWith(prefix)) ||
|
||||
name === '__platform_map_ground__',
|
||||
});
|
||||
} finally {
|
||||
geom.delete();
|
||||
}
|
||||
}
|
||||
const roots = Array.from({ length: model.njnt }, (_, id) => id).filter(
|
||||
(id) => Number(model.jnt_type[id]) === 0,
|
||||
);
|
||||
if (roots.length === 1) {
|
||||
const address = Number(model.jnt_qposadr[roots[0]]);
|
||||
this.trainingScenes.set(nextSession, {
|
||||
assets: JSON.stringify(placedMapAssets),
|
||||
geometries,
|
||||
initialPose: Array.from(data.qpos.slice(address, address + 7)),
|
||||
extent: Math.max(
|
||||
4,
|
||||
...placedMapAssets.map((asset) =>
|
||||
asset.selection.kind === 'builtin'
|
||||
? Math.max(
|
||||
Math.abs(asset.selection.config.positionX),
|
||||
Math.abs(asset.selection.config.positionY),
|
||||
) +
|
||||
asset.selection.config.size / 2
|
||||
: 0,
|
||||
),
|
||||
),
|
||||
});
|
||||
}
|
||||
}
|
||||
if (options.trainingPolicy) {
|
||||
if (!options.trainingDeployment?.terrain) throw new Error('事务策略加载需要配套训练地图');
|
||||
report(0.9, '校验候选场景的 ONNX 策略');
|
||||
await nextSession.loadRLPolicy(
|
||||
options.trainingPolicy.data,
|
||||
options.trainingPolicy.path,
|
||||
options.trainingDeployment,
|
||||
);
|
||||
// Configure/reset already supplied the correct spawn. Enable while still paused;
|
||||
// the candidate cannot advance until the viewer transaction commits.
|
||||
nextSession.setRLPolicyEnabled(true);
|
||||
}
|
||||
if (warnings.length) console.info('[MuJoCo] 兼容与地图处理', warnings);
|
||||
console.info('[MuJoCo] 模型编译完成');
|
||||
report(0.94, '生成初始仿真状态');
|
||||
@@ -263,6 +381,20 @@ export class MainThreadPhysicsAdapter implements PhysicsAdapter {
|
||||
);
|
||||
}
|
||||
}
|
||||
exportTrainingTerrain(
|
||||
assets: readonly PlacedMapAsset[],
|
||||
coordinates?: TrainingSceneCoordinates,
|
||||
): TrainingTerrain {
|
||||
const scene = this.session && this.trainingScenes.get(this.session);
|
||||
if (!scene || scene.assets !== JSON.stringify(assets))
|
||||
throw new Error('没有匹配的已应用碰撞场景/唯一浮动机器人,或场景已过时;请先应用地图');
|
||||
return trainingTerrainFromCompiledScene(
|
||||
scene.geometries,
|
||||
scene.initialPose,
|
||||
scene.extent,
|
||||
coordinates,
|
||||
);
|
||||
}
|
||||
advance(now: number): FrameResult {
|
||||
return this.session?.advance(now) ?? { steps: 0, stepMs: 0, overBudget: false };
|
||||
}
|
||||
@@ -315,13 +447,23 @@ export class MainThreadPhysicsAdapter implements PhysicsAdapter {
|
||||
removeController(): void {
|
||||
this.session?.removeController();
|
||||
}
|
||||
async loadRLPolicy(model: Uint8Array, path: string): Promise<RLPolicyStatus> {
|
||||
async loadRLPolicy(
|
||||
model: Uint8Array,
|
||||
path: string,
|
||||
deployment?: PolicyDeployment,
|
||||
): Promise<RLPolicyStatus> {
|
||||
if (!this.session) throw new Error('请先加载模型');
|
||||
return this.session.loadRLPolicy(model, path);
|
||||
return this.session.loadRLPolicy(model, path, deployment);
|
||||
}
|
||||
setRLPolicyEnabled(enabled: boolean): void {
|
||||
this.session?.setRLPolicyEnabled(enabled);
|
||||
}
|
||||
setNavigationTarget(target: [number, number]): void {
|
||||
this.session?.setNavigationTarget(target);
|
||||
}
|
||||
resetNavigationTarget(): void {
|
||||
this.session?.resetNavigationTarget();
|
||||
}
|
||||
setRLCommand(command: RLCommand): void {
|
||||
this.session?.setRLCommand(command);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user