feat(training): release V0.9.1 避障训练与基础策略迁移
This commit is contained in:
@@ -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();
|
||||
});
|
||||
|
||||
@@ -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 }));
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user