refactor(web-platform): release V0.7.2 项目常规优化
This commit is contained in:
@@ -0,0 +1,87 @@
|
||||
import { fireEvent, render, screen, waitFor } from '@testing-library/react';
|
||||
import { beforeEach, describe, expect, it, vi } from 'vitest';
|
||||
import { LocalTrainingPanel } from './LocalTrainingPanel';
|
||||
|
||||
beforeEach(() => {
|
||||
localStorage.clear();
|
||||
sessionStorage.clear();
|
||||
vi.unstubAllGlobals();
|
||||
});
|
||||
|
||||
describe('LocalTrainingPanel', () => {
|
||||
it('连接本地服务并从图形界面发起训练请求', async () => {
|
||||
const health = {
|
||||
version: '0.1.0',
|
||||
ready: true,
|
||||
trainerRoot: '/opt/unitree_rl_mjlab',
|
||||
python: '/env/bin/python',
|
||||
tasks: ['Unitree-Go2-Flat'],
|
||||
};
|
||||
const job = {
|
||||
id: 'a'.repeat(32),
|
||||
state: 'queued',
|
||||
taskId: 'Unitree-Go2-Flat',
|
||||
createdAt: '2025-01-01T00:00:00Z',
|
||||
iteration: 0,
|
||||
maxIterations: 2000,
|
||||
progress: 0,
|
||||
message: '等待启动',
|
||||
logs: [],
|
||||
artifactReady: false,
|
||||
};
|
||||
const fetchMock = vi
|
||||
.fn()
|
||||
.mockResolvedValueOnce(
|
||||
new Response(JSON.stringify(health), {
|
||||
status: 200,
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
}),
|
||||
)
|
||||
.mockResolvedValueOnce(
|
||||
new Response(JSON.stringify(job), {
|
||||
status: 202,
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
}),
|
||||
)
|
||||
.mockResolvedValueOnce(
|
||||
new Response(JSON.stringify(health), {
|
||||
status: 200,
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
}),
|
||||
)
|
||||
.mockResolvedValueOnce(
|
||||
new Response(JSON.stringify({ error: '训练任务不存在或服务已重启' }), {
|
||||
status: 404,
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
}),
|
||||
);
|
||||
vi.stubGlobal('fetch', fetchMock);
|
||||
render(<LocalTrainingPanel onPolicyReady={vi.fn()} />);
|
||||
fireEvent.change(screen.getByLabelText('训练服务访问令牌'), {
|
||||
target: { value: 'secret-token' },
|
||||
});
|
||||
fireEvent.click(screen.getByRole('button', { name: '连接' }));
|
||||
expect(await screen.findByText('/opt/unitree_rl_mjlab')).toBeInTheDocument();
|
||||
fireEvent.change(screen.getByLabelText('并行环境'), { target: { value: '32' } });
|
||||
fireEvent.click(screen.getByRole('button', { name: '发起本地训练' }));
|
||||
await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(2));
|
||||
const request = fetchMock.mock.calls[1][1] as RequestInit;
|
||||
expect(JSON.parse(String(request.body))).toMatchObject({
|
||||
taskId: 'Unitree-Go2-Flat',
|
||||
numEnvs: 32,
|
||||
device: 'gpu',
|
||||
gpuIds: [0],
|
||||
wandbMode: 'offline',
|
||||
});
|
||||
expect(new Headers(request.headers).get('Authorization')).toBe('Bearer secret-token');
|
||||
expect(await screen.findByText('排队中')).toBeInTheDocument();
|
||||
|
||||
const tokenInput = screen.getByLabelText('训练服务访问令牌');
|
||||
expect(tokenInput).toBeEnabled();
|
||||
fireEvent.change(tokenInput, { target: { value: 'new-secret-token' } });
|
||||
fireEvent.click(screen.getByRole('button', { name: '连接' }));
|
||||
await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(4));
|
||||
expect(await screen.findByRole('button', { name: '发起本地训练' })).toBeInTheDocument();
|
||||
expect(sessionStorage.getItem('mujoco-local-training-token')).toBe('new-secret-token');
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,430 @@
|
||||
import { useEffect, useState, type ReactNode } from 'react';
|
||||
import { Download, Link, Play, Server, Square } from 'lucide-react';
|
||||
import { Badge, Button, ProgressBar, PropertyRow, Select } from '../components/ui';
|
||||
import { LocalTrainingClient } from './LocalTrainingClient';
|
||||
import type { TrainingDevice, TrainingJob, TrainingServerInfo, WandbMode } from './types';
|
||||
|
||||
const ENDPOINT_KEY = 'mujoco-local-training-endpoint',
|
||||
JOB_KEY = 'mujoco-local-training-job',
|
||||
TOKEN_KEY = 'mujoco-local-training-token';
|
||||
const DEFAULT_ENDPOINT = 'http://127.0.0.1:8765';
|
||||
const ACTIVE_STATES = new Set(['queued', 'running']);
|
||||
function stored(key: string, fallback = ''): string {
|
||||
try {
|
||||
return localStorage.getItem(key) ?? fallback;
|
||||
} catch {
|
||||
return fallback;
|
||||
}
|
||||
}
|
||||
function sessionStored(key: string): string {
|
||||
try {
|
||||
return sessionStorage.getItem(key) ?? '';
|
||||
} catch {
|
||||
return '';
|
||||
}
|
||||
}
|
||||
function errorText(error: unknown): string {
|
||||
return error instanceof Error ? error.message : String(error);
|
||||
}
|
||||
function stateLabel(state: TrainingJob['state']): string {
|
||||
return {
|
||||
queued: '排队中',
|
||||
running: '训练中',
|
||||
succeeded: '已完成',
|
||||
failed: '失败',
|
||||
cancelled: '已取消',
|
||||
}[state];
|
||||
}
|
||||
|
||||
export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File): void }) {
|
||||
const [endpoint, setEndpoint] = useState(() => stored(ENDPOINT_KEY, DEFAULT_ENDPOINT));
|
||||
const [token, setToken] = useState(() => sessionStored(TOKEN_KEY));
|
||||
const [server, setServer] = useState<TrainingServerInfo>();
|
||||
const [job, setJob] = useState<TrainingJob>();
|
||||
const [busy, setBusy] = useState(false),
|
||||
[error, setError] = useState<string>();
|
||||
const [taskId, setTaskId] = useState('Unitree-Go2-Flat'),
|
||||
[numEnvs, setNumEnvs] = useState(4096),
|
||||
[maxIterations, setMaxIterations] = useState(2000),
|
||||
[seed, setSeed] = useState(42),
|
||||
[runName, setRunName] = useState('web'),
|
||||
[device, setDevice] = useState<TrainingDevice>('gpu'),
|
||||
[gpuIds, setGpuIds] = useState('0'),
|
||||
[wandbMode, setWandbMode] = useState<WandbMode>('offline');
|
||||
|
||||
const connect = async () => {
|
||||
setBusy(true);
|
||||
setError(undefined);
|
||||
try {
|
||||
const client = new LocalTrainingClient(endpoint, token),
|
||||
info = await client.health();
|
||||
setServer(info);
|
||||
try {
|
||||
localStorage.setItem(ENDPOINT_KEY, client.endpoint);
|
||||
sessionStorage.setItem(TOKEN_KEY, client.token);
|
||||
} catch {
|
||||
/* 当前会话仍可连接 */
|
||||
}
|
||||
if (info.tasks.length && !info.tasks.includes(taskId)) setTaskId(info.tasks[0]);
|
||||
const remembered = info.activeJobId ?? stored(JOB_KEY);
|
||||
if (remembered) {
|
||||
try {
|
||||
const recovered = await client.job(remembered);
|
||||
setJob(recovered);
|
||||
try {
|
||||
localStorage.setItem(JOB_KEY, recovered.id);
|
||||
} catch {
|
||||
/* ignore */
|
||||
}
|
||||
} catch {
|
||||
setJob(undefined);
|
||||
try {
|
||||
localStorage.removeItem(JOB_KEY);
|
||||
} catch {
|
||||
/* ignore */
|
||||
}
|
||||
}
|
||||
} else {
|
||||
setJob(undefined);
|
||||
}
|
||||
if (!info.ready) setError(info.error ?? '训练服务尚未就绪');
|
||||
} catch (value) {
|
||||
setServer(undefined);
|
||||
setError(errorText(value));
|
||||
} finally {
|
||||
setBusy(false);
|
||||
}
|
||||
};
|
||||
|
||||
const jobId = job?.id,
|
||||
jobState = job?.state;
|
||||
useEffect(() => {
|
||||
if (!jobId || !jobState || !ACTIVE_STATES.has(jobState)) return;
|
||||
let disposed = false;
|
||||
const refresh = async () => {
|
||||
try {
|
||||
const next = await new LocalTrainingClient(endpoint, token).job(jobId);
|
||||
if (!disposed) setJob(next);
|
||||
} catch (value) {
|
||||
if (!disposed) setError(errorText(value));
|
||||
}
|
||||
};
|
||||
const timer = window.setInterval(() => void refresh(), 1500);
|
||||
return () => {
|
||||
disposed = true;
|
||||
window.clearInterval(timer);
|
||||
};
|
||||
}, [endpoint, jobId, jobState, token]);
|
||||
|
||||
const start = async () => {
|
||||
setBusy(true);
|
||||
setError(undefined);
|
||||
try {
|
||||
const ids =
|
||||
device === 'gpu'
|
||||
? gpuIds
|
||||
.split(/[\s,]+/)
|
||||
.filter(Boolean)
|
||||
.map(Number)
|
||||
: [];
|
||||
if (ids.some((id) => !Number.isInteger(id) || id < 0))
|
||||
throw new Error('GPU 编号必须是非负整数');
|
||||
const next = await new LocalTrainingClient(endpoint, token).start({
|
||||
taskId,
|
||||
numEnvs,
|
||||
maxIterations,
|
||||
seed,
|
||||
runName,
|
||||
device,
|
||||
gpuIds: ids,
|
||||
wandbMode,
|
||||
});
|
||||
setJob(next);
|
||||
try {
|
||||
localStorage.setItem(JOB_KEY, next.id);
|
||||
} catch {
|
||||
/* ignore */
|
||||
}
|
||||
} catch (value) {
|
||||
setError(errorText(value));
|
||||
} finally {
|
||||
setBusy(false);
|
||||
}
|
||||
};
|
||||
const cancel = async () => {
|
||||
if (!job) return;
|
||||
setBusy(true);
|
||||
setError(undefined);
|
||||
try {
|
||||
setJob(await new LocalTrainingClient(endpoint, token).cancel(job.id));
|
||||
} catch (value) {
|
||||
setError(errorText(value));
|
||||
} finally {
|
||||
setBusy(false);
|
||||
}
|
||||
};
|
||||
const importResult = async () => {
|
||||
if (!job) return;
|
||||
setBusy(true);
|
||||
setError(undefined);
|
||||
try {
|
||||
onPolicyReady(await new LocalTrainingClient(endpoint, token).downloadPolicy(job.id));
|
||||
} catch (value) {
|
||||
setError(errorText(value));
|
||||
} finally {
|
||||
setBusy(false);
|
||||
}
|
||||
};
|
||||
const active = Boolean(job && ACTIVE_STATES.has(job.state));
|
||||
|
||||
return (
|
||||
<div>
|
||||
<label className="block text-xs text-text-secondary">
|
||||
<span className="mb-1 block">本地训练服务</span>
|
||||
<div className="flex gap-2">
|
||||
<input
|
||||
aria-label="本地训练服务地址"
|
||||
className="field h-7 min-w-0 flex-1 px-2 text-xs text-text-primary"
|
||||
value={endpoint}
|
||||
onChange={(event) => setEndpoint(event.target.value)}
|
||||
/>
|
||||
<Button
|
||||
icon={<Link className="h-3.5 w-3.5" />}
|
||||
disabled={busy || !token.trim()}
|
||||
onClick={() => void connect()}
|
||||
>
|
||||
连接
|
||||
</Button>
|
||||
</div>
|
||||
</label>
|
||||
<label className="mt-2 block text-[10px] text-text-tertiary">
|
||||
<span className="mb-1 block">访问令牌(服务启动时显示)</span>
|
||||
<input
|
||||
aria-label="训练服务访问令牌"
|
||||
type="password"
|
||||
autoComplete="off"
|
||||
className="field h-7 w-full px-2 text-xs text-text-primary"
|
||||
value={token}
|
||||
onChange={(event) => setToken(event.target.value)}
|
||||
/>
|
||||
</label>
|
||||
<div className="mt-2 flex items-center justify-between rounded-md border border-border bg-surface px-2 py-1.5 text-[10px] text-text-tertiary">
|
||||
<span className="flex min-w-0 items-center gap-1.5 truncate">
|
||||
<Server className="h-3.5 w-3.5" />
|
||||
{server?.trainerRoot ?? '请先启动本地训练服务'}
|
||||
</span>
|
||||
<Badge tone={server?.ready ? 'success' : 'warning'}>
|
||||
{server?.ready ? '可用' : '离线'}
|
||||
</Badge>
|
||||
</div>
|
||||
{server?.ready && !job && (
|
||||
<div className="mt-3 space-y-2">
|
||||
<Field label="训练任务">
|
||||
<Select
|
||||
aria-label="训练任务"
|
||||
className="w-full"
|
||||
value={taskId}
|
||||
onChange={(event) => setTaskId(event.target.value)}
|
||||
>
|
||||
{server.tasks.map((task) => (
|
||||
<option key={task} value={task}>
|
||||
{task}
|
||||
</option>
|
||||
))}
|
||||
</Select>
|
||||
</Field>
|
||||
<div className="grid grid-cols-2 gap-2">
|
||||
<NumberField
|
||||
label="并行环境"
|
||||
value={numEnvs}
|
||||
min={1}
|
||||
max={16384}
|
||||
onChange={setNumEnvs}
|
||||
/>
|
||||
<NumberField
|
||||
label="训练迭代"
|
||||
value={maxIterations}
|
||||
min={1}
|
||||
max={1000000}
|
||||
onChange={setMaxIterations}
|
||||
/>
|
||||
<NumberField
|
||||
label="随机种子"
|
||||
value={seed}
|
||||
min={0}
|
||||
max={2147483647}
|
||||
onChange={setSeed}
|
||||
/>
|
||||
<Field label="运行名称">
|
||||
<input
|
||||
aria-label="运行名称"
|
||||
className="field h-7 w-full px-2 text-xs text-text-primary"
|
||||
value={runName}
|
||||
onChange={(event) => setRunName(event.target.value)}
|
||||
/>
|
||||
</Field>
|
||||
</div>
|
||||
<div className="grid grid-cols-2 gap-2">
|
||||
<Field label="计算设备">
|
||||
<Select
|
||||
aria-label="计算设备"
|
||||
className="w-full"
|
||||
value={device}
|
||||
onChange={(event) => setDevice(event.target.value as TrainingDevice)}
|
||||
>
|
||||
<option value="gpu">GPU</option>
|
||||
<option value="cpu">CPU</option>
|
||||
</Select>
|
||||
</Field>
|
||||
<Field label="GPU 编号">
|
||||
<input
|
||||
aria-label="GPU 编号"
|
||||
className="field h-7 w-full px-2 text-xs text-text-primary disabled:opacity-40"
|
||||
value={gpuIds}
|
||||
disabled={device === 'cpu'}
|
||||
onChange={(event) => setGpuIds(event.target.value)}
|
||||
/>
|
||||
</Field>
|
||||
</div>
|
||||
<Field label="实验记录">
|
||||
<Select
|
||||
aria-label="W&B 模式"
|
||||
className="w-full"
|
||||
value={wandbMode}
|
||||
onChange={(event) => setWandbMode(event.target.value as WandbMode)}
|
||||
>
|
||||
<option value="offline">本地离线(默认,无需登录)</option>
|
||||
<option value="disabled">完全禁用 W&B</option>
|
||||
<option value="online">在线 W&B(需要 API Key)</option>
|
||||
</Select>
|
||||
</Field>
|
||||
<Button
|
||||
variant="primary"
|
||||
className="w-full"
|
||||
icon={<Play className="h-3.5 w-3.5" />}
|
||||
disabled={busy}
|
||||
onClick={() => void start()}
|
||||
>
|
||||
发起本地训练
|
||||
</Button>
|
||||
<p className="text-[10px] leading-4 text-text-tertiary">
|
||||
训练使用本地 mjlab
|
||||
任务资产,不会把浏览器中的模型上传到网络。服务一次只运行一个训练任务。
|
||||
</p>
|
||||
</div>
|
||||
)}
|
||||
{job && (
|
||||
<div className="mt-3 rounded-lg border border-border bg-surface p-2.5">
|
||||
<div className="mb-2 flex items-center justify-between gap-2">
|
||||
<span className="truncate text-xs font-medium text-text-primary" title={job.id}>
|
||||
{job.taskId}
|
||||
</span>
|
||||
<Badge
|
||||
tone={
|
||||
job.state === 'succeeded'
|
||||
? 'success'
|
||||
: job.state === 'failed' || job.state === 'cancelled'
|
||||
? 'warning'
|
||||
: 'accent'
|
||||
}
|
||||
>
|
||||
{stateLabel(job.state)}
|
||||
</Badge>
|
||||
</div>
|
||||
<ProgressBar value={job.progress} label="训练进度" />
|
||||
<div className="mt-2">
|
||||
<PropertyRow label="迭代" value={`${job.iteration} / ${job.maxIterations}`} />
|
||||
<PropertyRow label="状态" value={job.message} />
|
||||
</div>
|
||||
{job.logs.length > 0 && (
|
||||
<details className="mt-2">
|
||||
<summary className="cursor-pointer text-[10px] text-text-secondary">最近日志</summary>
|
||||
<pre className="mt-1 max-h-36 overflow-auto whitespace-pre-wrap break-all rounded bg-app p-2 text-[9px] leading-4 text-text-tertiary">
|
||||
{job.logs.slice(-40).join('\n')}
|
||||
</pre>
|
||||
</details>
|
||||
)}
|
||||
<div className="mt-3 grid grid-cols-2 gap-2">
|
||||
{active ? (
|
||||
<Button
|
||||
variant="danger"
|
||||
className="col-span-2"
|
||||
icon={<Square className="h-3.5 w-3.5" />}
|
||||
disabled={busy}
|
||||
onClick={() => void cancel()}
|
||||
>
|
||||
停止训练
|
||||
</Button>
|
||||
) : (
|
||||
<>
|
||||
<Button
|
||||
disabled={busy || !job.artifactReady}
|
||||
icon={<Download className="h-3.5 w-3.5" />}
|
||||
onClick={() => void importResult()}
|
||||
>
|
||||
导入策略
|
||||
</Button>
|
||||
<Button
|
||||
onClick={() => {
|
||||
setJob(undefined);
|
||||
try {
|
||||
localStorage.removeItem(JOB_KEY);
|
||||
} catch {
|
||||
/* ignore */
|
||||
}
|
||||
}}
|
||||
>
|
||||
新建任务
|
||||
</Button>
|
||||
</>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
{error && (
|
||||
<p
|
||||
role="alert"
|
||||
className="mt-2 break-words rounded bg-danger/10 p-2 text-[10px] leading-4 text-danger"
|
||||
>
|
||||
{error}
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
function Field({ label, children }: { label: string; children: ReactNode }) {
|
||||
return (
|
||||
<label className="block text-[10px] text-text-tertiary">
|
||||
<span className="mb-1 block">{label}</span>
|
||||
{children}
|
||||
</label>
|
||||
);
|
||||
}
|
||||
function NumberField({
|
||||
label,
|
||||
value,
|
||||
min,
|
||||
max,
|
||||
onChange,
|
||||
}: {
|
||||
label: string;
|
||||
value: number;
|
||||
min: number;
|
||||
max: number;
|
||||
onChange(value: number): void;
|
||||
}) {
|
||||
return (
|
||||
<Field label={label}>
|
||||
<input
|
||||
aria-label={label}
|
||||
type="number"
|
||||
className="field h-7 w-full px-2 text-xs text-text-primary"
|
||||
value={value}
|
||||
min={min}
|
||||
max={max}
|
||||
onChange={(event) => onChange(Number(event.target.value))}
|
||||
/>
|
||||
</Field>
|
||||
);
|
||||
}
|
||||
Reference in New Issue
Block a user