refactor: release v1.0.5 安全精简与网站发布
web-platform-ci / Standalone decision service (no cloud credentials) (push) Has been cancelled
web-platform-ci / TypeScript, lint, unit, build (push) Has been cancelled
web-platform-ci / Playwright E2E (push) Has been cancelled
lekiwi-compatibility / cpu-compatibility (push) Has been cancelled

This commit is contained in:
2026-09-29 09:55:08 +08:00
parent 7ebe9092ba
commit 26bb5634bd
64 changed files with 9048 additions and 3688 deletions
@@ -0,0 +1,96 @@
import { fireEvent, render, screen, waitFor } from '@testing-library/react';
import { beforeEach, describe, expect, it, vi } from 'vitest';
import { LocalTrainingClient } from './LocalTrainingClient';
import { LocalTrainingPanel } from './LocalTrainingPanel';
import type { TrainingJob, TrainingRequest } from './types';
const flatRequest: TrainingRequest = {
taskId: 'Unitree-Go2-Flat',
numEnvs: 4096,
maxIterations: 2000,
seed: 42,
runName: 'web',
device: 'gpu',
gpuIds: [0],
wandbMode: 'offline',
rewardPresetId: undefined,
};
const queuedJob: TrainingJob = {
id: 'a'.repeat(32),
taskId: 'Unitree-Go2-Flat',
state: 'queued',
createdAt: '2026-09-28T00:00:00Z',
iteration: 0,
maxIterations: 2000,
progress: 0,
message: '等待启动',
logs: [],
artifactReady: false,
};
beforeEach(() => {
vi.restoreAllMocks();
localStorage.clear();
sessionStorage.clear();
vi.spyOn(LocalTrainingClient.prototype, 'health').mockResolvedValue({
version: '0.1.0',
ready: true,
trainerRoot: '/测试训练目录',
python: '/测试虚拟环境/python',
tasks: ['Unitree-Go2-Flat'],
});
vi.spyOn(LocalTrainingClient.prototype, 'presets').mockResolvedValue([]);
vi.spyOn(LocalTrainingClient.prototype, 'start').mockResolvedValue(queuedJob);
vi.spyOn(LocalTrainingClient.prototype, 'job').mockResolvedValue(queuedJob);
});
async function connect() {
render(<LocalTrainingPanel onPolicyReady={vi.fn()} />);
fireEvent.change(screen.getByLabelText('训练服务访问令牌'), {
target: { value: '仅供单测的令牌' },
});
fireEvent.click(screen.getByRole('button', { name: /^连接$/ }));
const start = await screen.findByRole('button', { name: '发起本地训练' });
await waitFor(() => expect(start).toBeEnabled());
return start;
}
function submittedRequest() {
const calls = vi.mocked(LocalTrainingClient.prototype.start).mock.calls;
expect(calls).toHaveLength(1);
return calls[0][0];
}
describe('训练面板拆分前的请求表征', () => {
it('完整保留默认 Flat 请求字段,不隐式添加地形、传感器或预训练配置', async () => {
fireEvent.click(await connect());
await screen.findByText('排队中');
expect(submittedRequest()).toStrictEqual(flatRequest);
});
it('CPU 请求使用空 GPU 列表,即使输入框保留非法 GPU 文本', async () => {
const start = await connect();
fireEvent.change(screen.getByLabelText('GPU 编号'), { target: { value: '不是编号' } });
fireEvent.change(screen.getByLabelText('计算设备'), { target: { value: 'cpu' } });
fireEvent.click(start);
await screen.findByText('排队中');
expect(submittedRequest()).toStrictEqual({ ...flatRequest, device: 'cpu', gpuIds: [] });
});
it('GPU 输入接受空白与逗号,保留原顺序和重复编号', async () => {
const start = await connect();
fireEvent.change(screen.getByLabelText('GPU 编号'), { target: { value: ' 2, 1 2 ' } });
fireEvent.click(start);
await screen.findByText('排队中');
expect(submittedRequest()).toStrictEqual({ ...flatRequest, gpuIds: [2, 1, 2] });
});
it.each(['-1', '1.5', '未知'])('非法 GPU 输入 %s 不提交训练请求', async (value) => {
const start = await connect();
fireEvent.change(screen.getByLabelText('GPU 编号'), { target: { value } });
fireEvent.click(start);
expect(await screen.findByRole('alert')).toHaveTextContent('GPU 编号必须是非负整数');
expect(LocalTrainingClient.prototype.start).not.toHaveBeenCalled();
expect(start).toBeEnabled();
});
});
+119 -490
View File
@@ -6,10 +6,11 @@ import {
} from '../mobile/training';
import { MOBILE_TASK } from '../mobile/RobotDescriptor';
import type { TrainingStage } from '../mobile/TaskKernel';
import { PretrainedIdentity, PretrainedSourceSelect } from './PretrainedSourceSelect';
import { PretrainedSourceSelect } from './PretrainedSourceSelect';
import { pretrainedSelectionError } from './pretrainedSelection';
import { TrainingMetricsPanel } from './TrainingMetricsPanel';
import { trainingLosses } from './trainingLosses';
import { TrainingJobSection } from './TrainingJobSection';
import { ACTIVE_STATES, stateLabel } from './trainingPresentation';
import { MobileTaskFields, TerrainTaskFields } from './TrainingTaskFields';
import type { PolicyDeployment, TrainingTerrain } from '../rl/deployment';
import {
OBSTACLE_TASK_ID,
@@ -19,18 +20,17 @@ import {
} 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,
CollapsibleSection,
Tooltip,
} from '../components/ui';
import { useEffect, useRef, useState } from 'react';
import { ExternalLink, Link, Play, Server } from 'lucide-react';
import { Badge, Button, PropertyRow, Select, CollapsibleSection } from '../components/ui';
import { Field, NumberField } from './TrainingFields';
import { LocalTrainingClient } from './LocalTrainingClient';
import {
createTrainingRequest,
parseTrainingGpuIds,
trainingTaskDefaults,
validateTrainingParameters,
} from './trainingForm';
import type {
RewardPreset,
TrainingDevice,
@@ -48,19 +48,9 @@ import {
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,
@@ -128,19 +118,20 @@ export function LocalTrainingPanel({
? undefined
: pretrainedSelectionError(server?.pretrainedSources, taskId, pretrainedSourceId);
const selectTask = (id: string) => {
const defaults = trainingTaskDefaults(id);
setTaskId(id);
setMobileStage('navigate');
setSourceJobId('');
if (isMobileTrainingTask(id)) setPretrainedSourceId('');
setNumEnvs(isMobileTrainingTask(id) ? 1 : 4096);
setMaxIterations(isMobileTrainingTask(id) ? 1000 : 2000);
setDevice(isMobileTrainingTask(id) ? 'cpu' : 'gpu');
setNumEnvs(defaults.numEnvs);
setMaxIterations(defaults.maxIterations);
setDevice(defaults.device);
setCustomTerrainBoxes(undefined);
setRewardPresetId('');
setTerrainParams({});
setSensorCfg({});
setSensorMode('single_ring_raycast');
setTerrainPreset(id === OBSTACLE_TASK_ID ? 'discrete_obstacles' : '');
setTerrainPreset(defaults.terrainPreset);
};
const robotId = mobileTraining?.robotId;
const [observedRobotId, setObservedRobotId] = useState<string>();
@@ -397,35 +388,8 @@ export function LocalTrainingPanel({
!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('安全距离必须小于探测距离');
const ids = parseTrainingGpuIds(device, gpuIds);
validateTrainingParameters(terrainParams, sensorCfg, metadata);
let resolvedCustomTerrain = customTerrainBoxes;
let resolvedTerrainParams = terrainParams;
if (terrainPreset === 'custom_boxes') {
@@ -449,38 +413,37 @@ export function LocalTrainingPanel({
throw new Error('场景变体与任务不匹配');
mobilePackageId = uploaded.id;
}
const next = await client.start({
taskId,
numEnvs,
maxIterations,
seed,
runName,
device,
gpuIds: ids,
wandbMode: mobile ? 'disabled' : wandbMode,
...(mobile
? {
mobilePackageId,
mobileParams: {
rolloutSteps,
objectPosition,
goalPosition,
stage: mobileStage,
positionJitter,
evaluationEpisodes,
navigationBootstrapSteps,
...(sourceJobId ? { sourceJobId } : {}),
},
}
: {}),
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 } }
: {}),
});
const next = await client.start(
createTrainingRequest({
taskId,
numEnvs,
maxIterations,
seed,
runName,
device,
gpuIds: ids,
wandbMode,
mobile,
mobilePackageId,
mobileParams: {
rolloutSteps,
objectPosition,
goalPosition,
stage: mobileStage,
positionJitter,
evaluationEpisodes,
navigationBootstrapSteps,
sourceJobId,
},
rewardPresetId,
pretrainedSourceId,
terrainPreset,
terrainParams: resolvedTerrainParams,
customTerrainBoxes: resolvedCustomTerrain,
sensorCfg,
sensorMode,
}),
);
setJob(next);
try {
localStorage.setItem(TRAINING_JOB_KEY, next.id);
@@ -533,7 +496,6 @@ export function LocalTrainingPanel({
setBusy(false);
}
};
const active = Boolean(job && ACTIVE_STATES.has(job.state));
const summary = error ? '错误' : job ? stateLabel(job.state) : server?.ready ? '已连接' : '离线';
useEffect(() => {
onStatusChange?.(summary);
@@ -663,206 +625,47 @@ export function LocalTrainingPanel({
/>
)}
{metadata && !mobile && (
<>
<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}
title="从已应用实例编译并校验世界 AABB;旋转障碍会膨胀为轴对齐包围盒"
onClick={syncMap}
>
同步场景地图
</Button>
{terrainPreset === 'custom_boxes' && customTerrainBoxes && !sceneDirty && (
<p role="status">
已读取视口中 {customTerrainBoxes.actualObstacleCount}{' '}
个自定义障碍物;启动训练时会自动重新编译并校验地图
</p>
)}
{terrainPreset && terrainPreset !== 'custom_boxes' && (
<CollapsibleSection
title="地形详细参数"
defaultOpen={false}
keepMounted
forceOpen={Boolean(error)}
>
<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}
defaultValue={bounds.default}
min={bounds.min}
max={bounds.max}
step={bounds.integer ? 1 : 0.01}
onChange={(value) => setTerrainParams((old) => ({ ...old, [key]: value }))}
/>
))}
</div>
</CollapsibleSection>
)}
{['rough', 'wave', 'pyramid_stairs'].includes(terrainPreset) && (
<div title="训练使用 box 离散近似布局,不等于编辑器高度场">
<p className="text-xs text-warning">
近似地形:训练使用 box 离散布局,非编辑器高度场。
</p>
</div>
)}
{taskId === OBSTACLE_TASK_ID && (
<p className="text-xs text-text-secondary">
{sensorMode === 'single_ring_raycast' ? '水平32射线 · 81维' : '三层48射线 · 97维'}
</p>
)}
{taskId === OBSTACLE_TASK_ID && (
<CollapsibleSection
title="避障传感器高级设置"
defaultOpen={false}
keepMounted
forceOpen={Boolean(error)}
>
<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}
defaultValue={bounds.default}
min={bounds.min}
max={bounds.max}
step={0.01}
onChange={(value) => setSensorCfg((old) => ({ ...old, [key]: value }))}
/>
))}
</CollapsibleSection>
)}
{!metadata.browserCompatible && (
<div title="浏览器不支持当前观测契约,策略仅可在后端评测">
<p className="text-xs text-warning">仅后端部署:浏览器不支持当前观测契约。</p>
</div>
)}
</>
<TerrainTaskFields
taskId={taskId}
metadata={metadata}
terrainPreset={terrainPreset}
customTerrainBoxes={customTerrainBoxes}
terrainParams={terrainParams}
sensorMode={sensorMode}
sensorCfg={sensorCfg}
busy={busy}
sceneDirty={sceneDirty}
error={error}
setTerrainPreset={setTerrainPreset}
setCustomTerrainBoxes={setCustomTerrainBoxes}
setTerrainParams={setTerrainParams}
setSensorMode={setSensorMode}
setSensorCfg={setSensorCfg}
syncMap={syncMap}
/>
)}
{mobile && (
<>
<p className="text-xs text-text-secondary">
变体:{MOBILE_TRAINING_TASKS[taskId]} · 控制步长 {MOBILE_TASK.controlDt}s ·{' '}
{MOBILE_TASK.observationSize} → 12。原生 CPU 物理,设备选项控制 PPO
网络。场景自动同步,无需下载训练包。
</p>
<Field label="移动操作训练阶段">
<Select
aria-label="移动操作训练阶段"
value={mobileStage}
onChange={(e) => setMobileStage(e.target.value as TrainingStage)}
>
<option value="navigate">1 · 底盘接近(机械臂保持)</option>
<option value="reach">2 · 末端接近(先导航,再伸臂)</option>
<option value="pick-place">3 · 抓取放置</option>
</Select>
</Field>
<Field label="接续作业 ID(留空从头训练导航)">
<input
aria-label="接续作业 ID"
className="field h-7 w-full px-2 text-xs"
value={sourceJobId}
onChange={(e) => setSourceJobId(e.target.value.trim())}
/>
</Field>
<p className="text-xs text-warning">
导航接近位为物体前方 0.3 m、朝向世界 +X,并非放置目标点。机械臂目标限速{' '}
{MOBILE_TASK.armSpeedLimit} rad/s,实测超速 {MOBILE_TASK.jointSpeedStop} rad/s
安全终止。升级阶段需前一阶段至少 10 回合评估、成功率 ≥80%、无安全终止。
</p>
<NumberField
label="位置随机范围 m"
value={positionJitter}
min={0}
max={0.3}
step={0.01}
onChange={setPositionJitter}
/>
<NumberField
label="独立评估回合"
value={evaluationEpisodes}
min={2}
max={64}
onChange={setEvaluationEpisodes}
/>
<NumberField
label="导航启动示教步数(仅初训)"
value={navigationBootstrapSteps}
min={0}
max={10000}
onChange={setNavigationBootstrapSteps}
/>
<p className="text-xs text-text-secondary">
初训可先模仿闭环底盘控制器,再用 PPO 微调;0 表示纯
PPO。导出只包含训练后的神经网络,不包含示教控制器。
</p>
<NumberField
label="每环境采样步数"
value={rolloutSteps}
min={8}
max={4096}
onChange={setRolloutSteps}
/>
{(
[
['物体', objectPosition, setObjectPosition],
['目标', goalPosition, setGoalPosition],
] as const
).map(([label, position, setPosition]) => (
<div className="grid grid-cols-3 gap-2" key={label}>
{['X', 'Y', 'Z'].map((axis, i) => (
<NumberField
key={axis}
label={`${label} ${axis}`}
value={position[i]}
min={i === 2 ? MOBILE_TASK.objectStart[2] : -MOBILE_TASK.positionScale}
max={MOBILE_TASK.positionScale}
step={0.01}
onChange={(value) =>
setPosition((old) => old.map((v, j) => (j === i ? value : v)))
}
/>
))}
</div>
))}
<p className="text-xs">
总采样步数:{numEnvs * maxIterations * rolloutSteps};训练不保证学会抓取。
</p>
</>
<MobileTaskFields
taskId={taskId}
numEnvs={numEnvs}
maxIterations={maxIterations}
stage={mobileStage}
sourceJobId={sourceJobId}
positionJitter={positionJitter}
evaluationEpisodes={evaluationEpisodes}
navigationBootstrapSteps={navigationBootstrapSteps}
rolloutSteps={rolloutSteps}
objectPosition={objectPosition}
goalPosition={goalPosition}
setMobileStage={setMobileStage}
setSourceJobId={setSourceJobId}
setPositionJitter={setPositionJitter}
setEvaluationEpisodes={setEvaluationEpisodes}
setNavigationBootstrapSteps={setNavigationBootstrapSteps}
setRolloutSteps={setRolloutSteps}
setObjectPosition={setObjectPosition}
setGoalPosition={setGoalPosition}
/>
)}
<div className="grid grid-cols-2 gap-2">
<NumberField
@@ -970,209 +773,35 @@ export function LocalTrainingPanel({
</fieldset>
)}
{job && (
<div className="mt-3 border-t border-border pt-3">
<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} />
{mobileDeployment?.evaluation && (
<div className="mt-2 text-xs">
<PropertyRow label="阶段" value={mobileDeployment.trainingStage} />
<PropertyRow
label="独立评估成功率"
value={`${(mobileDeployment.evaluation.successRate * 100).toFixed(1)}% / ${mobileDeployment.evaluation.episodes} 回合`}
/>
<PropertyRow
label="实测关节峰值"
value={`${mobileDeployment.evaluation.maxJointVelocity.toFixed(3)} rad/s`}
/>
<PropertyRow label="安全终止次数" value={mobileDeployment.evaluation.safetyStops} />
<p className="text-warning">
导出成功不代表策略达标;未达标策略导入仅用于调试,请先同阶段续训。
</p>
</div>
)}
{job.taskId === 'Unitree-Go2-Rough' && (
<div>
<p className="text-xs text-warning">
仅后端评测:234 维 Rough 策略无法在浏览器加载。
</p>
</div>
)}
{job.deployment?.terrain && (
<div>
<p className="text-xs text-warning">
导入会替换当前物理地图并启动配套策略,请先保存场景。
</p>
<Badge tone="warning">
{job.deployment.terrain.approximation ? '近似碰撞布局' : '配套碰撞布局'}
</Badge>
</div>
)}
<ProgressBar value={job.progress} label="训练进度" />
<div className="mt-2">
<PropertyRow label="迭代" value={`${job.iteration} / ${job.maxIterations}`} />
<p role="status" className="break-words text-xs text-text-secondary">
{job.message}
</p>
{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-xs text-text-secondary">最近日志</summary>
<pre className="mt-1 max-h-36 overflow-auto whitespace-pre-wrap break-all rounded bg-app p-2 text-xs leading-4 text-text-tertiary">
{job.logs.slice(-40).join('\n')}
</pre>
</details>
)}
{mobileDeployment && job.state === 'succeeded' && (
<Button
disabled={busy}
onClick={() => {
setSourceJobId(job.id);
setMobileStage(mobileDeployment.trainingStage ?? 'navigate');
if (mobileDeployment.resetOptions) {
setObjectPosition(mobileDeployment.resetOptions.object);
setGoalPosition(mobileDeployment.resetOptions.goal);
}
setPositionJitter(mobileDeployment.trainingParams?.positionJitter ?? 0.1);
setJob(undefined);
}}
>
接续此作业(保留权重)
</Button>
)}
<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
disabled={busy}
onClick={() => {
setJob(undefined);
setSourceJobId('');
setMobileStage('navigate');
try {
localStorage.removeItem(TRAINING_JOB_KEY);
} catch {
/* ignore */
}
}}
>
新建任务
</Button>
</>
)}
</div>
</div>
<TrainingJobSection
job={job}
mobileDeployment={mobileDeployment}
busy={busy}
onContinue={() => {
if (!mobileDeployment) return;
setSourceJobId(job.id);
setMobileStage(mobileDeployment.trainingStage ?? 'navigate');
if (mobileDeployment.resetOptions) {
setObjectPosition(mobileDeployment.resetOptions.object);
setGoalPosition(mobileDeployment.resetOptions.goal);
}
setPositionJitter(mobileDeployment.trainingParams?.positionJitter ?? 0.1);
setJob(undefined);
}}
onCancel={cancel}
onImport={importResult}
onNew={() => {
setJob(undefined);
setSourceJobId('');
setMobileStage('navigate');
try {
localStorage.removeItem(TRAINING_JOB_KEY);
} catch {
/* ignore */
}
}}
/>
)}
</div>
);
}
function Field({ label, children }: { label: string; children: ReactNode }) {
return (
<label className="block text-xs text-text-tertiary">
<span className="mb-1 block">{label}</span>
{children}
</label>
);
}
function NumberField({
label,
value,
min,
max,
onChange,
step = 1,
defaultValue,
}: {
label: string;
value: number;
min: number;
max: number;
step?: number;
defaultValue?: number;
onChange(value: number): void;
}) {
return (
<Field label={label}>
<Tooltip
className="w-full"
content={`范围 ${min}–${max}${defaultValue === undefined ? '' : `;默认 ${defaultValue}`}`}
>
<input
aria-label={label}
type="number"
step={step}
className="field h-7 w-full px-2 text-xs tabular-nums text-text-primary"
value={value}
min={min}
max={max}
onChange={(event) => onChange(Number(event.target.value))}
/>
</Tooltip>
</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: '避障权重',
};
@@ -0,0 +1,49 @@
import type { ReactNode } from 'react';
import { Tooltip } from '../components/ui';
export function Field({ label, children }: { label: string; children: ReactNode }) {
return (
<label className="block text-xs text-text-tertiary">
<span className="mb-1 block">{label}</span>
{children}
</label>
);
}
export function NumberField({
label,
value,
min,
max,
onChange,
step = 1,
defaultValue,
}: {
label: string;
value: number;
min: number;
max: number;
step?: number;
defaultValue?: number;
onChange(value: number): void;
}) {
return (
<Field label={label}>
<Tooltip
className="w-full"
content={`范围 ${min}–${max}${defaultValue === undefined ? '' : `;默认 ${defaultValue}`}`}
>
<input
aria-label={label}
type="number"
step={step}
className="field h-7 w-full px-2 text-xs tabular-nums text-text-primary"
value={value}
min={min}
max={max}
onChange={(event) => onChange(Number(event.target.value))}
/>
</Tooltip>
</Field>
);
}
@@ -0,0 +1,68 @@
import { fireEvent, render, screen } from '@testing-library/react';
import { describe, expect, it, vi } from 'vitest';
import { TrainingJobSection } from './TrainingJobSection';
import type { TrainingJob } from './types';
const job: TrainingJob = {
id: '测试作业',
taskId: 'Unitree-Go2-Flat',
state: 'succeeded',
createdAt: '2026-09-28T00:00:00Z',
iteration: 2,
maxIterations: 2,
progress: 1,
message: '训练已完成',
logs: [],
artifactReady: true,
};
const actions = () => ({
onContinue: vi.fn(),
onCancel: vi.fn(),
onImport: vi.fn(),
onNew: vi.fn(),
});
describe('训练作业展示保持操作边界', () => {
it('排队和运行时只提供停止入口,busy 阻止重复操作', () => {
const callbacks = actions();
const { rerender } = render(
<TrainingJobSection job={{ ...job, state: 'queued' }} busy={false} {...callbacks} />,
);
fireEvent.click(screen.getByRole('button', { name: '停止训练' }));
expect(callbacks.onCancel).toHaveBeenCalledOnce();
expect(screen.queryByRole('button', { name: '导入策略' })).not.toBeInTheDocument();
rerender(<TrainingJobSection job={{ ...job, state: 'running' }} busy {...callbacks} />);
expect(screen.getByRole('button', { name: '停止训练' })).toBeDisabled();
});
it('完成后导入/新建委托父级处理,切换到 Rough 后拒绝导入', () => {
const callbacks = actions();
const { rerender } = render(<TrainingJobSection job={job} busy={false} {...callbacks} />);
fireEvent.click(screen.getByRole('button', { name: '导入策略' }));
fireEvent.click(screen.getByRole('button', { name: '新建任务' }));
expect(callbacks.onImport).toHaveBeenCalledOnce();
expect(callbacks.onNew).toHaveBeenCalledOnce();
rerender(
<TrainingJobSection
job={{ ...job, taskId: 'Unitree-Go2-Rough' }}
busy={false}
{...callbacks}
/>,
);
expect(screen.getByRole('button', { name: '导入策略' })).toBeDisabled();
expect(screen.getByText(/234 维 Rough/)).toBeInTheDocument();
});
it('无成果物时导入禁用,原样保留最后 40 行日志', () => {
const logs = Array.from({ length: 42 }, (_, i) => `第${i}行`);
const { container } = render(
<TrainingJobSection
job={{ ...job, artifactReady: false, logs }}
busy={false}
{...actions()}
/>,
);
expect(screen.getByRole('button', { name: '导入策略' })).toBeDisabled();
expect(container.querySelector('pre')?.textContent).toBe(logs.slice(-40).join('\n'));
});
});
@@ -0,0 +1,136 @@
import { Download, Square } from 'lucide-react';
import { Badge, Button, ProgressBar, PropertyRow } from '../components/ui';
import type { MobileDeployment } from '../mobile/training';
import { PretrainedIdentity } from './PretrainedSourceSelect';
import { TrainingMetricsPanel } from './TrainingMetricsPanel';
import { trainingLosses } from './trainingLosses';
import { ACTIVE_STATES, stateLabel } from './trainingPresentation';
import type { TrainingJob } from './types';
export function TrainingJobSection({
job,
mobileDeployment,
busy,
onContinue,
onCancel,
onImport,
onNew,
}: {
job: TrainingJob;
mobileDeployment?: MobileDeployment;
busy: boolean;
onContinue(): void;
onCancel(): void | Promise<void>;
onImport(): void | Promise<void>;
onNew(): void;
}) {
const active = ACTIVE_STATES.has(job.state);
return (
<div className="mt-3 border-t border-border pt-3">
<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} />
{mobileDeployment?.evaluation && (
<div className="mt-2 text-xs">
<PropertyRow label="阶段" value={mobileDeployment.trainingStage} />
<PropertyRow
label="独立评估成功率"
value={`${(mobileDeployment.evaluation.successRate * 100).toFixed(1)}% / ${mobileDeployment.evaluation.episodes} 回合`}
/>
<PropertyRow
label="实测关节峰值"
value={`${mobileDeployment.evaluation.maxJointVelocity.toFixed(3)} rad/s`}
/>
<PropertyRow label="安全终止次数" value={mobileDeployment.evaluation.safetyStops} />
<p className="text-warning">
导出成功不代表策略达标;未达标策略导入仅用于调试,请先同阶段续训。
</p>
</div>
)}
{job.taskId === 'Unitree-Go2-Rough' && (
<div>
<p className="text-xs text-warning">仅后端评测:234 维 Rough 策略无法在浏览器加载。</p>
</div>
)}
{job.deployment?.terrain && (
<div>
<p className="text-xs text-warning">
导入会替换当前物理地图并启动配套策略,请先保存场景。
</p>
<Badge tone="warning">
{job.deployment.terrain.approximation ? '近似碰撞布局' : '配套碰撞布局'}
</Badge>
</div>
)}
<ProgressBar value={job.progress} label="训练进度" />
<div className="mt-2">
<PropertyRow label="迭代" value={`${job.iteration} / ${job.maxIterations}`} />
<p role="status" className="break-words text-xs text-text-secondary">
{job.message}
</p>
{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-xs text-text-secondary">最近日志</summary>
<pre className="mt-1 max-h-36 overflow-auto whitespace-pre-wrap break-all rounded bg-app p-2 text-xs leading-4 text-text-tertiary">
{job.logs.slice(-40).join('\n')}
</pre>
</details>
)}
{mobileDeployment && job.state === 'succeeded' && (
<Button disabled={busy} onClick={onContinue}>
接续此作业(保留权重)
</Button>
)}
<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 onCancel()}
>
停止训练
</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 onImport()}
>
导入策略
</Button>
<Button disabled={busy} onClick={onNew}>
新建任务
</Button>
</>
)}
</div>
</div>
);
}
@@ -0,0 +1,281 @@
import type { Dispatch, SetStateAction } from 'react';
import { Button, CollapsibleSection, Select } from '../components/ui';
import { MOBILE_TASK } from '../mobile/RobotDescriptor';
import { MOBILE_TRAINING_TASKS, type MobileTrainingParams } from '../mobile/training';
import type { TrainingStage } from '../mobile/TaskKernel';
import { OBSTACLE_TASK_ID, type TrainingTerrain } from '../rl/deployment';
import type { TrainingTaskMetadata } from './types';
import { Field, NumberField } from './TrainingFields';
import { PARAMETER_LABELS, TERRAIN_LABELS } from './trainingPresentation';
type SensorMode = 'single_ring_raycast' | 'multi_ring_raycast';
export function TerrainTaskFields({
taskId,
metadata,
terrainPreset,
customTerrainBoxes,
terrainParams,
sensorMode,
sensorCfg,
busy,
sceneDirty,
error,
setTerrainPreset,
setCustomTerrainBoxes,
setTerrainParams,
setSensorMode,
setSensorCfg,
syncMap,
}: {
taskId: string;
metadata: TrainingTaskMetadata;
terrainPreset: string;
customTerrainBoxes?: TrainingTerrain;
terrainParams: Record<string, number>;
sensorMode: SensorMode;
sensorCfg: Record<string, number>;
busy: boolean;
sceneDirty: boolean;
error?: string;
setTerrainPreset(value: string): void;
setCustomTerrainBoxes(value: TrainingTerrain | undefined): void;
setTerrainParams: Dispatch<SetStateAction<Record<string, number>>>;
setSensorMode(value: SensorMode): void;
setSensorCfg: Dispatch<SetStateAction<Record<string, number>>>;
syncMap(): void;
}) {
return (
<>
<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}
title="从已应用实例编译并校验世界 AABB;旋转障碍会膨胀为轴对齐包围盒"
onClick={syncMap}
>
同步场景地图
</Button>
{terrainPreset === 'custom_boxes' && customTerrainBoxes && !sceneDirty && (
<p role="status">
已读取视口中 {customTerrainBoxes.actualObstacleCount}{' '}
个自定义障碍物;启动训练时会自动重新编译并校验地图
</p>
)}
{terrainPreset && terrainPreset !== 'custom_boxes' && (
<CollapsibleSection
title="地形详细参数"
defaultOpen={false}
keepMounted
forceOpen={Boolean(error)}
>
<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}
defaultValue={bounds.default}
min={bounds.min}
max={bounds.max}
step={bounds.integer ? 1 : 0.01}
onChange={(value) => setTerrainParams((old) => ({ ...old, [key]: value }))}
/>
))}
</div>
</CollapsibleSection>
)}
{['rough', 'wave', 'pyramid_stairs'].includes(terrainPreset) && (
<div title="训练使用 box 离散近似布局,不等于编辑器高度场">
<p className="text-xs text-warning">近似地形:训练使用 box 离散布局,非编辑器高度场。</p>
</div>
)}
{taskId === OBSTACLE_TASK_ID && (
<p className="text-xs text-text-secondary">
{sensorMode === 'single_ring_raycast' ? '水平32射线 · 81维' : '三层48射线 · 97维'}
</p>
)}
{taskId === OBSTACLE_TASK_ID && (
<CollapsibleSection
title="避障传感器高级设置"
defaultOpen={false}
keepMounted
forceOpen={Boolean(error)}
>
<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}
defaultValue={bounds.default}
min={bounds.min}
max={bounds.max}
step={0.01}
onChange={(value) => setSensorCfg((old) => ({ ...old, [key]: value }))}
/>
))}
</CollapsibleSection>
)}
{!metadata.browserCompatible && (
<div title="浏览器不支持当前观测契约,策略仅可在后端评测">
<p className="text-xs text-warning">仅后端部署:浏览器不支持当前观测契约。</p>
</div>
)}
</>
);
}
export function MobileTaskFields({
taskId,
numEnvs,
maxIterations,
stage: mobileStage,
sourceJobId,
positionJitter,
evaluationEpisodes,
navigationBootstrapSteps,
rolloutSteps,
objectPosition,
goalPosition,
setMobileStage,
setSourceJobId,
setPositionJitter,
setEvaluationEpisodes,
setNavigationBootstrapSteps,
setRolloutSteps,
setObjectPosition,
setGoalPosition,
}: Required<MobileTrainingParams> & {
taskId: string;
numEnvs: number;
maxIterations: number;
setMobileStage(value: TrainingStage): void;
setSourceJobId(value: string): void;
setPositionJitter(value: number): void;
setEvaluationEpisodes(value: number): void;
setNavigationBootstrapSteps(value: number): void;
setRolloutSteps(value: number): void;
setObjectPosition: Dispatch<SetStateAction<number[]>>;
setGoalPosition: Dispatch<SetStateAction<number[]>>;
}) {
return (
<>
<p className="text-xs text-text-secondary">
变体:{MOBILE_TRAINING_TASKS[taskId]} · 控制步长 {MOBILE_TASK.controlDt}s ·{' '}
{MOBILE_TASK.observationSize} → 12。原生 CPU 物理,设备选项控制 PPO
网络。场景自动同步,无需下载训练包。
</p>
<Field label="移动操作训练阶段">
<Select
aria-label="移动操作训练阶段"
value={mobileStage}
onChange={(e) => setMobileStage(e.target.value as TrainingStage)}
>
<option value="navigate">1 · 底盘接近(机械臂保持)</option>
<option value="reach">2 · 末端接近(先导航,再伸臂)</option>
<option value="pick-place">3 · 抓取放置</option>
</Select>
</Field>
<Field label="接续作业 ID(留空从头训练导航)">
<input
aria-label="接续作业 ID"
className="field h-7 w-full px-2 text-xs"
value={sourceJobId}
onChange={(e) => setSourceJobId(e.target.value.trim())}
/>
</Field>
<p className="text-xs text-warning">
导航接近位为物体前方 0.3 m、朝向世界 +X,并非放置目标点。机械臂目标限速{' '}
{MOBILE_TASK.armSpeedLimit} rad/s,实测超速 {MOBILE_TASK.jointSpeedStop} rad/s
安全终止。升级阶段需前一阶段至少 10 回合评估、成功率 ≥80%、无安全终止。
</p>
<NumberField
label="位置随机范围 m"
value={positionJitter}
min={0}
max={0.3}
step={0.01}
onChange={setPositionJitter}
/>
<NumberField
label="独立评估回合"
value={evaluationEpisodes}
min={2}
max={64}
onChange={setEvaluationEpisodes}
/>
<NumberField
label="导航启动示教步数(仅初训)"
value={navigationBootstrapSteps}
min={0}
max={10000}
onChange={setNavigationBootstrapSteps}
/>
<p className="text-xs text-text-secondary">
初训可先模仿闭环底盘控制器,再用 PPO 微调;0 表示纯
PPO。导出只包含训练后的神经网络,不包含示教控制器。
</p>
<NumberField
label="每环境采样步数"
value={rolloutSteps}
min={8}
max={4096}
onChange={setRolloutSteps}
/>
{(
[
['物体', objectPosition, setObjectPosition],
['目标', goalPosition, setGoalPosition],
] as const
).map(([label, position, setPosition]) => (
<div className="grid grid-cols-3 gap-2" key={label}>
{['X', 'Y', 'Z'].map((axis, i) => (
<NumberField
key={axis}
label={`${label} ${axis}`}
value={position[i]}
min={i === 2 ? MOBILE_TASK.objectStart[2] : -MOBILE_TASK.positionScale}
max={MOBILE_TASK.positionScale}
step={0.01}
onChange={(value) => setPosition((old) => old.map((v, j) => (j === i ? value : v)))}
/>
))}
</div>
))}
<p className="text-xs">
总采样步数:{numEnvs * maxIterations * rolloutSteps};训练不保证学会抓取。
</p>
</>
);
}
@@ -0,0 +1,168 @@
import { describe, expect, it } from 'vitest';
import { OBSTACLE_TASK_ID } from '../rl/deployment';
import {
createTrainingRequest,
parseTrainingGpuIds,
trainingTaskDefaults,
validateTrainingParameters,
} from './trainingForm';
import type { TrainingTaskMetadata } from './types';
type Form = Parameters<typeof createTrainingRequest>[0];
const form = (patch: Partial<Form> = {}): Form => ({
taskId: 'Unitree-Go2-Flat',
numEnvs: 4096,
maxIterations: 2000,
seed: 42,
runName: 'web',
device: 'gpu',
gpuIds: [0],
wandbMode: 'offline',
mobile: false,
mobileParams: {
rolloutSteps: 128,
objectPosition: [0.4, 0, 0.03],
goalPosition: [0.4, 0.5, 0.03],
stage: 'navigate',
sourceJobId: '',
},
rewardPresetId: '',
pretrainedSourceId: '',
terrainPreset: '',
terrainParams: {},
customTerrainBoxes: undefined,
sensorCfg: {},
sensorMode: 'single_ring_raycast',
...patch,
});
const metadata: TrainingTaskMetadata = {
id: OBSTACLE_TASK_ID,
name: '避障任务',
browserCompatible: true,
terrainPresets: ['discrete_obstacles'],
terrainParameters: {
obstacle_count: { min: 1, max: 50, default: 25, integer: true },
obstacle_height_min: { min: 0.1, max: 1, default: 0.2 },
obstacle_height_max: { min: 0.1, max: 1, default: 0.6 },
},
sensorTypes: ['raycast'],
sensorParameters: {
safetyDistance: { min: 0.1, max: 4, default: 0.5 },
maxDistance: { min: 1, max: 8, default: 4 },
},
mapSyncScope: 'custom_boxes',
};
describe('训练表单的纯转换边界', () => {
it.each(['MobileManipulator-LeKiwi-v1', 'MobileManipulator-LeKiwi-Bundle'])(
'移动任务 %s 的 CPU 默认值不变',
(id) => {
expect(trainingTaskDefaults(id)).toEqual({
numEnvs: 1,
maxIterations: 1000,
device: 'cpu',
terrainPreset: '',
});
},
);
it('Go2 默认 GPU 参数不变,只有避障任务默认使用离散障碍物', () => {
const defaults = { numEnvs: 4096, maxIterations: 2000, device: 'gpu', terrainPreset: '' };
expect(trainingTaskDefaults('Unitree-Go2-Flat')).toEqual(defaults);
expect(trainingTaskDefaults('Unitree-Go2-Rough')).toEqual(defaults);
expect(trainingTaskDefaults(OBSTACLE_TASK_ID)).toEqual({
...defaults,
terrainPreset: 'discrete_obstacles',
});
});
it('GPU 保留顺序和重复值,CPU 忽略文本', () => {
expect(parseTrainingGpuIds('gpu', '2, 1 2')).toEqual([2, 1, 2]);
expect(parseTrainingGpuIds('gpu', ' , ')).toEqual([]);
expect(parseTrainingGpuIds('cpu', '非法')).toEqual([]);
});
it.each(['-1', '1.5', '未知'])('拒绝非法 GPU 编号 %s', (text) => {
expect(() => parseTrainingGpuIds('gpu', text)).toThrow('GPU 编号必须是非负整数');
});
it('保持先地形后传感器的校验顺序,拒绝未知参数、非有限数和非整数', () => {
expect(() => validateTrainingParameters({}, {}, undefined)).not.toThrow();
expect(() => validateTrainingParameters({ missing: 0 }, { other: 0 }, metadata)).toThrow(
'参数 missing 超出允许范围',
);
for (const value of [NaN, Infinity, 0, 51, 1.5])
expect(() => validateTrainingParameters({ obstacle_count: value }, {}, metadata)).toThrow(
'参数 obstacle_count 超出允许范围',
);
for (const value of [1, 50])
expect(() =>
validateTrainingParameters({ obstacle_count: value }, {}, metadata),
).not.toThrow();
});
it('保留默认值比较及高度/感知距离约束', () => {
expect(() => validateTrainingParameters({ obstacle_height_min: 0.7 }, {}, metadata)).toThrow(
'障碍物最小高度不能超过最大高度',
);
expect(() =>
validateTrainingParameters({}, { safetyDistance: 4, maxDistance: 4 }, metadata),
).toThrow('安全距离必须小于探测距离');
});
it('Flat 请求不会添加未选择的地形和传感器字段,保留明确的 rewardPresetId', () => {
expect(createTrainingRequest(form())).toStrictEqual({
taskId: 'Unitree-Go2-Flat',
numEnvs: 4096,
maxIterations: 2000,
seed: 42,
runName: 'web',
device: 'gpu',
gpuIds: [0],
wandbMode: 'offline',
rewardPresetId: undefined,
});
});
it('避障请求不带 Flat 奖励,保留地形及传感器模式并且不修改输入', () => {
const input = form({
taskId: OBSTACLE_TASK_ID,
rewardPresetId: '仅属于Flat',
pretrainedSourceId: '预训练来源',
terrainPreset: 'discrete_obstacles',
terrainParams: { obstacle_count: 12 },
sensorCfg: { safetyDistance: 0.5 },
sensorMode: 'multi_ring_raycast',
});
const before = structuredClone(input);
const request = createTrainingRequest(input);
expect(request).toStrictEqual({
...createTrainingRequest(form()),
taskId: OBSTACLE_TASK_ID,
pretrainedSourceId: '预训练来源',
terrainPreset: 'discrete_obstacles',
terrainParams: { obstacle_count: 12 },
sensorType: 'raycast',
sensorCfg: { safetyDistance: 0.5, sensorMode: 'multi_ring_raycast' },
});
expect(input).toStrictEqual(before);
});
it('移动请求禁用 W&B,省略空接续 ID,非空 ID 原样保留', () => {
const input = form({
mobile: true,
taskId: 'MobileManipulator-LeKiwi-v1',
mobilePackageId: '包',
});
const mobileParams = { ...input.mobileParams };
delete mobileParams.sourceJobId;
expect(createTrainingRequest(input)).toStrictEqual({
...createTrainingRequest(form()),
taskId: input.taskId,
wandbMode: 'disabled',
mobilePackageId: '包',
mobileParams,
});
input.mobileParams.sourceJobId = '已完成作业';
expect(createTrainingRequest(input).mobileParams).toStrictEqual(input.mobileParams);
});
});
+116
View File
@@ -0,0 +1,116 @@
import { OBSTACLE_TASK_ID } from '../rl/deployment';
import { isMobileTrainingTask, type MobileTrainingParams } from '../mobile/training';
import type { TrainingDevice, TrainingRequest, TrainingTaskMetadata } from './types';
export function trainingTaskDefaults(taskId: string) {
const mobile = isMobileTrainingTask(taskId);
return {
numEnvs: mobile ? 1 : 4096,
maxIterations: mobile ? 1000 : 2000,
device: (mobile ? 'cpu' : 'gpu') as TrainingDevice,
terrainPreset: taskId === OBSTACLE_TASK_ID ? 'discrete_obstacles' : '',
};
}
/** 只解析界面输入;不去重、不排序,也不改变 CPU 模式忽略 GPU 文本的行为。 */
export function parseTrainingGpuIds(device: TrainingDevice, gpuIds: string): number[] {
const ids =
device === 'gpu'
? gpuIds
.split(/[\s,]+/)
.filter(Boolean)
.map(Number)
: [];
if (ids.some((id) => !Number.isInteger(id) || id < 0)) throw new Error('GPU 编号必须是非负整数');
return ids;
}
export function validateTrainingParameters(
terrainParams: Record<string, number>,
sensorCfg: Record<string, number>,
metadata: TrainingTaskMetadata | undefined,
): void {
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('安全距离必须小于探测距离');
}
type TrainingForm = Pick<
TrainingRequest,
'taskId' | 'numEnvs' | 'maxIterations' | 'seed' | 'runName' | 'device' | 'gpuIds' | 'wandbMode'
> & {
mobile: boolean;
mobilePackageId?: string;
mobileParams: MobileTrainingParams;
rewardPresetId: string;
pretrainedSourceId: string;
terrainPreset: string;
terrainParams: Record<string, number>;
customTerrainBoxes: TrainingRequest['customTerrainBoxes'];
sensorCfg: Record<string, number>;
sensorMode: 'single_ring_raycast' | 'multi_ring_raycast';
};
/** 场景编译/上传由调用方完成后再构造请求,保留原可选字段与插入顺序。 */
export function createTrainingRequest({
taskId,
numEnvs,
maxIterations,
seed,
runName,
device,
gpuIds,
wandbMode,
mobile,
mobilePackageId,
mobileParams,
rewardPresetId,
pretrainedSourceId,
terrainPreset,
terrainParams,
customTerrainBoxes,
sensorCfg,
sensorMode,
}: TrainingForm): TrainingRequest {
const { sourceJobId, ...mobileValues } = mobileParams;
return {
taskId,
numEnvs,
maxIterations,
seed,
runName,
device,
gpuIds,
wandbMode: mobile ? 'disabled' : wandbMode,
...(mobile
? {
mobilePackageId,
mobileParams: { ...mobileValues, ...(sourceJobId ? { sourceJobId } : {}) },
}
: {}),
rewardPresetId: taskId === 'Unitree-Go2-Flat' ? rewardPresetId || undefined : undefined,
...(pretrainedSourceId ? { pretrainedSourceId } : {}),
...(terrainPreset ? { terrainPreset, terrainParams } : {}),
...(terrainPreset === 'custom_boxes' ? { customTerrainBoxes } : {}),
...(taskId === OBSTACLE_TASK_ID
? { sensorType: 'raycast' as const, sensorCfg: { ...sensorCfg, sensorMode } }
: {}),
};
}
@@ -0,0 +1,37 @@
import type { TrainingJob } from './types';
export const ACTIVE_STATES = new Set(['queued', 'running']);
export function stateLabel(state: TrainingJob['state']): string {
return {
queued: '排队中',
running: '训练中',
succeeded: '已完成',
failed: '失败',
cancelled: '已取消',
}[state];
}
export const TERRAIN_LABELS: Record<string, string> = {
custom_boxes: '自定义场景碰撞布局(AABB近似)',
plane: '平地',
discrete_obstacles: '离散障碍物',
rough: '崎岖地面',
pyramid_stairs: '金字塔台阶',
wave: '波浪地形',
};
export 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: '避障权重',
};