import { PretrainedIdentity, PretrainedSourceSelect } from './PretrainedSourceSelect'; import { pretrainedSelectionError } from './pretrainedSelection'; import { TrainingMetricsPanel } from './TrainingMetricsPanel'; import { trainingLosses } from './trainingLosses'; import type { PolicyDeployment, TrainingTerrain } from '../rl/deployment'; import { OBSTACLE_TASK_ID, validatePolicyDeployment, validateCustomTerrain, readPolicyDeployment, } from '../rl/deployment'; import type { TrainingSceneCompiler } from '../map/trainingMap'; import type { PlacedMapAsset } from '../map/types'; import { useEffect, useRef, useState, type ReactNode } from 'react'; import { Download, ExternalLink, Link, Play, Server, Square } from 'lucide-react'; import { Badge, Button, ProgressBar, PropertyRow, Select } from '../components/ui'; import { LocalTrainingClient } from './LocalTrainingClient'; import type { RewardPreset, TrainingDevice, TrainingJob, TrainingServerInfo, WandbMode, } from './types'; import { DEFAULT_TRAINING_ENDPOINT, localStored, rememberTrainingConnection, sessionStored, TRAINING_ENDPOINT_KEY, TRAINING_JOB_KEY, TRAINING_TOKEN_KEY, } from './storage'; const ACTIVE_STATES = new Set(['queued', 'running']); function errorText(error: unknown): string { return error instanceof Error ? error.message : String(error); } function stateLabel(state: TrainingJob['state']): string { return { queued: '排队中', running: '训练中', succeeded: '已完成', failed: '失败', cancelled: '已取消', }[state]; } export function LocalTrainingPanel({ onPolicyReady, compileScene, sceneMaps = [], sceneDirty = false, }: { onPolicyReady(file: File, deployment?: PolicyDeployment): void | Promise; compileScene?: TrainingSceneCompiler; sceneMaps?: readonly PlacedMapAsset[]; sceneDirty?: boolean; }) { const [endpoint, setEndpoint] = useState(() => localStored(TRAINING_ENDPOINT_KEY, DEFAULT_TRAINING_ENDPOINT), ); const [token, setToken] = useState(() => sessionStored(TRAINING_TOKEN_KEY)); const [server, setServer] = useState(); const [job, setJob] = useState(); const [presets, setPresets] = useState([]); const [rewardPresetId, setRewardPresetId] = useState(''); const [pretrainedSourceId, setPretrainedSourceId] = useState(''); const [uploading, setUploading] = useState(false); const [connectionRevision, setConnectionRevision] = useState(0); const [busy, setBusy] = useState(false), [error, setError] = useState(); const [taskId, setTaskId] = useState('Unitree-Go2-Flat'), [numEnvs, setNumEnvs] = useState(4096), [maxIterations, setMaxIterations] = useState(2000), [seed, setSeed] = useState(42), [runName, setRunName] = useState('web'), [device, setDevice] = useState('gpu'), [gpuIds, setGpuIds] = useState('0'), [wandbMode, setWandbMode] = useState('offline'); const [terrainPreset, setTerrainPreset] = useState(''); const [customTerrainBoxes, setCustomTerrainBoxes] = useState(); const [terrainParams, setTerrainParams] = useState>({}); const [sensorMode, setSensorMode] = useState<'single_ring_raycast' | 'multi_ring_raycast'>( 'single_ring_raycast', ); const [sensorCfg, setSensorCfg] = useState>({}); const metadata = server?.taskMetadata?.find((item) => item.id === taskId); const sourceSelectionError = pretrainedSelectionError( server?.pretrainedSources, taskId, pretrainedSourceId, ); const selectTask = (id: string) => { setTaskId(id); setCustomTerrainBoxes(undefined); setRewardPresetId(''); setTerrainParams({}); setSensorCfg({}); setSensorMode('single_ring_raycast'); setTerrainPreset(id === OBSTACLE_TASK_ID ? 'discrete_obstacles' : ''); }; const syncMap = () => { try { if (sceneDirty) throw new Error('请先应用地图草稿,再同步训练地图'); if (!compileScene || !sceneMaps.length) throw new Error('没有已应用的碰撞地图'); if (!metadata?.terrainPresets.includes('custom_boxes')) throw new Error('当前服务/任务不支持custom_boxes,请升级训练服务'); const layout = validateCustomTerrain(compileScene()); setCustomTerrainBoxes(layout); setTerrainPreset('custom_boxes'); setTerrainParams({ size: layout.size, friction: layout.friction }); setError(undefined); } catch (value) { setError(errorText(value)); } }; const connectionEpoch = useRef(0); const connected = () => { connectionEpoch.current += 1; setConnectionRevision(connectionEpoch.current); setServer(undefined); setJob(undefined); setPresets([]); setRewardPresetId(''); setError(undefined); }; useEffect( () => () => { connectionEpoch.current += 1; }, [], ); const connect = async () => { setBusy(true); setError(undefined); const epoch = ++connectionEpoch.current; setConnectionRevision(epoch); try { const client = new LocalTrainingClient(endpoint, token), info = await client.health(); if (epoch !== connectionEpoch.current) return; setServer(info); try { const nextPresets = await client.presets(); if (epoch !== connectionEpoch.current) return; setPresets(nextPresets); } catch { setPresets([]); } try { rememberTrainingConnection(client.endpoint, client.token); } catch { /* 当前会话仍可连接 */ } if (info.tasks.length && !info.tasks.includes(taskId)) selectTask(info.tasks[0]); const remembered = info.activeJobId ?? localStored(TRAINING_JOB_KEY); if (remembered) { try { const recovered = await client.job(remembered); if (epoch !== connectionEpoch.current) return; setJob(recovered); try { localStorage.setItem(TRAINING_JOB_KEY, recovered.id); } catch { /* ignore */ } } catch { setJob(undefined); try { localStorage.removeItem(TRAINING_JOB_KEY); } catch { /* ignore */ } } } else { setJob(undefined); } if (!info.ready) setError(info.error ?? '训练服务尚未就绪'); } catch (value) { setServer(undefined); setError(errorText(value)); } finally { setBusy(false); } }; const jobId = job?.id, jobState = job?.state; useEffect(() => { if (!jobId || !jobState || !ACTIVE_STATES.has(jobState)) return; let disposed = false; const refresh = async () => { try { const next = await new LocalTrainingClient(endpoint, token).job(jobId); if (!disposed) setJob(next); } catch (value) { if (!disposed) setError(errorText(value)); } }; const timer = window.setInterval(() => void refresh(), 1500); return () => { disposed = true; window.clearInterval(timer); }; }, [endpoint, jobId, jobState, token]); useEffect(() => { const receive = (event: MessageEvent) => { if ( event.origin !== window.location.origin || !event.source || typeof event.data !== 'object' ) return; const data = event.data as { type?: string; sessionId?: string; policy?: unknown; taskId?: string; }; const source = event.source as Window; if (data.type === 'mujoco-tuning-ready') { try { let resolvedCustomTerrain = customTerrainBoxes; let resolvedTerrainParams = terrainParams; if (taskId === OBSTACLE_TASK_ID && terrainPreset === 'custom_boxes') { if (sceneDirty) throw new Error('请先应用地图草稿,再打开调参'); if (!compileScene || !sceneMaps.length) throw new Error('没有已应用的碰撞地图'); resolvedCustomTerrain = validateCustomTerrain(compileScene()); resolvedTerrainParams = { size: resolvedCustomTerrain.size, friction: resolvedCustomTerrain.friction, }; setCustomTerrainBoxes(resolvedCustomTerrain); setTerrainParams(resolvedTerrainParams); } source.postMessage( { type: 'mujoco-tuning-credentials', endpoint, token, trainingContext: { taskId: taskId === OBSTACLE_TASK_ID ? taskId : 'Unitree-Go2-Flat', seed, ...(pretrainedSourceId ? { pretrainedSourceId } : {}), ...(taskId === OBSTACLE_TASK_ID ? { taskConfig: { terrainPreset, terrainParams: resolvedTerrainParams, sensorType: 'raycast', sensorCfg: { ...sensorCfg, sensorMode }, ...(terrainPreset === 'custom_boxes' ? { customTerrainBoxes: resolvedCustomTerrain } : {}), }, } : {}), }, }, event.origin, ); } catch (value) { setError(errorText(value)); } } if (data.type === 'mujoco-tuning-import-policy' && data.sessionId) { const reply = (ok: boolean, message?: string) => { try { source.postMessage( { type: 'mujoco-tuning-import-policy-result', sessionId: data.sessionId, ok, error: message, }, event.origin, ); } catch { /* 调参窗口可能已关闭;不影响主工作台继续导入 */ } }; void (async () => { try { const policy = data.policy === undefined ? await new LocalTrainingClient(endpoint, token).downloadBestPolicy(data.sessionId!) : data.policy; if (!(policy instanceof File) || !/\.onnx$/i.test(policy.name)) throw new Error('调参工作台返回的 ONNX 策略无效'); if (policy.size > 64 * 1024 * 1024) throw new Error('ONNX 策略不能超过 64 MiB'); if (data.taskId === OBSTACLE_TASK_ID) { const deployment = readPolicyDeployment(new Uint8Array(await policy.arrayBuffer())); if (deployment?.taskId !== OBSTACLE_TASK_ID) throw new Error('避障最佳策略缺少匹配部署契约'); await onPolicyReady(policy, deployment); } else await onPolicyReady(policy); reply(true); } catch (value) { const message = errorText(value); setError(message); reply(false, message); } })(); } }; window.addEventListener('message', receive); return () => window.removeEventListener('message', receive); }, [ endpoint, pretrainedSourceId, onPolicyReady, token, taskId, seed, terrainPreset, terrainParams, sensorCfg, sensorMode, customTerrainBoxes, sceneMaps, sceneDirty, compileScene, ]); const openTuningDashboard = () => { rememberTrainingConnection(endpoint, token); window.open(new URL('tuning.html', document.baseURI), '_blank'); }; const start = async () => { if (uploading) return; if (sourceSelectionError) { setError(sourceSelectionError); return; } setBusy(true); setError(undefined); try { if ( taskId === OBSTACLE_TASK_ID && sensorMode === 'multi_ring_raycast' && !metadata?.sensorModes?.includes(sensorMode) ) throw new Error('训练服务不支持multi_ring_raycast,请升级服务'); const ids = device === 'gpu' ? gpuIds .split(/[\s,]+/) .filter(Boolean) .map(Number) : []; if (ids.some((id) => !Number.isInteger(id) || id < 0)) throw new Error('GPU 编号必须是非负整数'); for (const [values, schema] of [ [terrainParams, metadata?.terrainParameters], [sensorCfg, metadata?.sensorParameters], ] as const) { for (const [key, value] of Object.entries(values)) { const bounds = schema?.[key]; if ( !bounds || !Number.isFinite(value) || value < bounds.min || value > bounds.max || (bounds.integer && !Number.isInteger(value)) ) throw new Error(`参数 ${key} 超出允许范围`); } } if ((terrainParams.obstacle_height_min ?? 0.2) > (terrainParams.obstacle_height_max ?? 0.6)) throw new Error('障碍物最小高度不能超过最大高度'); if ((sensorCfg.safetyDistance ?? 0.5) >= (sensorCfg.maxDistance ?? 4)) throw new Error('安全距离必须小于探测距离'); let resolvedCustomTerrain = customTerrainBoxes; let resolvedTerrainParams = terrainParams; if (terrainPreset === 'custom_boxes') { if (sceneDirty) throw new Error('请先应用地图草稿,再训练'); if (!compileScene || !sceneMaps.length) throw new Error('没有已应用的碰撞地图'); resolvedCustomTerrain = validateCustomTerrain(compileScene()); resolvedTerrainParams = { size: resolvedCustomTerrain.size, friction: resolvedCustomTerrain.friction, }; setCustomTerrainBoxes(resolvedCustomTerrain); setTerrainParams(resolvedTerrainParams); } const next = await new LocalTrainingClient(endpoint, token).start({ taskId, numEnvs, maxIterations, seed, runName, device, gpuIds: ids, wandbMode, rewardPresetId: taskId === 'Unitree-Go2-Flat' ? rewardPresetId || undefined : undefined, ...(pretrainedSourceId ? { pretrainedSourceId } : {}), ...(terrainPreset ? { terrainPreset, terrainParams: resolvedTerrainParams } : {}), ...(terrainPreset === 'custom_boxes' ? { customTerrainBoxes: resolvedCustomTerrain } : {}), ...(taskId === OBSTACLE_TASK_ID ? { sensorType: 'raycast' as const, sensorCfg: { ...sensorCfg, sensorMode } } : {}), }); setJob(next); try { localStorage.setItem(TRAINING_JOB_KEY, next.id); } catch { /* ignore */ } } catch (value) { setError(errorText(value)); } finally { setBusy(false); } }; const cancel = async () => { if (!job) return; setBusy(true); setError(undefined); try { setJob(await new LocalTrainingClient(endpoint, token).cancel(job.id)); } catch (value) { setError(errorText(value)); } finally { setBusy(false); } }; const importResult = async () => { if (!job) return; setBusy(true); setError(undefined); try { if (job.taskId !== 'Unitree-Go2-Flat' && !job.deployment) throw new Error('该任务缺少浏览器部署契约'); const deployment = job.deployment ? validatePolicyDeployment(job.deployment) : undefined; const file = await new LocalTrainingClient(endpoint, token).downloadPolicy(job.id); await onPolicyReady(file, deployment); } catch (value) { setError(errorText(value)); } finally { setBusy(false); } }; const active = Boolean(job && ACTIVE_STATES.has(job.state)); return (
{server?.trainerRoot ?? '请先启动本地训练服务'} {server?.ready ? '可用' : '离线'}
{server?.ready && !job && (
{ setServer( (current) => current && { ...current, pretrainedSources: [ ...(current.pretrainedSources ?? []).filter((s) => s.id !== source.id), source, ], }, ); setPretrainedSourceId(source.id); }, }} /> {metadata && ( <>

启动训练时会自动从全部已应用实例重新编译并校验世界AABB,无需预先同步。旋转障碍会膨胀,底板标准化为z=[-0.2,0];仅保证训练与浏览器使用相同boxes,不等于原OBB。mesh/hfield、地下结构、混合摩擦明确拒绝。

{terrainPreset === 'custom_boxes' && customTerrainBoxes && !sceneDirty && (

已读取视口中 {customTerrainBoxes.actualObstacleCount}{' '} 个自定义障碍物;启动训练时会自动重新编译并校验地图

)} {terrainPreset && terrainPreset !== 'custom_boxes' && (
{Object.entries(metadata.terrainParameters).map(([key, bounds]) => ( setTerrainParams((old) => ({ ...old, [key]: value }))} /> ))}
)} {['rough', 'wave', 'pyramid_stairs'].includes(terrainPreset) && (

训练专用 box 离散近似布局,不等于编辑器高度场。

)} {taskId === OBSTACLE_TASK_ID && (
避障传感器高级设置 {Object.entries(metadata.sensorParameters).map(([key, bounds]) => ( setSensorCfg((old) => ({ ...old, [key]: value }))} /> ))}
)} {!metadata.browserCompatible && (

此任务可训练,但浏览器不支持其观测契约,不能一键部署。

)} )}
setRunName(event.target.value)} />
setGpuIds(event.target.value)} />

