feat(training): release V0.9.1 避障训练与基础策略迁移
web-platform-ci / TypeScript, lint, unit, build (push) Has been cancelled
web-platform-ci / Playwright E2E (push) Has been cancelled

This commit is contained in:
2026-09-08 10:50:13 +08:00
parent fa5485049a
commit 438e56bcc8
113 changed files with 15027 additions and 539 deletions
@@ -148,3 +148,66 @@ describe('LocalTrainingClient', () => {
expect(() => new LocalTrainingClient('http://localhost:8765', '')).toThrow('访问令牌');
});
});
it('LocalTrainingClient完整透传自定义任务与嵌套地形/传感器参数', async () => {
const fetchMock = vi.fn().mockResolvedValue(Response.json({ id: 'job' }));
vi.stubGlobal('fetch', fetchMock);
const request = {
taskId: 'Unitree-Go2-ObstacleAvoidance',
numEnvs: 16,
maxIterations: 1,
seed: 42,
runName: 'avoid',
device: 'gpu' as const,
gpuIds: [0],
wandbMode: 'disabled' as const,
terrainPreset: 'discrete_obstacles',
terrainParams: { friction: 0.9, obstacle_count: 12 },
sensorType: 'raycast' as const,
sensorCfg: { fov: 60, maxDistance: 3 },
};
await new LocalTrainingClient('http://127.0.0.1:8765', 'test').start(request);
expect(JSON.parse(String(fetchMock.mock.calls[0][1].body))).toEqual(request);
});
it('上传使用File原始body、显式模板、Bearer及AbortSignal,不发送JSON/base64/服务器路径', async () => {
const fetchMock = vi
.fn()
.mockResolvedValue(new Response(JSON.stringify({ id: 'source' }), { status: 201 }));
vi.stubGlobal('fetch', fetchMock);
const file = new File(['weights'], '单个 actor.ONNX');
const controller = new AbortController();
await new LocalTrainingClient('http://localhost:8765', 'secret-token').uploadPretrained(
file,
'go2-legacy47-v1',
controller.signal,
);
const [url, init] = fetchMock.mock.calls[0] as [string, RequestInit];
const query = new URL(url).searchParams;
expect(query.get('format')).toBe('onnx');
expect(query.get('template')).toBe('go2-legacy47-v1');
expect(query.get('name')).toBe(file.name);
expect(url).not.toContain('secret-token');
expect(init.body).toBe(file);
expect(init.signal).toBe(controller.signal);
expect(new Headers(init.headers).get('Content-Type')).toBe('application/octet-stream');
expect(new Headers(init.headers).get('Authorization')).toBe('Bearer secret-token');
expect(new Headers(init.headers).has('Content-Length')).toBe(false); // Browser owns this forbidden header.
});
it('上传前拒绝ZIP、空文件和超过各格式上限,仍由服务验证实际模型', () => {
const fetchMock = vi.fn();
vi.stubGlobal('fetch', fetchMock);
const client = new LocalTrainingClient('http://localhost:8765', 'token');
for (const [name, size] of [
['model.zip', 1],
['model.pt', 0],
['model.pt', 256 * 1024 ** 2 + 1],
['model.onnx', 64 * 1024 ** 2 + 1],
] as const) {
const file = new File(['x'], name);
Object.defineProperty(file, 'size', { value: size });
expect(() => client.uploadPretrained(file, 'go2-legacy47-v1')).toThrow();
}
expect(fetchMock).not.toHaveBeenCalled();
});
@@ -1,5 +1,6 @@
import type {
ParameterConstraint,
PretrainedSource,
RewardPreset,
TuningCapability,
TuningCreateRequest,
@@ -57,6 +58,26 @@ export class LocalTrainingClient {
health(): Promise<TrainingServerInfo> {
return this.json('/api/training/health');
}
uploadPretrained(
file: File,
template: 'go2-legacy47-v1',
signal?: AbortSignal,
): Promise<PretrainedSource> {
const format = file.name.split('.').pop()?.toLowerCase();
if (format !== 'pt' && format !== 'onnx')
throw new Error('请选择单个.pt或.onnx文件,不支持ZIP/目录');
const limit = (format === 'pt' ? 256 : 64) * 1024 ** 2;
if (!file.size || file.size > limit)
throw new Error(`文件不能为空,${format}上限为${limit / 1024 ** 2}MiB`);
if (template !== 'go2-legacy47-v1') throw new Error('请先确认Go2 legacy47模板');
const query = new URLSearchParams({ format, template, name: file.name });
return this.json(`/api/training/pretrained-sources/upload?${query}`, {
method: 'POST',
headers: { 'Content-Type': 'application/octet-stream' },
body: file,
signal,
});
}
start(request: TrainingRequest): Promise<TrainingJob> {
return this.json('/api/training/jobs', {
method: 'POST',
@@ -1,3 +1,8 @@
import customLayout from '../../../training_server/tests/fixtures/custom-boxes.json';
import { validateCustomTerrain, validatePolicyDeployment } from '../rl/deployment';
import fixture from '../rl/fixtures/obstacleDeployment.json';
import { DEFAULT_PHYSICAL_MAP_CONFIG } from '../map/types';
import { TRAINING_JOB_KEY } from './storage';
import { fireEvent, render, screen, waitFor } from '@testing-library/react';
import { beforeEach, describe, expect, it, vi } from 'vitest';
import { LocalTrainingPanel } from './LocalTrainingPanel';
@@ -161,3 +166,308 @@ describe('LocalTrainingPanel', () => {
);
});
});
const customTasks = ['Unitree-Go2-Flat', 'Unitree-Go2-ObstacleAvoidance', 'Unitree-Go2-Rough'];
const customMetadata = customTasks.map((id) => ({
id,
name: id,
browserCompatible: id !== 'Unitree-Go2-Rough',
terrainPresets: [
'plane',
'discrete_obstacles',
'rough',
'pyramid_stairs',
'wave',
'custom_boxes',
],
terrainParameters: {
size: { min: 8, max: 24, default: 12 },
friction: { min: 0.2, max: 2, default: 0.8 },
obstacle_count: { min: 1, max: 100, default: 24, integer: true },
},
sensorTypes: id.includes('Obstacle') ? ['raycast'] : [],
sensorModes: ['single_ring_raycast', 'multi_ring_raycast'],
sensorParameters: id.includes('Obstacle')
? { fov: { min: 30, max: 120, default: 90 }, maxDistance: { min: 1, max: 5, default: 4 } }
: {},
mapSyncScope: '单块预设参数',
}));
function customServer(job?: unknown) {
const fetchMock = vi.fn(async (url: string, init?: RequestInit) => {
if (url.endsWith('/health'))
return Response.json({
ready: true,
tasks: customTasks,
taskMetadata: customMetadata,
trainerRoot: '/local',
});
if (url.includes('/presets')) return Response.json({ presets: [] });
if (url.endsWith('/policy.onnx')) return new Response(new Uint8Array([8, 9]));
if (init?.method === 'POST' || job)
return Response.json(
job ?? {
id: 'c'.repeat(32),
taskId: 'Unitree-Go2-ObstacleAvoidance',
state: 'queued',
progress: 0,
iteration: 0,
maxIterations: 2,
message: '',
logs: [],
artifactReady: false,
},
);
return Response.json({ error: 'not found' }, { status: 404 });
});
vi.stubGlobal('fetch', fetchMock);
return fetchMock;
}
async function connectCustom() {
fireEvent.change(screen.getByLabelText('训练服务访问令牌'), { target: { value: 'token' } });
fireEvent.click(screen.getByRole('button', { name: '连接' }));
await screen.findByText('/local');
}
describe('LocalTrainingPanel 自定义任务', () => {
it('动态任务/地图/传感器配置传入payload,切换任务清掉sensor与reward状态', async () => {
const fetchMock = customServer();
render(<LocalTrainingPanel onPolicyReady={vi.fn()} />);
await connectCustom();
fireEvent.change(screen.getByLabelText('训练任务'), {
target: { value: 'Unitree-Go2-ObstacleAvoidance' },
});
expect(screen.getByLabelText('训练地形')).toHaveValue('discrete_obstacles');
expect(screen.getByLabelText('奖励配置 preset')).toBeDisabled();
fireEvent.change(screen.getByLabelText('感知角 FOV'), { target: { value: '60' } });
fireEvent.change(screen.getByLabelText('训练地形'), { target: { value: 'rough' } });
expect(screen.getByText(/训练专用 box/)).toBeInTheDocument();
fireEvent.change(screen.getByLabelText('训练任务'), { target: { value: 'Unitree-Go2-Flat' } });
expect(screen.queryByLabelText('感知角 FOV')).not.toBeInTheDocument();
expect(screen.getByLabelText('训练地形')).toHaveValue('');
fireEvent.change(screen.getByLabelText('训练任务'), {
target: { value: 'Unitree-Go2-ObstacleAvoidance' },
});
expect(screen.getByLabelText('感知角 FOV')).toHaveValue(90);
fireEvent.change(screen.getByLabelText('感知角 FOV'), { target: { value: '60' } });
fireEvent.click(screen.getByRole('button', { name: '发起本地训练' }));
await waitFor(() =>
expect(fetchMock.mock.calls.some((call) => call[1]?.method === 'POST')).toBe(true),
);
const payload = JSON.parse(
String(fetchMock.mock.calls.find((call) => call[1]?.method === 'POST')?.[1]?.body),
);
expect(payload).toMatchObject({
taskId: 'Unitree-Go2-ObstacleAvoidance',
terrainPreset: 'discrete_obstacles',
sensorType: 'raycast',
sensorCfg: { fov: 60 },
});
expect(payload.rewardPresetId).toBeUndefined();
});
it('无地图/多实例明确报错而不提交;超限参数阻止请求', async () => {
const fetchMock = customServer();
render(<LocalTrainingPanel onPolicyReady={vi.fn()} />);
await connectCustom();
fireEvent.click(screen.getByRole('button', { name: '同步当前场景地图' }));
expect(screen.getByRole('alert')).toHaveTextContent('没有已应用');
fireEvent.change(screen.getByLabelText('训练地形'), { target: { value: 'plane' } });
fireEvent.change(screen.getByLabelText('地图尺寸 m'), { target: { value: '100' } });
fireEvent.click(screen.getByRole('button', { name: '发起本地训练' }));
await waitFor(() => expect(screen.getByRole('alert')).toHaveTextContent('超出允许范围'));
expect(fetchMock.mock.calls.some((call) => call[1]?.method === 'POST')).toBe(false);
});
it('旧Rough作业禁用导入,不误当Flat', async () => {
customServer({
id: 'r'.repeat(32),
taskId: 'Unitree-Go2-Rough',
state: 'succeeded',
artifactReady: true,
progress: 1,
iteration: 1,
maxIterations: 1,
logs: [],
message: '',
});
localStorage.setItem(TRAINING_JOB_KEY, 'r'.repeat(32));
render(<LocalTrainingPanel onPolicyReady={vi.fn()} />);
await connectCustom();
expect(await screen.findByRole('button', { name: '导入策略' })).toBeDisabled();
expect(screen.getByText(/234维 Rough/)).toBeInTheDocument();
});
});
describe('LocalTrainingPanel 配套策略交接', () => {
it('同步权威布局并上传完整boxes,不再重跑预设种子', async () => {
const fetchMock = customServer();
render(
<LocalTrainingPanel
compileScene={() => validateCustomTerrain(customLayout)}
onPolicyReady={vi.fn()}
sceneMaps={[
{
id: 'one',
name: 'one',
selection: {
kind: 'builtin',
config: {
...DEFAULT_PHYSICAL_MAP_CONFIG,
preset: 'discrete_obstacles',
size: 16,
friction: 1.2,
seed: 123,
obstacleCount: 17,
},
},
},
]}
/>,
);
await connectCustom();
fireEvent.click(screen.getByRole('button', { name: '同步当前场景地图' }));
expect(screen.getByLabelText('训练地形')).toHaveValue('custom_boxes');
expect(screen.getByText('已将视口中 1 个自定义障碍物编译为训练地图布局')).toBeInTheDocument();
expect(screen.getByLabelText('出生 X')).toHaveValue(-2);
expect(screen.getByLabelText('目标 Y')).toHaveValue(-1);
expect(screen.getByText(/旋转障碍会膨胀/)).toBeInTheDocument();
fireEvent.click(screen.getByRole('button', { name: '发起本地训练' }));
await waitFor(() =>
expect(fetchMock.mock.calls.some((call) => call[1]?.method === 'POST')).toBe(true),
);
const request = JSON.parse(
String(fetchMock.mock.calls.find((call) => call[1]?.method === 'POST')?.[1]?.body),
);
expect(request.customTerrainBoxes).toEqual(customLayout);
expect(request.terrainPreset).toBe('custom_boxes');
expect(request.terrainParams).toEqual({ size: 12, friction: 0.8 });
});
it('等待地图/策略异步回调完成,失败显示错误,不提前解除busy', async () => {
localStorage.setItem(TRAINING_JOB_KEY, 'd'.repeat(32));
customServer({
id: 'd'.repeat(32),
taskId: fixture.taskId,
state: 'succeeded',
artifactReady: true,
progress: 1,
iteration: 1,
maxIterations: 1,
logs: ['Mean value loss: 0.0141'],
message: '',
deployment: fixture,
});
let reject!: (error: Error) => void;
const ready = vi.fn(
() =>
new Promise<void>((_resolve, failure) => {
reject = failure;
}),
);
render(<LocalTrainingPanel onPolicyReady={ready} />);
await connectCustom();
const button = await screen.findByRole('button', { name: '导入策略' });
fireEvent.click(button);
await waitFor(() => expect(ready).toHaveBeenCalledOnce());
expect(ready.mock.calls[0]).toEqual([expect.any(File), validatePolicyDeployment(fixture)]);
expect(button).toBeDisabled();
expect(screen.getByText('价值损失')).toBeInTheDocument();
reject(new Error('配套地图失败'));
await waitFor(() => expect(screen.getByRole('alert')).toHaveTextContent('配套地图失败'));
expect(button).toBeEnabled();
});
});
it('custom_boxes拒绝草稿、过时编译场景和未重新同步的坐标;安全区报错不删障碍', async () => {
const fetchMock = customServer();
const assets = [
{
id: 'one',
name: 'one',
selection: {
kind: 'builtin' as const,
config: { ...DEFAULT_PHYSICAL_MAP_CONFIG, preset: 'flat' as const },
},
},
];
let current = validateCustomTerrain(customLayout);
const compileScene = vi.fn(
(coordinates?: import('../map/trainingMap').TrainingSceneCoordinates) => ({
...structuredClone(current),
...(coordinates?.spawn ? { spawn: [...coordinates.spawn, 0.32] } : {}),
...(coordinates?.target ? { target: [...coordinates.target] } : {}),
}),
);
const { rerender } = render(
<LocalTrainingPanel
onPolicyReady={vi.fn()}
sceneMaps={assets}
sceneDirty
compileScene={compileScene}
/>,
);
await connectCustom();
fireEvent.click(screen.getByRole('button', { name: '同步当前场景地图' }));
expect(screen.getByRole('alert')).toHaveTextContent('草稿');
expect(compileScene).not.toHaveBeenCalled();
rerender(
<LocalTrainingPanel onPolicyReady={vi.fn()} sceneMaps={assets} compileScene={compileScene} />,
);
fireEvent.click(screen.getByRole('button', { name: '同步当前场景地图' }));
fireEvent.change(screen.getByLabelText('目标 X'), { target: { value: 1 } });
fireEvent.change(screen.getByLabelText('目标 Y'), { target: { value: 2 } });
fireEvent.click(screen.getByRole('button', { name: '发起本地训练' }));
await waitFor(() => expect(screen.getByRole('alert')).toHaveTextContent('重新同步'));
fireEvent.click(screen.getByRole('button', { name: '同步当前场景地图' }));
expect(screen.getByRole('alert')).toHaveTextContent('安全区');
expect(current.boxes).toHaveLength(2);
fireEvent.change(screen.getByLabelText('目标 X'), { target: { value: 2 } });
fireEvent.change(screen.getByLabelText('目标 Y'), { target: { value: -1 } });
fireEvent.click(screen.getByRole('button', { name: '同步当前场景地图' }));
current = { ...current, boxes: [current.boxes[0], { ...current.boxes[1], pos: [1, 3, 0.5] }] };
fireEvent.click(screen.getByRole('button', { name: '发起本地训练' }));
await waitFor(() => expect(screen.getByRole('alert')).toHaveTextContent('过时'));
expect(fetchMock.mock.calls.some((call) => call[1]?.method === 'POST')).toBe(false);
});
it('显式选择multi48传入训练请求,默认与切换任务仍single32', async () => {
const fetchMock = customServer();
render(<LocalTrainingPanel onPolicyReady={vi.fn()} />);
await connectCustom();
fireEvent.change(screen.getByLabelText('训练任务'), {
target: { value: 'Unitree-Go2-ObstacleAvoidance' },
});
expect(screen.getByLabelText('传感器模式')).toHaveValue('single_ring_raycast');
fireEvent.change(screen.getByLabelText('传感器模式'), {
target: { value: 'multi_ring_raycast' },
});
fireEvent.click(screen.getByRole('button', { name: '发起本地训练' }));
await waitFor(() => expect(fetchMock.mock.calls.some((c) => c[1]?.method === 'POST')).toBe(true));
const payload = JSON.parse(
String(fetchMock.mock.calls.find((c) => c[1]?.method === 'POST')?.[1]?.body),
);
expect(payload.sensorCfg.sensorMode).toBe('multi_ring_raycast');
});
it('Flat奖励菜单只展示服务确认归属Flat的preset,Obstacle和无身份项不展示', async () => {
const entries = [
{ id: 'a'.repeat(32), name: '合法历史Flat', taskId: 'Unitree-Go2-Flat' },
{ id: 'b'.repeat(32), name: 'Obstacle不能混入', taskId: 'Unitree-Go2-ObstacleAvoidance' },
{ id: 'c'.repeat(32), name: '缺少权威身份' },
];
vi.stubGlobal(
'fetch',
vi.fn(async (url: string) => {
if (url.endsWith('/health'))
return Response.json({
ready: true,
tasks: customTasks,
taskMetadata: customMetadata,
trainerRoot: '/local',
});
if (url.endsWith('/presets')) return Response.json({ presets: entries });
return Response.json({ error: 'not found' }, { status: 404 });
}),
);
render(<LocalTrainingPanel onPolicyReady={vi.fn()} />);
await connectCustom();
expect(await screen.findByRole('option', { name: '合法历史Flat' })).toBeInTheDocument();
expect(screen.queryByRole('option', { name: 'Obstacle不能混入' })).not.toBeInTheDocument();
expect(screen.queryByRole('option', { name: '缺少权威身份' })).not.toBeInTheDocument();
});
+451 -23
View File
@@ -1,4 +1,17 @@
import { useEffect, useState, type ReactNode } from 'react';
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';
@@ -33,7 +46,17 @@ function stateLabel(state: TrainingJob['state']): string {
}[state];
}
export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File): void }) {
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),
);
@@ -42,6 +65,9 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
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'),
@@ -53,15 +79,85 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
[gpuIds, setGpuIds] = useState('0'),
[wandbMode, setWandbMode] = useState<WandbMode>('offline');
const [terrainPreset, setTerrainPreset] = useState('');
const [customTerrainBoxes, setCustomTerrainBoxes] = useState<TrainingTerrain>();
const [syncedScene, setSyncedScene] = useState<string>();
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);
setSyncedScene(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,请升级训练服务');
setSyncedScene(undefined);
const layout = compileScene(
customTerrainBoxes
? {
spawn: [customTerrainBoxes.spawn[0], customTerrainBoxes.spawn[1]],
target: [customTerrainBoxes.target[0], customTerrainBoxes.target[1]],
}
: undefined,
);
setCustomTerrainBoxes(layout);
setTerrainPreset('custom_boxes');
setTerrainParams({ size: layout.size, friction: layout.friction });
validateCustomTerrain(layout);
setSyncedScene(JSON.stringify(sceneMaps));
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 {
setPresets(await client.presets());
const nextPresets = await client.presets();
if (epoch !== connectionEpoch.current) return;
setPresets(nextPresets);
} catch {
setPresets([]);
}
@@ -70,11 +166,12 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
} catch {
/* 当前会话仍可连接 */
}
if (info.tasks.length && !info.tasks.includes(taskId)) setTaskId(info.tasks[0]);
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);
@@ -129,10 +226,57 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
typeof event.data !== 'object'
)
return;
const data = event.data as { type?: string; sessionId?: string; policy?: unknown };
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') {
source.postMessage({ type: 'mujoco-tuning-credentials', endpoint, token }, event.origin);
try {
if (taskId === OBSTACLE_TASK_ID && terrainPreset === 'custom_boxes') {
if (
sceneDirty ||
syncedScene !== JSON.stringify(sceneMaps) ||
!compileScene ||
!customTerrainBoxes
)
throw new Error('自定义地图已过时,请重新同步后打开调参');
const current = compileScene({
spawn: [customTerrainBoxes.spawn[0], customTerrainBoxes.spawn[1]],
target: [customTerrainBoxes.target[0], customTerrainBoxes.target[1]],
});
if (JSON.stringify(current) !== JSON.stringify(customTerrainBoxes))
throw new Error('碰撞场景已过时,请重新同步');
}
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,
sensorType: 'raycast',
sensorCfg: { ...sensorCfg, sensorMode },
...(terrainPreset === 'custom_boxes' ? { customTerrainBoxes } : {}),
},
}
: {}),
},
},
event.origin,
);
} catch (value) {
setError(errorText(value));
}
}
if (data.type === 'mujoco-tuning-import-policy' && data.sessionId) {
const reply = (ok: boolean, message?: string) => {
@@ -159,7 +303,12 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
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');
onPolicyReady(policy);
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);
@@ -171,7 +320,23 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
};
window.addEventListener('message', receive);
return () => window.removeEventListener('message', receive);
}, [endpoint, onPolicyReady, token]);
}, [
endpoint,
pretrainedSourceId,
onPolicyReady,
token,
taskId,
seed,
terrainPreset,
terrainParams,
sensorCfg,
sensorMode,
customTerrainBoxes,
syncedScene,
sceneMaps,
sceneDirty,
compileScene,
]);
const openTuningDashboard = () => {
rememberTrainingConnection(endpoint, token);
@@ -179,9 +344,20 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
};
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
@@ -191,6 +367,43 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
: [];
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('安全距离必须小于探测距离');
if (terrainPreset === 'custom_boxes') {
if (
sceneDirty ||
!syncedScene ||
syncedScene !== JSON.stringify(sceneMaps) ||
!compileScene ||
!customTerrainBoxes
)
throw new Error('自定义地图未同步或场景/坐标已更改,请重新同步');
validateCustomTerrain(customTerrainBoxes);
const current = compileScene({
spawn: [customTerrainBoxes.spawn[0], customTerrainBoxes.spawn[1]],
target: [customTerrainBoxes.target[0], customTerrainBoxes.target[1]],
});
if (JSON.stringify(current) !== JSON.stringify(customTerrainBoxes))
throw new Error('已编译碰撞场景已过时,请重新同步');
}
const next = await new LocalTrainingClient(endpoint, token).start({
taskId,
numEnvs,
@@ -200,7 +413,13 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
device,
gpuIds: ids,
wandbMode,
rewardPresetId: rewardPresetId || undefined,
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 } }
: {}),
});
setJob(next);
try {
@@ -231,7 +450,11 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
setBusy(true);
setError(undefined);
try {
onPolicyReady(await new LocalTrainingClient(endpoint, token).downloadPolicy(job.id));
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 {
@@ -249,7 +472,12 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
aria-label="本地训练服务地址"
className="field h-7 min-w-0 flex-1 px-2 text-xs text-text-primary"
value={endpoint}
onChange={(event) => setEndpoint(event.target.value)}
disabled={busy}
onChange={(event) => {
if (busy) return;
connected();
setEndpoint(event.target.value);
}}
/>
<Button
icon={<Link className="h-3.5 w-3.5" />}
@@ -268,7 +496,12 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
autoComplete="off"
className="field h-7 w-full px-2 text-xs text-text-primary"
value={token}
onChange={(event) => setToken(event.target.value)}
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">
@@ -288,21 +521,166 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
打开自调参 Agent 工作台
</Button>
{server?.ready && !job && (
<div className="mt-3 space-y-2">
<fieldset disabled={busy} className="mt-3 space-y-2">
<Field label="训练任务">
<Select
aria-label="训练任务"
className="w-full"
value={taskId}
onChange={(event) => setTaskId(event.target.value)}
onChange={(event) => selectTask(event.target.value)}
>
{server.tasks.map((task) => (
<option key={task} value={task}>
{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);
setSyncedScene(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],出生高度标准化为0.32m。仅保证训练与浏览器使用相同boxes,不等于原OBB。mesh/hfield、地下结构、混合摩擦明确拒绝。
</p>
{terrainPreset === 'custom_boxes' && customTerrainBoxes && (
<>
<div className="grid grid-cols-2 gap-2">
{(['spawn', 'target'] as const).flatMap((key) =>
[0, 1].map((i) => (
<NumberField
key={`${key}${i}`}
label={`${key === 'spawn' ? '出生' : '目标'} ${i === 0 ? 'X' : 'Y'}`}
value={customTerrainBoxes[key][i]}
min={-12}
max={12}
step={0.1}
onChange={(value) => {
setSyncedScene(undefined);
setCustomTerrainBoxes(
(old) =>
old && {
...old,
[key]: old[key].map((v, j) => (i === j ? value : v)),
},
);
}}
/>
)),
)}
</div>
<p>
此处起终点仅用于固定评估和部署初始演示,训练会在同一连通自由区域内逐episode随机采样。参考点须保留0.5m圆形安全区;修改后请重新同步。
</p>
{syncedScene === JSON.stringify(sceneMaps) && !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="并行环境"
@@ -358,17 +736,20 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
</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.map((preset) => (
<option key={preset.id} value={preset.id}>
{preset.name}
</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="实验记录">
@@ -387,7 +768,7 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
variant="primary"
className="w-full"
icon={<Play className="h-3.5 w-3.5" />}
disabled={busy}
disabled={busy || uploading || Boolean(sourceSelectionError)}
onClick={() => void start()}
>
发起本地训练
@@ -396,7 +777,7 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
训练使用本地 mjlab
任务资产,不会把浏览器中的模型上传到网络。服务一次只运行一个训练任务。
</p>
</div>
</fieldset>
)}
{job && (
<div className="mt-3 rounded-lg border border-border bg-surface p-2.5">
@@ -416,11 +797,26 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
{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>
@@ -443,7 +839,12 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
) : (
<>
<Button
disabled={busy || !job.artifactReady}
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()}
>
@@ -492,11 +893,13 @@ function NumberField({
min,
max,
onChange,
step = 1,
}: {
label: string;
value: number;
min: number;
max: number;
step?: number;
onChange(value: number): void;
}) {
return (
@@ -504,6 +907,7 @@ function NumberField({
<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}
@@ -513,3 +917,27 @@ function NumberField({
</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,358 @@
import { fireEvent, render, screen, waitFor } from '@testing-library/react';
import { beforeEach, expect, it, vi } from 'vitest';
import { LocalTrainingPanel } from './LocalTrainingPanel';
import { TuningApp } from '../tuning/TuningApp';
import { useTuningStore } from '../tuning/tuningStore';
import type { PretrainedSource } from './types';
import { PretrainedSourceSelect } from './PretrainedSourceSelect';
const source: PretrainedSource = {
id: 'base',
label: '用户基础行走策略',
ready: true,
compatibleTasks: ['Unitree-Go2-Flat', 'Unitree-Go2-ObstacleAvoidance'],
observationSizes: [47, 81, 97],
initialization: {
sourceId: 'a'.repeat(64),
registeredId: 'base',
label: '用户基础行走策略',
manifest: {
source_iteration: 10000,
source_actor_dim: 47,
normalization: 'preserve-source-count/unit-new-features',
artifacts: Object.fromEntries(
['checkpoint', 'onnx', 'env', 'agent'].map((key) => [
key,
{ name: key === 'checkpoint' ? 'model_10000.pt' : key, sha256: 'b'.repeat(64), bytes: 1 },
]),
) as NonNullable<PretrainedSource['initialization']>['manifest']['artifacts'],
},
},
};
beforeEach(() => {
vi.restoreAllMocks();
vi.unstubAllGlobals();
localStorage.clear();
sessionStorage.clear();
useTuningStore.setState({
sessionId: undefined,
sessions: [],
capability: undefined,
connectionState: 'idle',
error: undefined,
});
});
const json = (value: unknown, status = 200) =>
new Response(JSON.stringify(value), { status, headers: { 'Content-Type': 'application/json' } });
it('普通训练可选来源、显示resolved checkpoint/SHA,任务切换保留选择,服务器错误不退回随机', async () => {
const requests: Record<string, unknown>[] = [];
vi.stubGlobal(
'fetch',
vi.fn((url: string, init?: RequestInit) => {
if (url.endsWith('/health'))
return Promise.resolve(
json({
ready: true,
tasks: ['Unitree-Go2-Flat', 'Unitree-Go2-Rough'],
pretrainedSources: [source],
}),
);
if (url.endsWith('/presets')) return Promise.resolve(json({ presets: [] }));
requests.push(JSON.parse(String(init?.body)) as Record<string, unknown>);
return Promise.resolve(json({ error: '基础策略快照SHA不匹配' }, 400));
}),
);
render(<LocalTrainingPanel onPolicyReady={vi.fn()} />);
fireEvent.change(screen.getByLabelText('训练服务访问令牌'), { target: { value: 'token' } });
fireEvent.click(screen.getByRole('button', { name: /^连接$/ }));
const select = await screen.findByLabelText('基础策略');
fireEvent.change(select, { target: { value: 'base' } });
expect(screen.getByText(/model_10000.pt/)).toBeVisible();
expect(screen.getByText(/checkpoint SHA256/)).toHaveTextContent('b'.repeat(64));
fireEvent.change(screen.getByLabelText('训练任务'), { target: { value: 'Unitree-Go2-Rough' } });
expect(select).toHaveValue('base');
expect(screen.getByRole('button', { name: '发起本地训练' })).toBeDisabled();
expect(screen.getByRole('option', { name: /任务不兼容/ })).toBeDisabled();
fireEvent.change(screen.getByLabelText('训练任务'), { target: { value: 'Unitree-Go2-Flat' } });
fireEvent.change(select, { target: { value: 'base' } });
fireEvent.click(screen.getByRole('button', { name: '发起本地训练' }));
expect(await screen.findByRole('alert')).toHaveTextContent('SHA不匹配');
expect(requests).toHaveLength(1);
expect(requests[0].pretrainedSourceId).toBe('base');
expect(requests[0]).not.toHaveProperty('pretrainedCheckpoint');
expect(select).toHaveValue('base');
});
it('自调参面板选择来源只提交注册ID,切换任务保留来源且保持approval模式', async () => {
const requests: Record<string, unknown>[] = [];
vi.stubGlobal(
'fetch',
vi.fn((url: string, init?: RequestInit) => {
if (url.endsWith('/capabilities'))
return Promise.resolve(
json({ ready: true, configured: true, model: 'stub', pretrainedSources: [source] }),
);
if (init?.method === 'POST') {
requests.push(JSON.parse(String(init.body)) as Record<string, unknown>);
return Promise.resolve(json({ error: '基础策略快照文件失效' }, 400));
}
return Promise.resolve(json({ sessions: [] }));
}),
);
render(<TuningApp />);
fireEvent.change(screen.getByLabelText('访问令牌(仅当前标签页)'), {
target: { value: 'token' },
});
fireEvent.click(screen.getByRole('button', { name: '连接/刷新' }));
await screen.findByRole('option', { name: /用户基础行走策略/ });
const select = screen.getByLabelText('基础策略');
fireEvent.change(select, { target: { value: 'base' } });
fireEvent.change(screen.getByLabelText('调参任务'), {
target: { value: 'Unitree-Go2-ObstacleAvoidance' },
});
expect(select).toHaveValue('base');
fireEvent.change(select, { target: { value: 'base' } });
fireEvent.click(screen.getByRole('button', { name: '启动自调参' }));
await waitFor(() => expect(requests).toHaveLength(1));
expect(requests[0]).toMatchObject({
pretrainedSourceId: 'base',
mode: 'approval',
taskId: 'Unitree-Go2-ObstacleAvoidance',
});
expect(await screen.findByText('基础策略快照文件失效')).toBeVisible();
});
it('缺checkpoint注册条目显式错误,不作为可用来源', () => {
render(
<PretrainedSourceSelect
taskId="Unitree-Go2-Flat"
value=""
onChange={vi.fn()}
sources={[
{
id: 'bad',
label: '错误来源',
ready: false,
compatibleTasks: [],
error: '缺少.pt,请配置匹配checkpoint',
},
]}
/>,
);
expect(screen.getByRole('alert')).toHaveTextContent('缺少.pt');
expect(screen.getByRole('option', { name: /验证失败/ })).toBeDisabled();
});
const sourceA: PretrainedSource = { ...source, id: 'a'.repeat(64) };
const sourceB: PretrainedSource = {
...source,
id: 'c'.repeat(64),
// Same registration alias/label is deliberately retained; content B is not A.
initialization: { ...source.initialization!, sourceId: 'c'.repeat(64) },
};
const invalidCatalogs: [string, PretrainedSource[] | undefined][] = [
['同别名A替换为B', [sourceB]],
['A变为not-ready', [{ ...sourceA, ready: false }, sourceB]],
['目录缺项', undefined],
['A变为任务不兼容', [{ ...sourceA, compatibleTasks: ['Unitree-Go2-Rough'] }, sourceB]],
];
for (const panel of ['普通训练', '自调参'] as const) {
for (const [scenario, invalidCatalog] of invalidCatalogs) {
for (const resolution of ['明确从头训练', '明确选择B'] as const) {
it(`${panel}真实刷新:${scenario}保留A并阻止请求,${resolution}后恢复`, async () => {
let catalog: PretrainedSource[] | undefined = [sourceA];
const requests: Record<string, unknown>[] = [];
vi.stubGlobal(
'fetch',
vi.fn((url: string, init?: RequestInit) => {
if (url.endsWith('/health'))
return Promise.resolve(
json({
ready: true,
tasks: ['Unitree-Go2-Flat'],
pretrainedSources: catalog,
}),
);
if (url.endsWith('/capabilities'))
return Promise.resolve(
json({
ready: true,
configured: true,
model: 'stub',
pretrainedSources: catalog,
}),
);
if (url.endsWith('/presets')) return Promise.resolve(json({ presets: [] }));
if (init?.method === 'POST') {
requests.push(JSON.parse(String(init.body)) as Record<string, unknown>);
return Promise.resolve(json({ error: '测试截获请求,未启动训练' }, 400));
}
return Promise.resolve(json({ sessions: [] }));
}),
);
const ordinary = panel === '普通训练';
render(ordinary ? <LocalTrainingPanel onPolicyReady={vi.fn()} /> : <TuningApp />);
fireEvent.change(
screen.getByLabelText(ordinary ? '训练服务访问令牌' : '访问令牌(仅当前标签页)'),
{ target: { value: 'token' } },
);
const refresh = screen.getByRole('button', { name: ordinary ? /^连接$/ : '连接/刷新' });
fireEvent.click(refresh);
await screen.findByRole('option', { name: /用户基础行走策略/ });
const select = screen.getByLabelText('基础策略');
fireEvent.change(select, { target: { value: sourceA.id } });
const start = screen.getByRole('button', {
name: ordinary ? '发起本地训练' : '启动自调参',
});
await waitFor(() => expect(start).toBeEnabled());
catalog = invalidCatalog;
fireEvent.click(refresh);
await screen.findByText(/所选基础策略已失效:/);
await waitFor(() => expect(refresh).toBeEnabled());
expect(select).toHaveValue(sourceA.id);
expect(start).toBeDisabled();
fireEvent.click(start);
expect(requests).toHaveLength(0);
if (resolution === '明确选择B') {
catalog = [sourceB];
fireEvent.click(refresh);
await screen.findByRole('option', { name: /用户基础行走策略.*已验证/ });
await waitFor(() => expect(refresh).toBeEnabled());
expect(select).toHaveValue(sourceA.id);
expect(start).toBeDisabled();
expect(requests).toHaveLength(0);
fireEvent.change(select, { target: { value: sourceB.id } });
} else fireEvent.change(select, { target: { value: '' } });
expect(screen.queryByText(/所选基础策略已失效:/)).not.toBeInTheDocument();
await waitFor(() => expect(start).toBeEnabled());
fireEvent.click(start);
await waitFor(() => expect(requests).toHaveLength(1));
if (resolution === '明确选择B') expect(requests[0].pretrainedSourceId).toBe(sourceB.id);
else expect(requests[0]).not.toHaveProperty('pretrainedSourceId');
if (!ordinary) expect(requests[0].mode).toBe('approval');
});
}
}
}
const uploadCapability = {
enabled: true,
templateId: 'go2-legacy47-v1',
formats: { pt: 256 * 1024 ** 2, onnx: 64 * 1024 ** 2 },
endpoint: '/api/training/pretrained-sources/upload',
};
for (const ordinary of [true, false]) {
for (const outcome of ['pt', 'onnx', 'failure', 'cancel', 'task', 'connection'] as const) {
it(`${ordinary ? '普通' : '自调参'}上传${outcome}:阻止未完成启动并隔离旧epoch,保留旧选择`, async () => {
let finish!: (response: Response) => void;
const uploads: { url: string; init?: RequestInit }[] = [];
const starts: Record<string, unknown>[] = [];
const format = outcome === 'onnx' ? 'onnx' : 'pt';
const uploaded: PretrainedSource = {
...sourceB,
label: `single.${format}`,
initialization: {
...sourceB.initialization!,
label: `single.${format}`,
manifest: {
...sourceB.initialization!.manifest,
sourceFormat: format,
source_iteration: format === 'onnx' ? null : 10000,
contract: 'go2-legacy47-v1',
derived_fields: { normalizer_count: { policy: 'synthetic', value: 1000000 } },
artifacts: {
checkpoint: source.initialization!.manifest.artifacts.checkpoint,
upload: { name: `upload.${format}`, sha256: 'd'.repeat(64), bytes: 3 },
},
},
},
};
vi.stubGlobal(
'fetch',
vi.fn((url: string, init?: RequestInit) => {
if (url.includes('/pretrained-sources/upload?')) {
uploads.push({ url, init });
return new Promise<Response>((resolve) => {
finish = resolve;
});
}
if (url.endsWith('/health') || url.endsWith('/capabilities'))
return Promise.resolve(
json({
ready: true,
configured: true,
tasks: ['Unitree-Go2-Flat', 'Unitree-Go2-Rough'],
pretrainedSources: [sourceA],
pretrainedUpload: uploadCapability,
}),
);
if (url.endsWith('/presets')) return Promise.resolve(json({ presets: [] }));
if (init?.method === 'POST') {
starts.push(JSON.parse(String(init.body)) as Record<string, unknown>);
return Promise.resolve(json({ error: '测试禁止真实训练' }, 400));
}
return Promise.resolve(json({ sessions: [] }));
}),
);
render(ordinary ? <LocalTrainingPanel onPolicyReady={vi.fn()} /> : <TuningApp />);
fireEvent.change(
screen.getByLabelText(ordinary ? '训练服务访问令牌' : '访问令牌(仅当前标签页)'),
{ target: { value: 'token' } },
);
const connect = () =>
fireEvent.click(screen.getByRole('button', { name: ordinary ? /^连接$/ : '连接/刷新' }));
connect();
await screen.findByRole('option', { name: /用户基础行走策略/ });
fireEvent.change(screen.getByLabelText('基础策略'), { target: { value: sourceA.id } });
const input = screen.getByLabelText('选择基础策略文件');
expect(input).toBeDisabled();
fireEvent.click(screen.getByLabelText('确认Go2 legacy47模板'));
const file = new File(['abc'], `single.${format}`);
fireEvent.change(input, { target: { files: [file] } });
await waitFor(() => expect(uploads).toHaveLength(1));
expect(uploads[0].init?.body).toBe(file);
expect(uploads[0].url).toContain('template=go2-legacy47-v1');
const startName = ordinary ? '发起本地训练' : '启动自调参';
expect(screen.getByRole('button', { name: startName })).toBeDisabled();
expect(starts).toHaveLength(0);
if (outcome === 'cancel') fireEvent.click(screen.getByRole('button', { name: '取消上传' }));
if (outcome === 'task')
fireEvent.change(screen.getByLabelText(ordinary ? '训练任务' : '调参任务'), {
target: { value: ordinary ? 'Unitree-Go2-Rough' : 'Unitree-Go2-ObstacleAvoidance' },
});
if (outcome === 'connection') {
fireEvent.change(screen.getByLabelText(ordinary ? '本地训练服务地址' : '训练服务地址'), {
target: { value: 'http://127.0.0.1:9999' },
});
connect();
await screen.findByRole('option', { name: /用户基础行走策略/ });
}
finish(outcome === 'failure' ? json({ error: '不支持该模型' }, 400) : json(uploaded, 201));
if (outcome === 'pt' || outcome === 'onnx') {
await waitFor(() => expect(screen.getByLabelText('基础策略')).toHaveValue(sourceB.id));
expect(screen.getByText(/原文件 SHA256/)).toHaveTextContent('d'.repeat(64));
if (outcome === 'onnx')
expect(screen.getByText(/ONNX统计count合成/)).toHaveTextContent('1000000');
fireEvent.click(screen.getByRole('button', { name: startName }));
await waitFor(() => expect(starts).toHaveLength(1));
expect(starts[0].pretrainedSourceId).toBe(sourceB.id);
if (!ordinary) expect(starts[0].mode).toBe('approval');
} else {
if (outcome === 'failure') {
await screen.findByText(/不支持该模型/);
fireEvent.change(input, { target: { files: [file] } });
await waitFor(() => expect(uploads).toHaveLength(2));
fireEvent.click(screen.getByRole('button', { name: '取消上传' }));
finish(json(uploaded, 201));
}
await waitFor(() => expect(screen.getByLabelText('基础策略')).toHaveValue(sourceA.id));
expect(uploads[0].init?.signal?.aborted).toBe(outcome !== 'failure');
expect(starts).toHaveLength(0);
}
});
}
}
@@ -0,0 +1,123 @@
import { Select } from '../components/ui';
import type { PretrainedInitialization, PretrainedSource } from './types';
import { pretrainedSelectionError } from './pretrainedSelection';
import { PretrainedUpload, type PretrainedUploadConnection } from './PretrainedUpload';
export function PretrainedIdentity({ source }: { source?: PretrainedInitialization }) {
if (!source)
return <p className="text-xs text-text-tertiary">初始化:随机新策略(未选择基础策略)</p>;
return (
<div className="break-all text-xs text-text-secondary" aria-label="基础策略身份">
<p>
初始化:{source.label} · {source.manifest.artifacts.checkpoint.name}
</p>
<p>checkpoint SHA256:{source.manifest.artifacts.checkpoint.sha256}</p>
{source.manifest.artifacts.upload ? (
<>
<p>
上传格式:{source.manifest.sourceFormat} · 原文件:{source.label}
</p>
<p>原文件 SHA256:{source.manifest.artifacts.upload.sha256}</p>
<p>
模板:{source.manifest.contract}
(用户确认缺失的物理语义);继承actor权重,非完整resume。
</p>
{source.manifest.sourceFormat === 'onnx' && (
<p>
ONNX统计count合成:{source.manifest.derived_fields?.normalizer_count?.value ?? '未知'}
; 探索std使用新训练默认,critic/optimizer重新初始化。
</p>
)}
</>
) : (
<p>ONNX SHA256:{source.manifest.artifacts.onnx?.sha256 ?? '未知'}</p>
)}
<p>
来源迭代 {source.manifest.source_iteration ?? '未知'}
;新训练从0开始,critic/optimizer重新初始化;同trial续训保留自身checkpoint。
</p>
</div>
);
}
export function PretrainedSourceSelect({
sources = [],
taskId,
value,
onChange,
disabled = false,
upload,
}: {
sources?: PretrainedSource[];
taskId: string;
value: string;
onChange(value: string): void;
disabled?: boolean;
upload?: PretrainedUploadConnection;
}) {
const selected = sources.find((source) => source.id === value);
const selectionError = pretrainedSelectionError(sources, taskId, value);
return (
<div className="space-y-1">
<label className="block text-xs text-text-secondary">
基础策略
<Select
aria-label="基础策略"
className="w-full"
value={value}
disabled={disabled}
onChange={(event) => onChange(event.target.value)}
>
<option value="">不选择(随机初始化)</option>
{value && !selected && (
<option value={value} disabled>
所选基础策略已失效({value})
</option>
)}
{sources.map((source) => (
<option
key={source.id}
value={source.id}
disabled={!source.ready || !source.compatibleTasks.includes(taskId)}
>
{source.label}
{!source.ready
? '(验证失败)'
: !source.compatibleTasks.includes(taskId)
? '(任务不兼容)'
: '(已验证)'}
</option>
))}
</Select>
</label>
{upload && (
<PretrainedUpload
key={JSON.stringify([upload.endpoint, upload.token, upload.revision, taskId])}
connection={upload}
disabled={disabled}
/>
)}
{!sources.length && (
<p className="text-xs">尚无基础策略,可直接上传单个文件;不选择则随机初始化。</p>
)}
{sources
.filter((source) => source.error)
.map((source) => (
<p role="alert" key={source.id}>
{source.label}:{source.error}
</p>
))}
{selected && (
<p className="text-xs">
兼容:{selected.compatibleTasks.join(' / ')};观测 {selected.observationSizes?.join('/')}{' '}
→ 12动作。只继承actor权重,新训练从0开始。
</p>
)}
{selectionError ? (
<p role="alert">{selectionError}</p>
) : (
<PretrainedIdentity source={selected?.initialization} />
)}
</div>
);
}
@@ -0,0 +1,111 @@
import { useEffect, useRef, useState } from 'react';
import { Button } from '../components/ui';
import { LocalTrainingClient } from './LocalTrainingClient';
import type { PretrainedSource } from './types';
export interface PretrainedUploadConnection {
endpoint: string;
token: string;
revision: number;
enabled: boolean;
onUploaded(source: PretrainedSource): void;
onBusyChange(busy: boolean): void;
}
/** Remounted for each connection/task epoch by the shared source selector. */
export function PretrainedUpload({
connection,
disabled,
}: {
connection: PretrainedUploadConnection;
disabled: boolean;
}) {
const [confirmed, setConfirmed] = useState(false);
const [pending, setPending] = useState(false);
const [message, setMessage] = useState('');
const [error, setError] = useState('');
const request = useRef<AbortController | null>(null);
const onBusyChange = connection.onBusyChange;
useEffect(
() => () => {
request.current?.abort();
onBusyChange(false);
},
[onBusyChange],
);
const upload = async (file: File) => {
if (!confirmed || disabled || pending || !connection.enabled) return;
const controller = new AbortController();
request.current = controller;
setPending(true);
onBusyChange(true);
setError('');
setMessage(`正在上传并验证 ${file.name},请稍候…`);
try {
const source = await new LocalTrainingClient(
connection.endpoint,
connection.token,
).uploadPretrained(file, 'go2-legacy47-v1', controller.signal);
if (controller.signal.aborted) return;
if (!source.ready || !source.initialization)
throw new Error(source.error ?? '上传来源未通过验证');
connection.onUploaded(source);
setMessage(`已选择 ${source.label};仅继承策略权重,不是完整训练resume。`);
} catch (value) {
if (!controller.signal.aborted) {
setMessage('');
setError(
`${value instanceof Error ? value.message : String(value)};保留原基础策略选择,未启动训练。`,
);
}
} finally {
if (!controller.signal.aborted) {
request.current = null;
setPending(false);
onBusyChange(false);
}
}
};
const cancel = () => {
request.current?.abort();
request.current = null;
setPending(false);
onBusyChange(false);
setMessage('已取消等待,保留原基础策略选择;服务若已完成验证,可刷新目录查看。');
setError('');
};
return (
<div className="space-y-2 rounded border border-border p-2 text-xs">
<p>上传基础策略:单个.pt(≤256MiB)或.onnx(≤64MiB),无需服务器路径或配套文件。</p>
<label className="flex gap-2">
<input
type="checkbox"
aria-label="确认Go2 legacy47模板"
checked={confirmed}
disabled={disabled || pending || !connection.enabled}
onChange={(event) => setConfirmed(event.target.checked)}
/>
按Go2 legacy47观测及FL/FR/RL/RR关节顺序解释文件;缺失的物理语义由我确认。
</label>
<p>
仅支持47维Go2行走actor;ONNX仅继承推理网络,统计count合成,探索/价值网络/优化器重新初始化。
</p>
<input
type="file"
accept=".pt,.onnx"
aria-label="选择基础策略文件"
disabled={!confirmed || disabled || pending || !connection.enabled}
onChange={(event) => {
const file = event.target.files?.[0];
event.target.value = ''; // Same-file retry must still dispatch change.
if (file) void upload(file);
}}
/>
{!connection.enabled && <p>请先连接支持单文件上传的训练服务。</p>}
{pending && <Button onClick={cancel}>取消上传</Button>}
{message && <p role="status">{message}</p>}
{error && <p role="alert">{error}</p>}
</div>
);
}
@@ -0,0 +1,63 @@
import { TrainingMetricHistory } from './TrainingMetricHistory';
const header = (i: number) => `Learning iteration ${i} / 1000`;
it('解析多个迭代全部指标、重复poll去重、同迭代分批补全', () => {
const h = new TrainingMetricHistory();
const logs = [
header(1),
'Mean value loss: 1e-3',
'Mean surrogate loss: -0.1',
header(2),
'Mean value loss: 2',
];
expect(h.update('a', logs)).toBe(true);
expect(h.update('a', logs)).toBe(false);
expect(h.series()[0].points.map((p) => [p.step, p.value])).toEqual([
[1, 0.001],
[2, 2],
]);
expect(
h.update('a', [...logs, 'Mean entropy loss: -3', 'Mean reward: 4', 'Mean episode length: 50']),
).toBe(true);
expect(h.series()).toHaveLength(5);
expect(h.series().find((s) => s.tag === '平均奖励')!.points[0]).toEqual({
step: 2,
value: 4,
wallTime: 0,
});
});
it('滚动截断按重叠上下文补全,无header/overlap不猜迭代,忽略非有限/缺失/其它数字', () => {
const h = new TrainingMetricHistory();
h.update('a', [header(3), 'Mean value loss: 3']);
h.update('a', ['Mean value loss: 3', 'Mean reward: 7']);
expect(h.series()[1].points[0].step).toBe(3);
h.update('a', ['Mean entropy loss: 5']);
expect(h.series()).toHaveLength(2);
h.update('a', [
header(4),
'Mean value loss: NaN',
'Mean reward: Infinity',
'Mean reward: 1e999',
'Mean reward: 3 ms',
'Total timesteps: 40',
'Iteration time: 2',
]);
expect(h.series()[0].points).toHaveLength(1);
h.update('b', ['Mean reward: 1']);
expect(h.series()).toEqual([]);
h.update('b', [header(0), 'Mean reward: .5']);
expect(h.series()[0].points[0].step).toBe(0);
});
it('容量和去重索引有界,旧重发不挤掉最新迭代,job切换清空', () => {
const h = new TrainingMetricHistory();
for (let i = 0; i < 650; i++) h.update('a', [header(i), `Mean reward: ${i}`]);
const points = h.series()[0].points;
expect(points).toHaveLength(500);
expect(points[0].step).toBe(150);
expect(points.at(-1)!.step).toBe(649);
h.update('a', [header(1), 'Mean reward: 999']);
expect(h.series()[0].points).toEqual(points);
h.update('new-job', [header(0), 'Mean value loss: 1']);
expect(h.series()[0].points).toHaveLength(1);
expect(() => new TrainingMetricHistory(1)).toThrow();
});
@@ -0,0 +1,93 @@
import type { ScalarSeries } from './types';
export const TRAINING_METRICS = {
value: '价值损失',
surrogate: '策略损失',
entropy: '熵损失',
reward: '平均奖励',
episodeLength: '平均回合长度',
} as const;
type Metric = keyof typeof TRAINING_METRICS;
const METRIC_PATTERN =
/^\s*Mean (value loss|surrogate loss|entropy loss|reward|episode length):\s*([-+]?(?:\d+(?:\.\d*)?|\.\d+)(?:e[-+]?\d+)?)\s*$/i;
const KEYS: Record<string, Metric> = {
'value loss': 'value',
'surrogate loss': 'surrogate',
'entropy loss': 'entropy',
reward: 'reward',
'episode length': 'episodeLength',
};
/** Bounded iteration merge. Headerless lines are accepted only with proven suffix/prefix overlap.
* On reconnect without overlap, orphan scalars are skipped rather than guessed onto a new step. */
export class TrainingMetricHistory {
private readonly rows = new Map<number, Partial<Record<Metric, number>>>();
private previous: { line: string; iteration?: number }[] = [];
private jobId?: string;
constructor(private readonly capacity = 500) {
if (!Number.isInteger(capacity) || capacity < 2) throw new Error('指标容量必须至少为2');
}
update(jobId: string, logs: readonly string[]): boolean {
let changed = false;
if (this.jobId !== jobId) {
this.rows.clear();
this.previous = [];
this.jobId = jobId;
changed = true;
}
const lines = logs
.flatMap((line) => line.split('\n'))
.slice(-1000)
// rsl_rl uses ANSI SGR colors around iteration headers.
// eslint-disable-next-line no-control-regex
.map((line) => line.replace(/\u001b\[[0-9;]*m/g, '').trim());
let iteration: number | undefined;
for (let overlap = Math.min(lines.length, this.previous.length); overlap > 0; overlap--) {
const start = this.previous.length - overlap;
if (lines.slice(0, overlap).every((line, i) => this.previous[start + i].line === line)) {
iteration = this.previous[start].iteration;
break;
}
}
const contexts: typeof this.previous = [];
for (const line of lines) {
const header = /\bLearning iteration\s+(\d+)\s*\/\s*\d+/i.exec(line);
if (header) {
const step = Number(header[1]);
iteration = Number.isSafeInteger(step) ? step : undefined;
}
contexts.push({ line, iteration });
const scalar = METRIC_PATTERN.exec(line);
if (iteration === undefined || !scalar) continue;
const value = Number(scalar[2]);
if (!Number.isFinite(value)) continue;
const key = KEYS[scalar[1].toLowerCase()];
if (!this.rows.has(iteration) && this.rows.size >= this.capacity) {
const oldest = Math.min(...this.rows.keys());
if (iteration < oldest) continue;
this.rows.delete(oldest);
}
const row = this.rows.get(iteration) ?? {};
if (row[key] !== value) {
row[key] = value;
changed = true;
}
this.rows.set(iteration, row);
}
this.previous = contexts;
return changed;
}
series(): ScalarSeries[] {
const rows = [...this.rows].sort(([a], [b]) => a - b);
return Object.entries(TRAINING_METRICS)
.map(([key, tag]) => ({
tag,
points: rows.flatMap(([step, row]) =>
row[key as Metric] === undefined
? []
: [{ step, wallTime: 0, value: row[key as Metric]! }],
),
}))
.filter((series) => series.points.length > 0);
}
}
@@ -0,0 +1,36 @@
import { fireEvent, render, screen } from '@testing-library/react';
import { TrainingMetricsPanel } from './TrainingMetricsPanel';
import { ScalarChart } from '../components/charts/ScalarChart';
vi.mock('../components/charts/ScalarChart', () => ({
ScalarChart: vi.fn(({ title }: { title: string }) => (
<div data-testid="metric-chart">{title}曲线</div>
)),
}));
it('折叠不挂载图表,价值/策略/综合独立缩放,job切换清空历史', () => {
const logs = [
'Learning iteration 1 / 10',
'Mean value loss: 1',
'Mean surrogate loss: -2',
'Mean entropy loss: -3',
'Mean reward: 4',
'Mean episode length: 50',
];
const view = render(<TrainingMetricsPanel jobId="a" logs={logs} />);
expect(screen.queryByTestId('metric-chart')).not.toBeInTheDocument();
fireEvent.click(screen.getByRole('button', { name: /训练指标趋势/ }));
expect(screen.getByText('价值损失曲线')).toBeInTheDocument();
vi.mocked(ScalarChart).mockClear();
view.rerender(<TrainingMetricsPanel jobId="a" logs={[...logs]} />);
expect(ScalarChart).not.toHaveBeenCalled();
fireEvent.click(screen.getByRole('tab', { name: '策略损失' }));
expect(screen.getByText('策略损失曲线')).toBeInTheDocument();
fireEvent.click(screen.getByRole('tab', { name: '综合' }));
expect(screen.getAllByTestId('metric-chart')).toHaveLength(5);
expect(screen.getByText('平均回合长度曲线')).toBeInTheDocument();
view.rerender(<TrainingMetricsPanel jobId="b" logs={[]} />);
expect(screen.queryByTestId('metric-chart')).not.toBeInTheDocument();
expect(screen.getByText(/尚无带迭代编号/)).toBeInTheDocument();
fireEvent.click(screen.getByRole('button', { name: /训练指标趋势/ }));
expect(screen.queryByRole('tabpanel')).not.toBeInTheDocument();
});
@@ -0,0 +1,81 @@
import { memo, useEffect, useState } from 'react';
import { ScalarChart } from '../components/charts/ScalarChart';
import { TrainingMetricHistory, TRAINING_METRICS } from './TrainingMetricHistory';
import type { ScalarSeries } from './types';
const MetricChart = memo(function MetricChart({ series }: { series: ScalarSeries }) {
return (
<div>
<p className="mb-1 text-xs text-text-secondary">
{series.tag} · 最新原值 {series.points.at(-1)!.value.toPrecision(4)}
</p>
<ScalarChart series={[series]} smoothing={0.4} title={series.tag} xLabel="Iteration" />
</div>
);
});
export const TrainingMetricsPanel = memo(function TrainingMetricsPanel({
jobId,
logs,
}: {
jobId: string;
logs: readonly string[];
}) {
const [history] = useState(() => new TrainingMetricHistory());
const [series, setSeries] = useState<ScalarSeries[]>([]);
const [open, setOpen] = useState(false);
const [tab, setTab] = useState('value');
useEffect(() => {
if (history.update(jobId, logs)) setSeries(history.series());
}, [history, jobId, logs]);
const selected = series.filter(
(item) =>
tab === 'all' ||
item.tag === (tab === 'value' ? TRAINING_METRICS.value : TRAINING_METRICS.surrogate),
);
return (
<section className="mt-3 min-w-0 rounded-lg border border-border">
<button
type="button"
className="w-full p-2 text-left text-xs text-text-primary"
aria-expanded={open}
onClick={() => setOpen(!open)}
>
{open ? '▾' : '▸'} 训练指标趋势
</button>
{open && (
<div className="min-w-0 space-y-2 p-2 pt-0">
<div role="tablist" aria-label="训练指标" className="flex gap-2">
{[
['value', '价值损失'],
['surrogate', '策略损失'],
['all', '综合'],
].map(([id, label]) => (
<button
key={id}
type="button"
role="tab"
aria-selected={tab === id}
className={`rounded px-2 py-1 text-xs ${tab === id ? 'bg-accent/10 text-accent' : 'text-text-secondary'}`}
onClick={() => setTab(id)}
>
{label}
</button>
))}
</div>
<div role="tabpanel" aria-label="训练指标曲线" className="space-y-2">
{selected.length ? (
selected.map((item) => <MetricChart key={item.tag} series={item} />)
) : (
<p className="text-xs text-text-tertiary">尚无带迭代编号的指标日志</p>
)}
</div>
<p className="text-[10px] text-text-tertiary">
各指标独立纵轴;EMA 0.4
仅用于曲线,悬停显示原值。最多保留最近500个有指标的迭代,重连仅恢复服务端日志尾部。
</p>
</div>
)}
</section>
);
});
@@ -0,0 +1,21 @@
import type { PretrainedSource } from './types';
/** A nonempty selection is an initialization intent, never a random-training fallback. */
export function pretrainedSelectionError(
sources: readonly PretrainedSource[] | undefined,
taskId: string,
selectedId: string,
): string | undefined {
if (!selectedId) return undefined;
const source = sources?.find((item) => item.id === selectedId);
const reason = !source
? '目录缺少该内容ID'
: !source.ready
? '来源未通过验证'
: !source.compatibleTasks.includes(taskId)
? '来源与当前任务不兼容'
: undefined;
return reason
? `所选基础策略已失效:${reason}。请明确选择其他有效基础策略,或选择“不选择(随机初始化)”;不会自动退回随机初始化。`
: undefined;
}
@@ -0,0 +1,18 @@
import { expect, it } from 'vitest';
import { trainingLosses } from './trainingLosses';
it('提取最新有限PPO损失,兼容无指标旧日志', () => {
expect(
trainingLosses([
'Mean value loss: 1.2',
'Mean surrogate loss: -0.03',
'Mean entropy loss: 1e-3',
'Mean value loss: 0.1',
'Mean value loss: Infinity',
]),
).toEqual([
{ label: '价值损失', value: 0.1 },
{ label: '策略损失', value: -0.03 },
{ label: '熵损失', value: 0.001 },
]);
expect(trainingLosses(['Starting training'])).toEqual([]);
});
@@ -0,0 +1,14 @@
/** rsl_rl console scalars, newest finite value per loss; old servers need no new endpoint. */
export function trainingLosses(logs: readonly string[]): { label: string; value: number }[] {
const values = new Map<string, number>();
for (const line of logs) {
const match = /Mean (value|surrogate|entropy) loss:\s*([-+\d.eE]+)/.exec(line);
if (match && Number.isFinite(Number(match[2]))) values.set(match[1], Number(match[2]));
}
const labels: Record<string, string> = {
value: '价值损失',
surrogate: '策略损失',
entropy: '熵损失',
};
return Array.from(values, ([key, value]) => ({ label: labels[key], value }));
}
+85 -10
View File
@@ -1,13 +1,75 @@
import type { PolicyDeployment, TrainingTerrain } from '../rl/deployment';
export interface TrainingParameter {
min: number;
max: number;
default: number;
integer?: boolean;
}
export interface TrainingTaskMetadata {
id: string;
name: string;
browserCompatible: boolean;
terrainPresets: string[];
terrainParameters: Record<string, TrainingParameter>;
sensorTypes: string[];
sensorModes?: string[];
sensorParameters: Record<string, TrainingParameter>;
mapSyncScope: string;
}
export type TrainingJobState = 'queued' | 'running' | 'succeeded' | 'failed' | 'cancelled';
export type TrainingDevice = 'cpu' | 'gpu';
export type WandbMode = 'offline' | 'online' | 'disabled';
export interface PretrainedInitialization {
sourceId: string;
registeredId: string;
label: string;
manifest: {
source_iteration: number | null;
sourceFormat?: 'pt' | 'onnx';
contract?: string;
derived_fields?: {
normalizer_count?: { policy: string; value: number };
exploration_std?: { policy: string; value: number };
};
source_actor_dim: number;
normalization: string;
artifacts: { checkpoint: PretrainedArtifact } & Partial<
Record<'upload' | 'onnx' | 'env' | 'agent', PretrainedArtifact>
>;
};
}
export interface PretrainedArtifact {
name: string;
sha256: string;
bytes: number;
}
export interface PretrainedUploadCapability {
enabled: boolean;
templateId: 'go2-legacy47-v1';
formats: { pt: number; onnx: number };
endpoint: string;
}
export interface PretrainedSource {
id: string;
label: string;
ready: boolean;
compatibleTasks: string[];
observationSizes?: number[];
initialization?: PretrainedInitialization;
error?: string;
}
export interface TrainingServerInfo {
pretrainedSources?: PretrainedSource[];
pretrainedUpload?: PretrainedUploadCapability;
version: string;
ready: boolean;
trainerRoot: string;
python: string;
tasks: string[];
taskMetadata?: TrainingTaskMetadata[];
activeJobId?: string;
resourceOwner?: string;
tuning?: TuningCapability;
@@ -15,6 +77,11 @@ export interface TrainingServerInfo {
}
export interface TrainingRequest {
customTerrainBoxes?: TrainingTerrain;
terrainPreset?: string;
terrainParams?: Record<string, number>;
sensorType?: 'raycast';
sensorCfg?: Partial<import('../rl/deployment').ObstacleSensorConfig>;
taskId: string;
numEnvs: number;
maxIterations: number;
@@ -24,9 +91,12 @@ export interface TrainingRequest {
gpuIds: number[];
wandbMode: WandbMode;
rewardPresetId?: string;
pretrainedSourceId?: string;
}
export interface TrainingJob {
pretrained?: PretrainedInitialization;
deployment?: PolicyDeployment;
id: string;
state: TrainingJobState;
taskId: string;
@@ -60,16 +130,11 @@ export interface RewardConfiguration {
params: Record<string, number>;
}
export interface ObjectiveWeights {
velocity_tracking: number;
action_smoothness: number;
posture_stability: number;
fall_avoidance: number;
foot_slip: number;
energy: number;
}
export type ObjectiveWeights = Record<string, number>;
export interface TuningCapability {
pretrainedSources?: PretrainedSource[];
pretrainedUpload?: PretrainedUploadCapability;
ready: boolean;
configured: boolean;
apiKeyConfigured: boolean;
@@ -79,7 +144,12 @@ export interface TuningCapability {
}
export interface TuningCreateRequest {
taskId: 'Unitree-Go2-Flat';
pretrainedSourceId?: string;
taskId: 'Unitree-Go2-Flat' | 'Unitree-Go2-ObstacleAvoidance';
taskConfig?: Pick<
TrainingRequest,
'terrainPreset' | 'terrainParams' | 'sensorType' | 'sensorCfg' | 'customTerrainBoxes'
> & { seed?: number };
mode: TuningMode;
runName: string;
numEnvs: number;
@@ -160,7 +230,11 @@ export interface TuningSession {
mode: TuningMode;
createdAt: string;
updatedAt: string;
config: TuningCreateRequest & { rungs: number[]; promote: number[] };
config: TuningCreateRequest & {
rungs: number[];
promote: number[];
pretrained?: PretrainedInitialization;
};
objectiveWeights: ObjectiveWeights;
message: string;
currentTrialId?: string;
@@ -194,6 +268,7 @@ export interface TuningMetricsResponse {
export interface RewardPreset {
id: string;
taskId: string;
name: string;
sessionId: string;
trialId: string;