895 lines
32 KiB
TypeScript
895 lines
32 KiB
TypeScript
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<void>;
|
||
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<TrainingServerInfo>();
|
||
const [job, setJob] = useState<TrainingJob>();
|
||
const [presets, setPresets] = useState<RewardPreset[]>([]);
|
||
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<string>();
|
||
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<TrainingDevice>('gpu'),
|
||
[gpuIds, setGpuIds] = useState('0'),
|
||
[wandbMode, setWandbMode] = useState<WandbMode>('offline');
|
||
|
||
const [terrainPreset, setTerrainPreset] = useState('');
|
||
const [customTerrainBoxes, setCustomTerrainBoxes] = useState<TrainingTerrain>();
|
||
const [terrainParams, setTerrainParams] = useState<Record<string, number>>({});
|
||
const [sensorMode, setSensorMode] = useState<'single_ring_raycast' | 'multi_ring_raycast'>(
|
||
'single_ring_raycast',
|
||
);
|
||
const [sensorCfg, setSensorCfg] = useState<Record<string, number>>({});
|
||
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 (
|
||
<div>
|
||
<label className="block text-xs text-text-secondary">
|
||
<span className="mb-1 block">本地训练服务</span>
|
||
<div className="flex gap-2">
|
||
<input
|
||
aria-label="本地训练服务地址"
|
||
className="field h-7 min-w-0 flex-1 px-2 text-xs text-text-primary"
|
||
value={endpoint}
|
||
disabled={busy}
|
||
onChange={(event) => {
|
||
if (busy) return;
|
||
connected();
|
||
setEndpoint(event.target.value);
|
||
}}
|
||
/>
|
||
<Button
|
||
icon={<Link className="h-3.5 w-3.5" />}
|
||
disabled={busy || !token.trim()}
|
||
onClick={() => void connect()}
|
||
>
|
||
连接
|
||
</Button>
|
||
</div>
|
||
</label>
|
||
<label className="mt-2 block text-[10px] text-text-tertiary">
|
||
<span className="mb-1 block">访问令牌(服务启动时显示)</span>
|
||
<input
|
||
aria-label="训练服务访问令牌"
|
||
type="password"
|
||
autoComplete="off"
|
||
className="field h-7 w-full px-2 text-xs text-text-primary"
|
||
value={token}
|
||
disabled={busy}
|
||
onChange={(event) => {
|
||
if (busy) return;
|
||
connected();
|
||
setToken(event.target.value);
|
||
}}
|
||
/>
|
||
</label>
|
||
<div className="mt-2 flex items-center justify-between rounded-md border border-border bg-surface px-2 py-1.5 text-[10px] text-text-tertiary">
|
||
<span className="flex min-w-0 items-center gap-1.5 truncate">
|
||
<Server className="h-3.5 w-3.5" />
|
||
{server?.trainerRoot ?? '请先启动本地训练服务'}
|
||
</span>
|
||
<Badge tone={server?.ready ? 'success' : 'warning'}>
|
||
{server?.ready ? '可用' : '离线'}
|
||
</Badge>
|
||
</div>
|
||
<Button
|
||
className="mt-2 w-full"
|
||
icon={<ExternalLink className="h-3.5 w-3.5" />}
|
||
onClick={openTuningDashboard}
|
||
>
|
||
打开自调参 Agent 工作台
|
||
</Button>
|
||
{server?.ready && !job && (
|
||
<fieldset disabled={busy} className="mt-3 space-y-2">
|
||
<Field label="训练任务">
|
||
<Select
|
||
aria-label="训练任务"
|
||
className="w-full"
|
||
value={taskId}
|
||
onChange={(event) => selectTask(event.target.value)}
|
||
>
|
||
{server.tasks.map((task) => (
|
||
<option key={task} value={task}>
|
||
{server.taskMetadata?.find((item) => item.id === task)?.name ?? task}
|
||
</option>
|
||
))}
|
||
</Select>
|
||
</Field>
|
||
<PretrainedSourceSelect
|
||
sources={server.pretrainedSources}
|
||
taskId={taskId}
|
||
value={pretrainedSourceId}
|
||
onChange={setPretrainedSourceId}
|
||
disabled={busy || uploading}
|
||
upload={{
|
||
endpoint,
|
||
token,
|
||
revision: connectionRevision,
|
||
enabled: Boolean(server.pretrainedUpload?.enabled),
|
||
onBusyChange: setUploading,
|
||
onUploaded: (source) => {
|
||
setServer(
|
||
(current) =>
|
||
current && {
|
||
...current,
|
||
pretrainedSources: [
|
||
...(current.pretrainedSources ?? []).filter((s) => s.id !== source.id),
|
||
source,
|
||
],
|
||
},
|
||
);
|
||
setPretrainedSourceId(source.id);
|
||
},
|
||
}}
|
||
/>
|
||
{metadata && (
|
||
<>
|
||
<Field label="训练地形">
|
||
<Select
|
||
aria-label="训练地形"
|
||
value={terrainPreset}
|
||
onChange={(e) => {
|
||
setTerrainPreset(e.target.value);
|
||
setCustomTerrainBoxes(undefined);
|
||
setTerrainParams({});
|
||
}}
|
||
>
|
||
{taskId !== OBSTACLE_TASK_ID && <option value="">原任务默认地形</option>}
|
||
{metadata.terrainPresets.map((preset) => (
|
||
<option key={preset} value={preset}>
|
||
{TERRAIN_LABELS[preset] ?? preset}
|
||
</option>
|
||
))}
|
||
</Select>
|
||
</Field>
|
||
<Button disabled={busy} onClick={syncMap}>
|
||
检查当前场景地图(可选)
|
||
</Button>
|
||
<p className="text-[10px] text-text-tertiary">
|
||
启动训练时会自动从全部已应用实例重新编译并校验世界AABB,无需预先同步。旋转障碍会膨胀,底板标准化为z=[-0.2,0];仅保证训练与浏览器使用相同boxes,不等于原OBB。mesh/hfield、地下结构、混合摩擦明确拒绝。
|
||
</p>
|
||
{terrainPreset === 'custom_boxes' && customTerrainBoxes && !sceneDirty && (
|
||
<p role="status">
|
||
已读取视口中 {customTerrainBoxes.actualObstacleCount}{' '}
|
||
个自定义障碍物;启动训练时会自动重新编译并校验地图
|
||
</p>
|
||
)}
|
||
{terrainPreset && terrainPreset !== 'custom_boxes' && (
|
||
<div className="grid grid-cols-2 gap-2">
|
||
{Object.entries(metadata.terrainParameters).map(([key, bounds]) => (
|
||
<NumberField
|
||
key={key}
|
||
label={PARAMETER_LABELS[key] ?? key}
|
||
value={terrainParams[key] ?? bounds.default}
|
||
min={bounds.min}
|
||
max={bounds.max}
|
||
step={bounds.integer ? 1 : 0.01}
|
||
onChange={(value) => setTerrainParams((old) => ({ ...old, [key]: value }))}
|
||
/>
|
||
))}
|
||
</div>
|
||
)}
|
||
{['rough', 'wave', 'pyramid_stairs'].includes(terrainPreset) && (
|
||
<p>训练专用 box 离散近似布局,不等于编辑器高度场。</p>
|
||
)}
|
||
{taskId === OBSTACLE_TASK_ID && (
|
||
<details open>
|
||
<summary>避障传感器高级设置</summary>
|
||
<Field label="传感器模式">
|
||
<Select
|
||
aria-label="传感器模式"
|
||
value={sensorMode}
|
||
onChange={(event) => setSensorMode(event.target.value as typeof sensorMode)}
|
||
>
|
||
<option value="single_ring_raycast">水平32射线 / 81维(默认)</option>
|
||
<option
|
||
value="multi_ring_raycast"
|
||
disabled={!metadata?.sensorModes?.includes('multi_ring_raycast')}
|
||
>
|
||
三层48射线 / 97维(非高程图)
|
||
</option>
|
||
</Select>
|
||
</Field>
|
||
{Object.entries(metadata.sensorParameters).map(([key, bounds]) => (
|
||
<NumberField
|
||
key={key}
|
||
label={PARAMETER_LABELS[key] ?? key}
|
||
value={sensorCfg[key] ?? bounds.default}
|
||
min={bounds.min}
|
||
max={bounds.max}
|
||
step={0.01}
|
||
onChange={(value) => setSensorCfg((old) => ({ ...old, [key]: value }))}
|
||
/>
|
||
))}
|
||
</details>
|
||
)}
|
||
{!metadata.browserCompatible && (
|
||
<p>此任务可训练,但浏览器不支持其观测契约,不能一键部署。</p>
|
||
)}
|
||
</>
|
||
)}
|
||
<div className="grid grid-cols-2 gap-2">
|
||
<NumberField
|
||
label="并行环境"
|
||
value={numEnvs}
|
||
min={1}
|
||
max={16384}
|
||
onChange={setNumEnvs}
|
||
/>
|
||
<NumberField
|
||
label="训练迭代"
|
||
value={maxIterations}
|
||
min={1}
|
||
max={1000000}
|
||
onChange={setMaxIterations}
|
||
/>
|
||
<NumberField
|
||
label="随机种子"
|
||
value={seed}
|
||
min={0}
|
||
max={2147483647}
|
||
onChange={setSeed}
|
||
/>
|
||
<Field label="运行名称">
|
||
<input
|
||
aria-label="运行名称"
|
||
className="field h-7 w-full px-2 text-xs text-text-primary"
|
||
value={runName}
|
||
onChange={(event) => setRunName(event.target.value)}
|
||
/>
|
||
</Field>
|
||
</div>
|
||
<div className="grid grid-cols-2 gap-2">
|
||
<Field label="计算设备">
|
||
<Select
|
||
aria-label="计算设备"
|
||
className="w-full"
|
||
value={device}
|
||
onChange={(event) => setDevice(event.target.value as TrainingDevice)}
|
||
>
|
||
<option value="gpu">GPU</option>
|
||
<option value="cpu">CPU</option>
|
||
</Select>
|
||
</Field>
|
||
<Field label="GPU 编号">
|
||
<input
|
||
aria-label="GPU 编号"
|
||
className="field h-7 w-full px-2 text-xs text-text-primary disabled:opacity-40"
|
||
value={gpuIds}
|
||
disabled={device === 'cpu'}
|
||
onChange={(event) => setGpuIds(event.target.value)}
|
||
/>
|
||
</Field>
|
||
</div>
|
||
<Field label="奖励配置">
|
||
<Select
|
||
disabled={taskId !== 'Unitree-Go2-Flat'}
|
||
aria-label="奖励配置 preset"
|
||
className="w-full"
|
||
value={rewardPresetId}
|
||
onChange={(event) => setRewardPresetId(event.target.value)}
|
||
>
|
||
<option value="">仓库默认奖励</option>
|
||
{presets
|
||
.filter((preset) => preset.taskId === 'Unitree-Go2-Flat')
|
||
.map((preset) => (
|
||
<option key={preset.id} value={preset.id}>
|
||
{preset.name}
|
||
</option>
|
||
))}
|
||
</Select>
|
||
</Field>
|
||
<Field label="实验记录">
|
||
<Select
|
||
aria-label="W&B 模式"
|
||
className="w-full"
|
||
value={wandbMode}
|
||
onChange={(event) => setWandbMode(event.target.value as WandbMode)}
|
||
>
|
||
<option value="offline">本地离线(默认,无需登录)</option>
|
||
<option value="disabled">完全禁用 W&B</option>
|
||
<option value="online">在线 W&B(需要 API Key)</option>
|
||
</Select>
|
||
</Field>
|
||
<Button
|
||
variant="primary"
|
||
className="w-full"
|
||
icon={<Play className="h-3.5 w-3.5" />}
|
||
disabled={busy || uploading || Boolean(sourceSelectionError)}
|
||
onClick={() => void start()}
|
||
>
|
||
发起本地训练
|
||
</Button>
|
||
<p className="text-[10px] leading-4 text-text-tertiary">
|
||
训练使用本地 mjlab
|
||
任务资产,不会把浏览器中的模型上传到网络。服务一次只运行一个训练任务。
|
||
</p>
|
||
</fieldset>
|
||
)}
|
||
{job && (
|
||
<div className="mt-3 rounded-lg border border-border bg-surface p-2.5">
|
||
<div className="mb-2 flex items-center justify-between gap-2">
|
||
<span className="truncate text-xs font-medium text-text-primary" title={job.id}>
|
||
{job.taskId}
|
||
</span>
|
||
<Badge
|
||
tone={
|
||
job.state === 'succeeded'
|
||
? 'success'
|
||
: job.state === 'failed' || job.state === 'cancelled'
|
||
? 'warning'
|
||
: 'accent'
|
||
}
|
||
>
|
||
{stateLabel(job.state)}
|
||
</Badge>
|
||
</div>
|
||
<PretrainedIdentity source={job.pretrained} />
|
||
{job.taskId === 'Unitree-Go2-Rough' && (
|
||
<p>234维 Rough 策略仅支持后端评测,浏览器不可加载。</p>
|
||
)}
|
||
{job.deployment?.terrain && (
|
||
<p className="text-[10px] text-text-tertiary">
|
||
导入将替换当前物理地图并启动配套策略;
|
||
{job.deployment.terrain.approximation ? '训练专用近似布局' : '配套碰撞布局'}
|
||
。请先保存场景。
|
||
</p>
|
||
)}
|
||
<ProgressBar value={job.progress} label="训练进度" />
|
||
<div className="mt-2">
|
||
<PropertyRow label="迭代" value={`${job.iteration} / ${job.maxIterations}`} />
|
||
<PropertyRow label="状态" value={job.message} />
|
||
{trainingLosses(job.logs).map(({ label, value }) => (
|
||
<PropertyRow key={label} label={label} value={value} />
|
||
))}
|
||
</div>
|
||
<TrainingMetricsPanel key={job.id} jobId={job.id} logs={job.logs} />
|
||
{job.logs.length > 0 && (
|
||
<details className="mt-2">
|
||
<summary className="cursor-pointer text-[10px] text-text-secondary">最近日志</summary>
|
||
<pre className="mt-1 max-h-36 overflow-auto whitespace-pre-wrap break-all rounded bg-app p-2 text-[9px] leading-4 text-text-tertiary">
|
||
{job.logs.slice(-40).join('\n')}
|
||
</pre>
|
||
</details>
|
||
)}
|
||
<div className="mt-3 grid grid-cols-2 gap-2">
|
||
{active ? (
|
||
<Button
|
||
variant="danger"
|
||
className="col-span-2"
|
||
icon={<Square className="h-3.5 w-3.5" />}
|
||
disabled={busy}
|
||
onClick={() => void cancel()}
|
||
>
|
||
停止训练
|
||
</Button>
|
||
) : (
|
||
<>
|
||
<Button
|
||
disabled={
|
||
busy ||
|
||
!job.artifactReady ||
|
||
job.taskId === 'Unitree-Go2-Rough' ||
|
||
(job.deployment && !job.deployment.browserCompatible)
|
||
}
|
||
icon={<Download className="h-3.5 w-3.5" />}
|
||
onClick={() => void importResult()}
|
||
>
|
||
导入策略
|
||
</Button>
|
||
<Button
|
||
onClick={() => {
|
||
setJob(undefined);
|
||
try {
|
||
localStorage.removeItem(TRAINING_JOB_KEY);
|
||
} catch {
|
||
/* ignore */
|
||
}
|
||
}}
|
||
>
|
||
新建任务
|
||
</Button>
|
||
</>
|
||
)}
|
||
</div>
|
||
</div>
|
||
)}
|
||
{error && (
|
||
<p
|
||
role="alert"
|
||
className="mt-2 break-words rounded bg-danger/10 p-2 text-[10px] leading-4 text-danger"
|
||
>
|
||
{error}
|
||
</p>
|
||
)}
|
||
</div>
|
||
);
|
||
}
|
||
|
||
function Field({ label, children }: { label: string; children: ReactNode }) {
|
||
return (
|
||
<label className="block text-[10px] text-text-tertiary">
|
||
<span className="mb-1 block">{label}</span>
|
||
{children}
|
||
</label>
|
||
);
|
||
}
|
||
function NumberField({
|
||
label,
|
||
value,
|
||
min,
|
||
max,
|
||
onChange,
|
||
step = 1,
|
||
}: {
|
||
label: string;
|
||
value: number;
|
||
min: number;
|
||
max: number;
|
||
step?: number;
|
||
onChange(value: number): void;
|
||
}) {
|
||
return (
|
||
<Field label={label}>
|
||
<input
|
||
aria-label={label}
|
||
type="number"
|
||
step={step}
|
||
className="field h-7 w-full px-2 text-xs text-text-primary"
|
||
value={value}
|
||
min={min}
|
||
max={max}
|
||
onChange={(event) => onChange(Number(event.target.value))}
|
||
/>
|
||
</Field>
|
||
);
|
||
}
|
||
|
||
const TERRAIN_LABELS: Record<string, string> = {
|
||
custom_boxes: '自定义场景碰撞布局(AABB近似)',
|
||
plane: '平地',
|
||
discrete_obstacles: '离散障碍物',
|
||
rough: '崎岖地面',
|
||
pyramid_stairs: '金字塔台阶',
|
||
wave: '波浪地形',
|
||
};
|
||
const PARAMETER_LABELS: Record<string, string> = {
|
||
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: '避障权重',
|
||
};
|