训练使用本地 mjlab 任务资产,不会把浏览器中的模型上传到网络。服务一次只运行一个训练任务。

)} {job && (
{job.taskId} {stateLabel(job.state)}
{job.taskId === 'Unitree-Go2-Rough' && (

234维 Rough 策略仅支持后端评测,浏览器不可加载。

)} {job.deployment?.terrain && (

导入将替换当前物理地图并启动配套策略; {job.deployment.terrain.approximation ? '训练专用近似布局' : '配套碰撞布局'} 。请先保存场景。

)}
{trainingLosses(job.logs).map(({ label, value }) => ( ))}
{job.logs.length > 0 && (
最近日志
                {job.logs.slice(-40).join('\n')}
              
)}
{active ? ( ) : ( <> )}
)} {error && (

{error}

)}
); } function Field({ label, children }: { label: string; children: ReactNode }) { return ( ); } function NumberField({ label, value, min, max, onChange, step = 1, }: { label: string; value: number; min: number; max: number; step?: number; onChange(value: number): void; }) { return ( onChange(Number(event.target.value))} /> ); } const TERRAIN_LABELS: Record = { custom_boxes: '自定义场景碰撞布局(AABB近似)', plane: '平地', discrete_obstacles: '离散障碍物', rough: '崎岖地面', pyramid_stairs: '金字塔台阶', wave: '波浪地形', }; const PARAMETER_LABELS: Record = { size: '地图尺寸 m', obstacle_count: '障碍物数量', obstacle_height_min: '最小障碍高度 m', obstacle_height_max: '最大障碍高度 m', spacing: '障碍物间距 m', friction: '地面摩擦', roughness: '崎岖高度 m', step_height: '台阶高度 m', wave_amplitude: '波浪幅度 m', fov: '感知角 FOV', maxDistance: '探测距离 m', safetyDistance: '安全距离 m', avoidanceWeight: '避障权重', };