feat(training): release V0.8 自调参 Agent
web-platform-ci / TypeScript, lint, unit, build (push) Has been cancelled
web-platform-ci / Playwright E2E (push) Has been cancelled

This commit is contained in:
2026-09-02 13:49:34 +08:00
parent cffac29a03
commit deead17a9a
47 changed files with 4986 additions and 96 deletions
+4 -1
View File
@@ -14,7 +14,8 @@ src/
├── rl/ ONNX 策略运行时、任务绑定、类型和面板
├── simulation/ MuJoCo 会话、物理适配器和仿真控制组件
├── telemetry/ 数据源抽象、记录器、导出和数据面板
├── training/ 本地训练客户端、类型和面板
├── training/ 本地训练/调参客户端、类型、共享连接和面板
├── tuning/ 独立 tuning.html 的 Agent dashboard 与 scalar 图表
├── viewer/ Three.js 场景、渲染和交互
├── stores/ 跨域应用状态
└── test/ 全局测试初始化
@@ -38,4 +39,6 @@ src/
SimulationSession snapshot → app → viewer / 各业务面板
```
主工作台由 `index.html → src/main.tsx` 启动;自调参工作台由 Vite MPA 入口 `tuning.html → src/tuning/main.tsx` 启动,避免把 MuJoCo/Three.js 主应用依赖打入监控页面。两页仅通过训练 HTTP API和严格同源的短消息交接训练服务凭据/策略导入请求,不在 URL 中传 token。
测试文件使用 `*.test.ts(x)` 与被测模块共置;端到端测试统一保存在 `e2e/`
+4 -2
View File
@@ -15,7 +15,7 @@
- 内置平地、坡道、楼梯、可复现随机障碍物及 9 类系统参数化地形,可配置尺寸、摩擦、难度、种子与高度场采样精度
- 导入单文件 `.py` 控制器,通过本地 Pyodide 在 `mj_step` 前按仿真时间同步执行
- 导入 mjlab 导出的 `policy.onnx`,在浏览器本地执行 Go2-W 平衡/速度策略推理
- 从图形界面向本机训练桥接服务发起 mjlab 强化学习训练、查看进度/日志、停止任务并导入训练生成的 ONNX
- 从图形界面向本机训练桥接服务发起 mjlab 强化学习训练、查看进度/日志、停止任务并导入训练生成的 ONNX;可在独立 TensorBoard 风格页面运行 DeepSeek 奖励函数自调参
- 可配置仿真遥测记录,实时查看速度、机身姿态、位置、驱动力等指标并导出 CSV/JSON
- FPS、物理耗时和主线程步进预算提示
@@ -119,7 +119,9 @@ npm run training-server -- \
--trainer-python /path/to/training-env/bin/python
```
服务启动时会在终端输出一个随机访问令牌;在界面中填写该令牌后连接。令牌仅保存在当前标签页的 `sessionStorage`。界面默认连接 `http://127.0.0.1:8765`,可选择服务端允许的任务、并行环境数、训练迭代、随机种子、CPU/GPU、GPU 编号和实验记录方式。W&B 默认为本地离线模式,无需登录或 API Key;也可完全禁用,只有明确选择在线模式时才会联网登录。训练期间页面轮询迭代进度与最近日志,可以停止任务;训练成功后点击“导入策略”,生成的 `policy.onnx` 会进入现有 ONNX 加载流程。
服务启动时会在终端输出一个随机访问令牌;在界面中填写该令牌后连接。令牌仅保存在当前标签页的 `sessionStorage`。界面默认连接 `http://127.0.0.1:8765`,可选择服务端允许的任务、并行环境数、训练迭代、随机种子、CPU/GPU、GPU 编号和实验记录方式。W&B 默认为本地离线模式,无需登录或 API Key;也可完全禁用,只有明确选择在线模式时才会联网登录。训练期间页面轮询迭代进度与最近日志,可以停止任务;训练成功后点击“导入策略”,生成的 `policy.onnx` 会进入现有 ONNX 加载流程。普通训练还可以选择自调参产生的命名 reward preset,而不会改写仓库默认配置。
连接服务后点击“打开自调参 Agent 工作台”会打开独立 `tuning.html`。该页面提供自动/逐轮审批模式、目标权重与预算配置、TensorBoard scalar 筛选/平滑/缩放、trial/rung 排行、固定评估指标、Agent 决策时间线、参数 patch 修改审批、暂停/恢复/停止、最佳 preset JSON 和 ONNX 导出。新标签页 URL 不包含 token;同源 opener 会一次性交接凭据,直接打开页面时也可手工输入。DeepSeek key 始终由本地 Python 服务的 `DEEPSEEK_API_KEY` 环境变量读取,浏览器不会接触该 key。
桥接服务只监听本机回环地址,并检查 Host、Origin 和 Bearer Token;仅接受允许列表中的任务和经过范围校验的参数,不执行前端提供的 Shell 命令;一次只运行一个训练进程。默认任务使用仓库内置的 Go2 机器人资产与环境配置,**不会自动把浏览器中临时编辑的 MJCF/URDF 作为训练环境**。自定义浏览器模型训练需要在兼容的外部训练工程中注册 task,并通过服务的 `--trainer-root` 指定该工程。服务配置、接口和安全边界见 [`../training_server/README.md`](../training_server/README.md)。
+8
View File
@@ -77,6 +77,14 @@ const LARGE_MODEL = `
</worldbody>
</mujoco>`;
test('独立自调参工作台不需要加载 MuJoCo 主应用即可打开', async ({ page }) => {
await page.goto('/tuning.html');
await expect(page.getByRole('heading', { name: 'Go2 奖励函数自调参 Agent' })).toBeVisible();
await expect(page.getByText('新建 Unitree-Go2-Flat 调参 Session')).toBeVisible();
await expect(page.getByLabel('访问令牌(仅当前标签页)')).toBeVisible();
await expect(page.getByRole('button', { name: '启动自调参' })).toBeVisible();
});
test('显示中文平台骨架并加载单文件模型', async ({ page }) => {
page.on('console', (message) => console.log(`[browser:${message.type()}] ${message.text()}`));
page.on('pageerror', (error) => console.log(`[browser:error] ${error.message}`));
+21
View File
@@ -1 +1,22 @@
import '@testing-library/jest-dom/vitest';
Object.defineProperty(window, 'matchMedia', {
writable: true,
value: (query: string) => ({
matches: false,
media: query,
onchange: null,
addListener: () => undefined,
removeListener: () => undefined,
addEventListener: () => undefined,
removeEventListener: () => undefined,
dispatchEvent: () => false,
}),
});
if (!globalThis.ResizeObserver)
globalThis.ResizeObserver = class ResizeObserver {
observe(): void {}
unobserve(): void {}
disconnect(): void {}
};
@@ -51,6 +51,81 @@ describe('LocalTrainingClient', () => {
).rejects.toThrow('已有训练任务正在运行');
});
it('调用调参 session、metrics 与审批接口且不泄露令牌到 URL', async () => {
const fetchMock = vi.fn().mockImplementation(() =>
Promise.resolve(
new Response(JSON.stringify({ id: 'b'.repeat(32), state: 'running', series: [] }), {
status: 200,
headers: { 'Content-Type': 'application/json' },
}),
),
);
vi.stubGlobal('fetch', fetchMock);
const client = new LocalTrainingClient('http://127.0.0.1:8765', 'deep-secret');
await client.tuningMetrics('a'.repeat(32), 'b'.repeat(32), ['Evaluation/score'], 500);
await client.decideProposal('a'.repeat(32), 'c'.repeat(32), 'approve', {
feedback: 'ok',
patch: { weights: { pose: 1.1 }, params: {} },
});
expect(fetchMock.mock.calls[0][0]).toContain('/api/tuning/sessions/');
expect(fetchMock.mock.calls[0][0]).toContain('maxPoints=500');
expect(fetchMock.mock.calls[0][0]).not.toContain('deep-secret');
const approval = fetchMock.mock.calls[1][1] as RequestInit;
expect(approval.method).toBe('POST');
expect(new Headers(approval.headers).get('Authorization')).toBe('Bearer deep-secret');
});
it('覆盖调参生命周期、preset 与最佳策略下载客户端方法', async () => {
const fetchMock = vi.fn().mockImplementation((input: string) => {
if (String(input).endsWith('/policy.onnx'))
return Promise.resolve(new Response(new Uint8Array([1, 2, 3]), { status: 200 }));
return Promise.resolve(
new Response(
JSON.stringify({ sessions: [], presets: [], configured: true, id: 'a'.repeat(32) }),
{
status: 200,
headers: { 'Content-Type': 'application/json' },
},
),
);
});
vi.stubGlobal('fetch', fetchMock);
const client = new LocalTrainingClient('http://127.0.0.1:8765', 'secret-token');
await client.tuningCapability();
await client.testTuningAgent();
await client.tuningSessions();
await client.startTuning({
taskId: 'Unitree-Go2-Flat',
mode: 'automatic',
runName: 'test',
numEnvs: 16,
seed: 42,
gpuIds: [0],
trialCount: 4,
initialIterations: 1,
middleIterations: 2,
finalIterations: 3,
evalNumEnvs: 8,
evalSteps: 10,
objectiveWeights: {
velocity_tracking: 0.35,
action_smoothness: 0.2,
posture_stability: 0.15,
fall_avoidance: 0.15,
foot_slip: 0.1,
energy: 0.05,
},
fallbackEnabled: false,
});
await client.tuningSession('a'.repeat(32));
await client.tuningAction('a'.repeat(32), 'pause');
await client.cancelTuning('a'.repeat(32));
await client.presets();
const policy = await client.downloadBestPolicy('a'.repeat(32));
expect(policy.size).toBe(3);
expect(fetchMock).toHaveBeenCalledTimes(9);
});
it('拒绝非 HTTP 地址和空访问令牌', () => {
expect(() => new LocalTrainingClient('file:///tmp/socket', 'secret-token')).toThrow(
'http 或 https',
@@ -1,4 +1,13 @@
import type { TrainingJob, TrainingRequest, TrainingServerInfo } from './types';
import type {
RewardPreset,
TuningCapability,
TuningCreateRequest,
TuningMetricsResponse,
TuningSession,
TrainingJob,
TrainingRequest,
TrainingServerInfo,
} from './types';
function normalizeEndpoint(value: string): string {
const endpoint = value.trim().replace(/\/+$/, '');
@@ -61,12 +70,85 @@ export class LocalTrainingClient {
return this.json(`/api/training/jobs/${encodeURIComponent(id)}`, { method: 'DELETE' });
}
async downloadPolicy(id: string): Promise<File> {
const response = await fetch(
`${this.endpoint}/api/training/jobs/${encodeURIComponent(id)}/artifacts/policy.onnx`,
this.requestInit(),
return this.download(
`/api/training/jobs/${encodeURIComponent(id)}/artifacts/policy.onnx`,
`policy-${id.slice(0, 8)}.onnx`,
);
}
tuningCapability(): Promise<TuningCapability> {
return this.json('/api/tuning/capabilities');
}
testTuningAgent(): Promise<{ ok: boolean; model: string; outputType: string }> {
return this.json('/api/tuning/agent/test', { method: 'POST' });
}
tuningSessions(): Promise<TuningSession[]> {
return this.json<{ sessions: TuningSession[] }>('/api/tuning/sessions').then(
(value) => value.sessions,
);
}
startTuning(request: TuningCreateRequest): Promise<TuningSession> {
return this.json('/api/tuning/sessions', {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify(request),
});
}
tuningSession(id: string): Promise<TuningSession> {
return this.json(`/api/tuning/sessions/${encodeURIComponent(id)}`);
}
tuningAction(id: string, action: 'pause' | 'resume'): Promise<TuningSession> {
return this.json(`/api/tuning/sessions/${encodeURIComponent(id)}/${action}`, {
method: 'POST',
});
}
cancelTuning(id: string): Promise<TuningSession> {
return this.json(`/api/tuning/sessions/${encodeURIComponent(id)}`, { method: 'DELETE' });
}
decideProposal(
sessionId: string,
proposalId: string,
action: 'approve' | 'reject',
payload: {
feedback?: string;
patch?: { weights: Record<string, number>; params: Record<string, number> };
},
): Promise<TuningSession> {
return this.json(
`/api/tuning/sessions/${encodeURIComponent(sessionId)}/proposals/${encodeURIComponent(proposalId)}/${action}`,
{
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify(payload),
},
);
}
tuningMetrics(
sessionId: string,
trialId: string,
tags: string[] = [],
maxPoints = 1000,
): Promise<TuningMetricsResponse> {
const query = new URLSearchParams({ maxPoints: String(maxPoints) });
if (tags.length) query.set('tags', tags.join(','));
return this.json(
`/api/tuning/sessions/${encodeURIComponent(sessionId)}/trials/${encodeURIComponent(trialId)}/metrics?${query}`,
);
}
presets(): Promise<RewardPreset[]> {
return this.json<{ presets: RewardPreset[] }>('/api/tuning/presets').then(
(value) => value.presets,
);
}
downloadBestPolicy(sessionId: string): Promise<File> {
return this.download(
`/api/tuning/sessions/${encodeURIComponent(sessionId)}/artifacts/best/policy.onnx`,
`best-policy-${sessionId.slice(0, 8)}.onnx`,
);
}
private async download(path: string, name: string): Promise<File> {
const response = await fetch(`${this.endpoint}${path}`, this.requestInit());
if (!response.ok) throw await responseError(response);
const blob = await response.blob();
return new File([blob], `policy-${id.slice(0, 8)}.onnx`, { type: 'application/octet-stream' });
return new File([blob], name, { type: 'application/octet-stream' });
}
}
@@ -37,6 +37,12 @@ describe('LocalTrainingPanel', () => {
headers: { 'Content-Type': 'application/json' },
}),
)
.mockResolvedValueOnce(
new Response(JSON.stringify({ presets: [] }), {
status: 200,
headers: { 'Content-Type': 'application/json' },
}),
)
.mockResolvedValueOnce(
new Response(JSON.stringify(job), {
status: 202,
@@ -49,6 +55,12 @@ describe('LocalTrainingPanel', () => {
headers: { 'Content-Type': 'application/json' },
}),
)
.mockResolvedValueOnce(
new Response(JSON.stringify({ presets: [] }), {
status: 200,
headers: { 'Content-Type': 'application/json' },
}),
)
.mockResolvedValueOnce(
new Response(JSON.stringify({ error: '训练任务不存在或服务已重启' }), {
status: 404,
@@ -62,10 +74,15 @@ describe('LocalTrainingPanel', () => {
});
fireEvent.click(screen.getByRole('button', { name: '连接' }));
expect(await screen.findByText('/opt/unitree_rl_mjlab')).toBeInTheDocument();
const open = vi.spyOn(window, 'open').mockImplementation(() => null);
fireEvent.click(screen.getByRole('button', { name: '打开自调参 Agent 工作台' }));
expect(open).toHaveBeenCalledWith(expect.any(URL), '_blank');
expect(String(open.mock.calls[0][0])).toContain('tuning.html');
expect(String(open.mock.calls[0][0])).not.toContain('secret-token');
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;
await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(3));
const request = fetchMock.mock.calls[2][1] as RequestInit;
expect(JSON.parse(String(request.body))).toMatchObject({
taskId: 'Unitree-Go2-Flat',
numEnvs: 32,
@@ -80,7 +97,7 @@ describe('LocalTrainingPanel', () => {
expect(tokenInput).toBeEnabled();
fireEvent.change(tokenInput, { target: { value: 'new-secret-token' } });
fireEvent.click(screen.getByRole('button', { name: '连接' }));
await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(4));
await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(6));
expect(await screen.findByRole('button', { name: '发起本地训练' })).toBeInTheDocument();
expect(sessionStorage.getItem('mujoco-local-training-token')).toBe('new-secret-token');
});
@@ -1,28 +1,25 @@
import { useEffect, useState, type ReactNode } from 'react';
import { Download, Link, Play, Server, Square } from 'lucide-react';
import { Download, ExternalLink, Link, Play, Server, Square } from 'lucide-react';
import { Badge, Button, ProgressBar, PropertyRow, Select } from '../components/ui';
import { LocalTrainingClient } from './LocalTrainingClient';
import type { TrainingDevice, TrainingJob, TrainingServerInfo, WandbMode } from './types';
import type {
RewardPreset,
TrainingDevice,
TrainingJob,
TrainingServerInfo,
WandbMode,
} from './types';
import {
DEFAULT_TRAINING_ENDPOINT,
localStored,
rememberTrainingConnection,
sessionStored,
TRAINING_ENDPOINT_KEY,
TRAINING_JOB_KEY,
TRAINING_TOKEN_KEY,
} from './storage';
const 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);
}
@@ -37,10 +34,14 @@ function stateLabel(state: TrainingJob['state']): string {
}
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 [endpoint, setEndpoint] = useState(() =>
localStored(TRAINING_ENDPOINT_KEY, DEFAULT_TRAINING_ENDPOINT),
);
const [token, setToken] = useState(() => sessionStored(TRAINING_TOKEN_KEY));
const [server, setServer] = useState<TrainingServerInfo>();
const [job, setJob] = useState<TrainingJob>();
const [presets, setPresets] = useState<RewardPreset[]>([]);
const [rewardPresetId, setRewardPresetId] = useState('');
const [busy, setBusy] = useState(false),
[error, setError] = useState<string>();
const [taskId, setTaskId] = useState('Unitree-Go2-Flat'),
@@ -60,26 +61,30 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
info = await client.health();
setServer(info);
try {
localStorage.setItem(ENDPOINT_KEY, client.endpoint);
sessionStorage.setItem(TOKEN_KEY, client.token);
setPresets(await client.presets());
} catch {
setPresets([]);
}
try {
rememberTrainingConnection(client.endpoint, client.token);
} catch {
/* 当前会话仍可连接 */
}
if (info.tasks.length && !info.tasks.includes(taskId)) setTaskId(info.tasks[0]);
const remembered = info.activeJobId ?? stored(JOB_KEY);
const remembered = info.activeJobId ?? localStored(TRAINING_JOB_KEY);
if (remembered) {
try {
const recovered = await client.job(remembered);
setJob(recovered);
try {
localStorage.setItem(JOB_KEY, recovered.id);
localStorage.setItem(TRAINING_JOB_KEY, recovered.id);
} catch {
/* ignore */
}
} catch {
setJob(undefined);
try {
localStorage.removeItem(JOB_KEY);
localStorage.removeItem(TRAINING_JOB_KEY);
} catch {
/* ignore */
}
@@ -116,6 +121,37 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
};
}, [endpoint, jobId, jobState, token]);
useEffect(() => {
const receive = (event: MessageEvent) => {
if (
event.origin !== window.location.origin ||
!event.source ||
typeof event.data !== 'object'
)
return;
const data = event.data as { type?: string; sessionId?: string };
if (data.type === 'mujoco-tuning-ready') {
(event.source as Window).postMessage(
{ type: 'mujoco-tuning-credentials', endpoint, token },
event.origin,
);
}
if (data.type === 'mujoco-tuning-import-policy' && data.sessionId) {
void new LocalTrainingClient(endpoint, token)
.downloadBestPolicy(data.sessionId)
.then(onPolicyReady)
.catch((value: unknown) => setError(errorText(value)));
}
};
window.addEventListener('message', receive);
return () => window.removeEventListener('message', receive);
}, [endpoint, onPolicyReady, token]);
const openTuningDashboard = () => {
rememberTrainingConnection(endpoint, token);
window.open(new URL('tuning.html', document.baseURI), '_blank');
};
const start = async () => {
setBusy(true);
setError(undefined);
@@ -138,10 +174,11 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
device,
gpuIds: ids,
wandbMode,
rewardPresetId: rewardPresetId || undefined,
});
setJob(next);
try {
localStorage.setItem(JOB_KEY, next.id);
localStorage.setItem(TRAINING_JOB_KEY, next.id);
} catch {
/* ignore */
}
@@ -217,6 +254,15 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
{server?.ready ? '可用' : '离线'}
</Badge>
</div>
{server?.ready && (
<Button
className="mt-2 w-full"
icon={<ExternalLink className="h-3.5 w-3.5" />}
onClick={openTuningDashboard}
>
Agent
</Button>
)}
{server?.ready && !job && (
<div className="mt-3 space-y-2">
<Field label="训练任务">
@@ -286,6 +332,21 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
/>
</Field>
</div>
<Field label="奖励配置">
<Select
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>
))}
</Select>
</Field>
<Field label="实验记录">
<Select
aria-label="W&B 模式"
@@ -368,7 +429,7 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
onClick={() => {
setJob(undefined);
try {
localStorage.removeItem(JOB_KEY);
localStorage.removeItem(TRAINING_JOB_KEY);
} catch {
/* ignore */
}
+30
View File
@@ -0,0 +1,30 @@
export const TRAINING_ENDPOINT_KEY = 'mujoco-local-training-endpoint';
export const TRAINING_JOB_KEY = 'mujoco-local-training-job';
export const TRAINING_TOKEN_KEY = 'mujoco-local-training-token';
export const TUNING_SESSION_KEY = 'mujoco-tuning-session';
export const DEFAULT_TRAINING_ENDPOINT = 'http://127.0.0.1:8765';
export function localStored(key: string, fallback = ''): string {
try {
return localStorage.getItem(key) ?? fallback;
} catch {
return fallback;
}
}
export function sessionStored(key: string): string {
try {
return sessionStorage.getItem(key) ?? '';
} catch {
return '';
}
}
export function rememberTrainingConnection(endpoint: string, token: string): void {
try {
localStorage.setItem(TRAINING_ENDPOINT_KEY, endpoint);
sessionStorage.setItem(TRAINING_TOKEN_KEY, token);
} catch {
/* 当前内存会话仍可继续 */
}
}
+142
View File
@@ -9,6 +9,8 @@ export interface TrainingServerInfo {
python: string;
tasks: string[];
activeJobId?: string;
resourceOwner?: string;
tuning?: TuningCapability;
error?: string;
}
@@ -21,6 +23,7 @@ export interface TrainingRequest {
device: TrainingDevice;
gpuIds: number[];
wandbMode: WandbMode;
rewardPresetId?: string;
}
export interface TrainingJob {
@@ -38,3 +41,142 @@ export interface TrainingJob {
artifactReady: boolean;
artifactName?: string;
}
export type TuningMode = 'automatic' | 'approval';
export type TuningSessionState =
| 'queued'
| 'running'
| 'evaluating'
| 'awaiting_approval'
| 'paused'
| 'interrupted'
| 'succeeded'
| 'failed'
| 'cancelled';
export interface RewardConfiguration {
weights: Record<string, number>;
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 interface TuningCapability {
ready: boolean;
configured: boolean;
apiKeyConfigured: boolean;
frameworkInstalled: boolean;
model: string;
baseUrl: string;
}
export interface TuningCreateRequest {
taskId: 'Unitree-Go2-Flat';
mode: TuningMode;
runName: string;
numEnvs: number;
seed: number;
gpuIds: number[];
trialCount: number;
initialIterations: number;
middleIterations: number;
finalIterations: number;
evalNumEnvs: number;
evalSteps: number;
objectiveWeights: ObjectiveWeights;
fallbackEnabled: boolean;
}
export interface TuningTrial {
id: string;
sessionId: string;
number: number;
state: string;
rung: number;
targetIterations: number;
rewardConfig: RewardConfiguration;
proposalId?: string;
score?: number;
eligible?: boolean;
evaluation?: {
metrics: Record<string, number>;
metricStd?: Record<string, number>;
score?: { score: number; eligible: boolean; components: Record<string, number> };
};
createdAt: string;
startedAt?: string;
endedAt?: string;
message: string;
}
export interface TuningProposal {
id: string;
sessionId: string;
baseTrialId?: string;
state: 'pending' | 'approved' | 'rejected';
source: 'agent' | 'fallback';
patch: { weights: Record<string, number>; params: Record<string, number> };
rationale: string;
expectedImpact: Record<string, string>;
confidence: number;
createdAt: string;
decidedAt?: string;
feedback?: string;
}
export interface TuningAuditEvent {
id: number;
type: string;
payload: Record<string, unknown>;
createdAt: string;
}
export interface TuningSession {
id: string;
state: TuningSessionState;
mode: TuningMode;
createdAt: string;
updatedAt: string;
config: TuningCreateRequest & { rungs: number[]; promote: number[] };
objectiveWeights: ObjectiveWeights;
message: string;
currentTrialId?: string;
bestTrialId?: string;
consecutiveNoImprove: number;
fallbackEnabled: boolean;
trials: TuningTrial[];
proposals: TuningProposal[];
audit: TuningAuditEvent[];
}
export interface ScalarPoint {
step: number;
wallTime: number;
value: number;
}
export interface ScalarSeries {
tag: string;
points: ScalarPoint[];
}
export interface TuningMetricsResponse {
trialId: string;
series: ScalarSeries[];
}
export interface RewardPreset {
id: string;
name: string;
sessionId: string;
trialId: string;
rewardConfig: RewardConfiguration;
createdAt: string;
}
+81
View File
@@ -0,0 +1,81 @@
import { useEffect, useMemo, useRef } from 'react';
import uPlot from 'uplot';
import type { ScalarSeries } from '../training/types';
const COLORS = ['#38d39f', '#60a5fa', '#f59e0b', '#f472b6', '#a78bfa', '#fb7185'];
function smooth(values: (number | null)[], factor: number): (number | null)[] {
if (factor <= 0) return values;
let previous: number | undefined;
return values.map((value) => {
if (value === null) return null;
previous = previous === undefined ? value : factor * previous + (1 - factor) * value;
return previous;
});
}
export function ScalarChart({ series, smoothing }: { series: ScalarSeries[]; smoothing: number }) {
const host = useRef<HTMLDivElement>(null);
const prepared = useMemo(() => {
const steps = Array.from(
new Set(series.flatMap((item) => item.points.map((point) => point.step))),
).sort((a, b) => a - b);
const columns: uPlot.AlignedData = [steps];
for (const item of series) {
const byStep = new Map(item.points.map((point) => [point.step, point.value]));
columns.push(
smooth(
steps.map((step) => byStep.get(step) ?? null),
smoothing,
),
);
}
return columns;
}, [series, smoothing]);
useEffect(() => {
if (!host.current || series.length === 0 || prepared[0].length === 0) return;
const element = host.current;
const chart = new uPlot(
{
width: Math.max(320, element.clientWidth),
height: 360,
title: '训练与评估 Scalars',
cursor: { drag: { x: true, y: true, setScale: true } },
scales: { x: { time: false } },
axes: [
{ stroke: '#8fa0b5', grid: { stroke: '#213044' } },
{ stroke: '#8fa0b5', grid: { stroke: '#213044' } },
],
series: [
{ label: 'Step' },
...series.map((item, index) => ({
label: item.tag,
stroke: COLORS[index % COLORS.length],
width: 2,
spanGaps: true,
})),
],
},
prepared,
element,
);
const observer = new ResizeObserver((entries) => {
const width = entries[0]?.contentRect.width;
if (width) chart.setSize({ width: Math.max(320, Math.floor(width)), height: 360 });
});
observer.observe(element);
return () => {
observer.disconnect();
chart.destroy();
};
}, [prepared, series]);
if (series.length === 0)
return (
<div className="grid h-[360px] place-items-center rounded-lg border border-border bg-app text-xs text-text-tertiary">
trial scalar
</div>
);
return <div ref={host} className="min-w-0 overflow-hidden rounded-lg bg-app p-2" />;
}
@@ -0,0 +1,53 @@
import { fireEvent, render, screen, waitFor } from '@testing-library/react';
import { beforeEach, describe, expect, it, vi } from 'vitest';
import { TuningApp } from './TuningApp';
beforeEach(() => {
localStorage.clear();
sessionStorage.clear();
vi.unstubAllGlobals();
});
describe('TuningApp', () => {
it('连接调参服务并显示 DeepSeek 能力与新建表单', async () => {
const fetchMock = vi.fn((input: string | URL | Request, _init?: RequestInit) => {
void _init;
const url = String(input);
if (url.endsWith('/api/tuning/capabilities'))
return Promise.resolve(
new Response(
JSON.stringify({
ready: true,
configured: true,
apiKeyConfigured: true,
frameworkInstalled: true,
model: 'deepseek-v4-flash',
baseUrl: 'https://api.deepseek.com',
}),
{ status: 200, headers: { 'Content-Type': 'application/json' } },
),
);
return Promise.resolve(
new Response(JSON.stringify({ sessions: [] }), {
status: 200,
headers: { 'Content-Type': 'application/json' },
}),
);
});
vi.stubGlobal('fetch', fetchMock);
render(<TuningApp />);
fireEvent.change(screen.getByLabelText('访问令牌(仅当前标签页)'), {
target: { value: 'training-secret' },
});
fireEvent.click(screen.getByRole('button', { name: '连接/刷新' }));
expect(await screen.findByText(/deepseek-v4-flash/)).toBeInTheDocument();
expect(screen.getByText('新建 Unitree-Go2-Flat 调参 Session')).toBeInTheDocument();
await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(2));
for (const call of fetchMock.mock.calls) {
expect(String(call[0])).not.toContain('training-secret');
const init = call[1];
expect(init).toBeDefined();
expect(new Headers(init?.headers).get('Authorization')).toBe('Bearer training-secret');
}
});
});
+846
View File
@@ -0,0 +1,846 @@
import { useEffect, useMemo, useState, type ReactNode } from 'react';
import {
Bot,
Check,
Download,
FlaskConical,
Pause,
Play,
RefreshCw,
Square,
Upload,
X,
} from 'lucide-react';
import { Badge, Button, ProgressBar, Select } from '../components/ui';
import { LocalTrainingClient } from '../training/LocalTrainingClient';
import {
DEFAULT_TRAINING_ENDPOINT,
localStored,
rememberTrainingConnection,
sessionStored,
TRAINING_ENDPOINT_KEY,
TRAINING_TOKEN_KEY,
TUNING_SESSION_KEY,
} from '../training/storage';
import type {
ObjectiveWeights,
ScalarSeries,
TuningCapability,
TuningMode,
TuningProposal,
TuningSession,
TuningTrial,
} from '../training/types';
import { ScalarChart } from './ScalarChart';
const ACTIVE = new Set(['queued', 'running', 'evaluating', 'awaiting_approval', 'paused']);
const DEFAULT_OBJECTIVES: ObjectiveWeights = {
velocity_tracking: 0.35,
action_smoothness: 0.2,
posture_stability: 0.15,
fall_avoidance: 0.15,
foot_slip: 0.1,
energy: 0.05,
};
const OBJECTIVE_LABELS: Record<keyof ObjectiveWeights, string> = {
velocity_tracking: '速度跟踪',
action_smoothness: '动作平滑',
posture_stability: '姿态稳定',
fall_avoidance: '减少跌倒',
foot_slip: '足端滑移',
energy: '能耗',
};
function errorText(value: unknown): string {
return value instanceof Error ? value.message : String(value);
}
function stateLabel(value: string): string {
return (
{
queued: '排队',
running: '训练中',
evaluating: '评估中',
awaiting_approval: '等待审批',
paused: '已暂停',
interrupted: '已中断',
succeeded: '已完成',
failed: '失败',
cancelled: '已取消',
completed: '完成',
}[value] ?? value
);
}
function downloadFile(file: File): void {
const url = URL.createObjectURL(file);
const anchor = document.createElement('a');
anchor.href = url;
anchor.download = file.name;
anchor.click();
URL.revokeObjectURL(url);
}
function downloadJson(name: string, value: unknown): void {
downloadFile(
new File([JSON.stringify(value, null, 2) + '\n'], name, { type: 'application/json' }),
);
}
export function TuningApp() {
const [endpoint, setEndpoint] = useState(() =>
localStored(TRAINING_ENDPOINT_KEY, DEFAULT_TRAINING_ENDPOINT),
);
const [token, setToken] = useState(() => sessionStored(TRAINING_TOKEN_KEY));
const [capability, setCapability] = useState<TuningCapability>();
const [sessions, setSessions] = useState<TuningSession[]>([]);
const [session, setSession] = useState<TuningSession>();
const [selectedTrialId, setSelectedTrialId] = useState<string>();
const [series, setSeries] = useState<ScalarSeries[]>([]);
const [busy, setBusy] = useState(false);
const [error, setError] = useState<string>();
const [mode, setMode] = useState<TuningMode>('approval');
const [runName, setRunName] = useState('go2-auto-tune');
const [numEnvs, setNumEnvs] = useState(4096);
const [trialCount, setTrialCount] = useState(12);
const [gpuIds, setGpuIds] = useState('0');
const [fallbackEnabled, setFallbackEnabled] = useState(false);
const [objectives, setObjectives] = useState(DEFAULT_OBJECTIVES);
const [smoothing, setSmoothing] = useState(0.3);
const [tagFilter, setTagFilter] = useState('');
const [feedback, setFeedback] = useState('');
const [patchText, setPatchText] = useState('');
useEffect(() => {
const receive = (event: MessageEvent) => {
if (event.origin !== location.origin || typeof event.data !== 'object') return;
const data = event.data as { type?: string; endpoint?: string; token?: string };
if (data.type === 'mujoco-tuning-credentials' && data.endpoint && data.token) {
setEndpoint(data.endpoint);
setToken(data.token);
rememberTrainingConnection(data.endpoint, data.token);
}
};
window.addEventListener('message', receive);
window.opener?.postMessage({ type: 'mujoco-tuning-ready' }, location.origin);
return () => window.removeEventListener('message', receive);
}, []);
const client = () => new LocalTrainingClient(endpoint, token);
const connect = async () => {
setBusy(true);
setError(undefined);
try {
const api = client();
const [nextCapability, nextSessions] = await Promise.all([
api.tuningCapability(),
api.tuningSessions(),
]);
setCapability(nextCapability);
setSessions(nextSessions);
rememberTrainingConnection(api.endpoint, api.token);
const remembered = localStored(TUNING_SESSION_KEY);
const target = nextSessions.find((item) => item.id === remembered) ?? nextSessions[0];
if (target) {
const detail = await api.tuningSession(target.id);
setSession(detail);
setSelectedTrialId(detail.currentTrialId ?? detail.bestTrialId ?? detail.trials.at(-1)?.id);
}
} catch (value) {
setError(errorText(value));
} finally {
setBusy(false);
}
};
const sessionId = session?.id;
const sessionState = session?.state;
useEffect(() => {
if (!sessionId || !sessionState || !ACTIVE.has(sessionState)) return;
const timer = window.setInterval(() => {
void new LocalTrainingClient(endpoint, token)
.tuningSession(sessionId)
.then((next) => {
setSession(next);
setSelectedTrialId((current) => current ?? next.currentTrialId ?? next.trials.at(-1)?.id);
})
.catch((value: unknown) => setError(errorText(value)));
}, 2000);
return () => window.clearInterval(timer);
}, [endpoint, sessionId, sessionState, token]);
useEffect(() => {
if (!sessionId || !selectedTrialId) return;
let disposed = false;
const refresh = () =>
new LocalTrainingClient(endpoint, token)
.tuningMetrics(sessionId, selectedTrialId, [], 1200)
.then((value) => {
if (!disposed) setSeries(value.series);
})
.catch((value: unknown) => {
if (!disposed) setError(errorText(value));
});
void refresh();
const timer = window.setInterval(() => void refresh(), 3000);
return () => {
disposed = true;
window.clearInterval(timer);
};
}, [endpoint, selectedTrialId, sessionId, token]);
const start = async () => {
setBusy(true);
setError(undefined);
try {
const ids = gpuIds
.split(/[\s,]+/)
.filter(Boolean)
.map(Number);
if (!ids.length || ids.some((value) => !Number.isInteger(value) || value < 0))
throw new Error('GPU 编号必须是非负整数');
const next = await client().startTuning({
taskId: 'Unitree-Go2-Flat',
mode,
runName,
numEnvs,
seed: 42,
gpuIds: ids,
trialCount,
initialIterations: 300,
middleIterations: 900,
finalIterations: 2000,
evalNumEnvs: 256,
evalSteps: 1000,
objectiveWeights: objectives,
fallbackEnabled,
});
setSession(next);
setSelectedTrialId(next.currentTrialId ?? next.trials[0]?.id);
localStorage.setItem(TUNING_SESSION_KEY, next.id);
} catch (value) {
setError(errorText(value));
} finally {
setBusy(false);
}
};
const runAction = async (action: 'pause' | 'resume' | 'cancel') => {
if (!session) return;
setBusy(true);
try {
setSession(
action === 'cancel'
? await client().cancelTuning(session.id)
: await client().tuningAction(session.id, action),
);
} catch (value) {
setError(errorText(value));
} finally {
setBusy(false);
}
};
const pending = session?.proposals
.slice()
.reverse()
.find((item) => item.state === 'pending');
const decide = async (proposal: TuningProposal, action: 'approve' | 'reject') => {
if (!session) return;
setBusy(true);
try {
const payload: { feedback?: string; patch?: TuningProposal['patch'] } = { feedback };
if (action === 'approve')
payload.patch = JSON.parse(
patchText || JSON.stringify(proposal.patch),
) as TuningProposal['patch'];
setSession(await client().decideProposal(session.id, proposal.id, action, payload));
setFeedback('');
setPatchText('');
} catch (value) {
setError(errorText(value));
} finally {
setBusy(false);
}
};
const completed = session?.trials.filter((trial) => trial.state === 'completed').length ?? 0;
const totalStages = session
? session.config.trialCount + session.config.promote.slice(1).reduce((a, b) => a + b, 0)
: 1;
const selectedTrial = session?.trials.find((trial) => trial.id === selectedTrialId);
const filteredSeries = useMemo(
() => series.filter((item) => item.tag.toLowerCase().includes(tagFilter.toLowerCase())),
[series, tagFilter],
);
return (
<main className="min-h-full bg-app text-text-primary">
<header className="flex min-h-14 flex-wrap items-center justify-between gap-3 border-b border-border bg-surface px-5 py-3">
<div>
<h1 className="flex items-center gap-2 text-base font-semibold">
<Bot className="h-5 w-5 text-accent" /> Go2 Agent
</h1>
<p className="mt-0.5 text-[11px] text-text-tertiary">
DeepSeek · · TensorBoard Scalars
</p>
</div>
<div className="flex items-center gap-2">
{capability && (
<Badge tone={capability.configured ? 'success' : 'warning'}>
{capability.model} · {capability.configured ? '已配置' : '未配置'}
</Badge>
)}
<Button
icon={<RefreshCw className="h-3.5 w-3.5" />}
disabled={busy || !token}
onClick={() => void connect()}
>
/
</Button>
</div>
</header>
<section className="grid gap-3 border-b border-border bg-surface/60 p-3 lg:grid-cols-[1fr_1fr_auto]">
<Field label="训练服务地址">
<input
className="field h-8 w-full px-2 text-xs"
value={endpoint}
onChange={(event) => setEndpoint(event.target.value)}
/>
</Field>
<Field label="访问令牌(仅当前标签页)">
<input
type="password"
className="field h-8 w-full px-2 text-xs"
value={token}
onChange={(event) => setToken(event.target.value)}
/>
</Field>
<Button
className="self-end"
icon={<FlaskConical className="h-3.5 w-3.5" />}
disabled={!capability?.configured || busy}
onClick={() =>
void client()
.testTuningAgent()
.then(() => setError(undefined))
.catch((value: unknown) => setError(errorText(value)))
}
>
Agent
</Button>
</section>
{!session ? (
<NewSessionForm
mode={mode}
setMode={setMode}
runName={runName}
setRunName={setRunName}
numEnvs={numEnvs}
setNumEnvs={setNumEnvs}
trialCount={trialCount}
setTrialCount={setTrialCount}
gpuIds={gpuIds}
setGpuIds={setGpuIds}
fallback={fallbackEnabled}
setFallback={setFallbackEnabled}
objectives={objectives}
setObjectives={setObjectives}
start={start}
busy={busy}
sessions={sessions}
open={async (id) => {
const next = await client().tuningSession(id);
setSession(next);
setSelectedTrialId(next.currentTrialId ?? next.bestTrialId ?? next.trials.at(-1)?.id);
}}
/>
) : (
<div className="grid min-h-[calc(100vh-132px)] grid-cols-1 xl:grid-cols-[280px_minmax(0,1fr)_360px]">
<aside className="border-r border-border bg-surface p-3">
<div className="mb-3 flex items-center justify-between">
<div>
<p className="text-xs font-semibold">{session.config.runName}</p>
<p className="font-mono text-[9px] text-text-tertiary">{session.id}</p>
</div>
<Badge
tone={
session.state === 'succeeded'
? 'success'
: session.state === 'failed'
? 'warning'
: 'accent'
}
>
{stateLabel(session.state)}
</Badge>
</div>
<ProgressBar
value={Math.min(1, completed / totalStages)}
label={`${completed} / ${totalStages} 阶段`}
/>
<p className="mt-2 rounded bg-app p-2 text-[10px] leading-4 text-text-secondary">
{session.message}
</p>
<div className="mt-3 grid grid-cols-3 gap-1">
{session.state !== 'paused' ? (
<Button
icon={<Pause className="h-3 w-3" />}
disabled={busy || !ACTIVE.has(session.state)}
onClick={() => void runAction('pause')}
>
</Button>
) : (
<Button
icon={<Play className="h-3 w-3" />}
disabled={busy}
onClick={() => void runAction('resume')}
>
</Button>
)}
<Button
variant="danger"
icon={<Square className="h-3 w-3" />}
disabled={busy || !ACTIVE.has(session.state)}
onClick={() => void runAction('cancel')}
>
</Button>
<Button
onClick={() => {
setSession(undefined);
setSeries([]);
}}
>
</Button>
</div>
<h2 className="mb-2 mt-4 text-[10px] font-semibold uppercase tracking-wider text-text-tertiary">
Trials
</h2>
<div className="max-h-[58vh] space-y-1 overflow-auto panel-scroll">
{session.trials.map((trial) => (
<button
key={trial.id}
className={`w-full rounded border p-2 text-left ${selectedTrialId === trial.id ? 'border-accent bg-accent/10' : 'border-border bg-app hover:bg-element-hover'}`}
onClick={() => setSelectedTrialId(trial.id)}
>
<div className="flex justify-between text-[10px]">
<span>
T{trial.number} · R{trial.rung}
</span>
<span>{stateLabel(trial.state)}</span>
</div>
<div className="mt-1 flex justify-between font-mono text-[9px] text-text-tertiary">
<span>{trial.targetIterations} it</span>
<span>
{trial.score === undefined || trial.score === null
? '—'
: trial.score.toFixed(4)}
</span>
</div>
</button>
))}
</div>
</aside>
<section className="min-w-0 space-y-3 p-4">
<div className="flex flex-wrap items-end gap-3">
<Field label="Tag 过滤">
<input
className="field h-7 w-64 px-2 text-xs"
value={tagFilter}
onChange={(event) => setTagFilter(event.target.value)}
placeholder="Episode_Reward / Evaluation"
/>
</Field>
<Field label={`平滑 ${smoothing.toFixed(2)}`}>
<input
type="range"
min="0"
max="0.95"
step="0.05"
value={smoothing}
onChange={(event) => setSmoothing(Number(event.target.value))}
/>
</Field>
<span className="text-[10px] text-text-tertiary">
{selectedTrial
? `Trial ${selectedTrial.number} / rung ${selectedTrial.rung}`
: '请选择 trial'}
</span>
</div>
<ScalarChart series={filteredSeries} smoothing={smoothing} />
{selectedTrial?.evaluation && <EvaluationCard trial={selectedTrial} />}
<Leaderboard trials={session.trials} select={setSelectedTrialId} />
</section>
<aside className="space-y-3 border-l border-border bg-surface p-3">
{pending && (
<ApprovalCard
proposal={pending}
patchText={patchText || JSON.stringify(pending.patch, null, 2)}
setPatchText={setPatchText}
feedback={feedback}
setFeedback={setFeedback}
decide={decide}
busy={busy}
/>
)}
<AgentTimeline session={session} />
{session.bestTrialId && (
<div className="rounded-lg border border-border bg-app p-3">
<h2 className="text-xs font-semibold"></h2>
<div className="mt-2 grid grid-cols-2 gap-2">
<Button
icon={<Download className="h-3.5 w-3.5" />}
onClick={() =>
void client()
.downloadBestPolicy(session.id)
.then(downloadFile)
.catch((value: unknown) => setError(errorText(value)))
}
>
ONNX
</Button>
<Button
icon={<Upload className="h-3.5 w-3.5" />}
onClick={() => {
if (window.opener)
window.opener.postMessage(
{ type: 'mujoco-tuning-import-policy', sessionId: session.id },
location.origin,
);
else void client().downloadBestPolicy(session.id).then(downloadFile);
}}
>
</Button>
<Button
className="col-span-2"
onClick={() => {
const best = session.trials.find((trial) => trial.id === session.bestTrialId);
if (best)
downloadJson(
`reward-preset-${session.id.slice(0, 8)}.json`,
best.rewardConfig,
);
}}
>
Reward Preset
</Button>
</div>
</div>
)}
</aside>
</div>
)}
{error && (
<div
role="alert"
className="fixed bottom-4 left-1/2 z-50 max-w-2xl -translate-x-1/2 rounded-lg border border-danger-border bg-danger-soft px-4 py-3 text-xs text-danger shadow-xl"
>
{error}
</div>
)}
</main>
);
}
function NewSessionForm(props: {
mode: TuningMode;
setMode(value: TuningMode): void;
runName: string;
setRunName(value: string): void;
numEnvs: number;
setNumEnvs(value: number): void;
trialCount: number;
setTrialCount(value: number): void;
gpuIds: string;
setGpuIds(value: string): void;
fallback: boolean;
setFallback(value: boolean): void;
objectives: ObjectiveWeights;
setObjectives(value: ObjectiveWeights): void;
start(): Promise<void>;
busy: boolean;
sessions: TuningSession[];
open(id: string): Promise<void>;
}) {
return (
<div className="mx-auto grid max-w-6xl gap-4 p-5 lg:grid-cols-[2fr_1fr]">
<section className="rounded-xl border border-border bg-surface p-5">
<h2 className="text-sm font-semibold"> Unitree-Go2-Flat Session</h2>
<div className="mt-4 grid gap-3 sm:grid-cols-2">
<Field label="运行名称">
<input
className="field h-8 w-full px-2 text-xs"
value={props.runName}
onChange={(e) => props.setRunName(e.target.value)}
/>
</Field>
<Field label="模式">
<Select
className="w-full"
value={props.mode}
onChange={(e) => props.setMode(e.target.value as TuningMode)}
>
<option value="approval"></option>
<option value="automatic"></option>
</Select>
</Field>
<NumberInput
label="并行环境"
value={props.numEnvs}
min={1}
max={16384}
change={props.setNumEnvs}
/>
<NumberInput
label="唯一配置数"
value={props.trialCount}
min={4}
max={20}
change={props.setTrialCount}
/>
<Field label="GPU 编号">
<input
className="field h-8 w-full px-2 text-xs"
value={props.gpuIds}
onChange={(e) => props.setGpuIds(e.target.value)}
/>
</Field>
<label className="flex items-end gap-2 pb-2 text-xs text-text-secondary">
<input
type="checkbox"
checked={props.fallback}
onChange={(e) => props.setFallback(e.target.checked)}
/>
Agent Optuna fallback
</label>
</div>
<h3 className="mb-2 mt-5 text-xs font-semibold"></h3>
<div className="grid gap-2 sm:grid-cols-2">
{(Object.keys(props.objectives) as (keyof ObjectiveWeights)[]).map((key) => (
<label
key={key}
className="grid grid-cols-[90px_1fr_42px] items-center gap-2 text-[10px] text-text-secondary"
>
<span>{OBJECTIVE_LABELS[key]}</span>
<input
type="range"
min="0"
max="1"
step="0.05"
value={props.objectives[key]}
onChange={(e) =>
props.setObjectives({ ...props.objectives, [key]: Number(e.target.value) })
}
/>
<span>{Math.round(props.objectives[key] * 100)}%</span>
</label>
))}
</div>
<p
className={`mt-2 text-[10px] ${Math.abs(Object.values(props.objectives).reduce((a, b) => a + b, 0) - 1) < 1e-6 ? 'text-success' : 'text-danger'}`}
>
{Math.round(Object.values(props.objectives).reduce((a, b) => a + b, 0) * 100)}%
</p>
<Button
variant="primary"
className="mt-4 w-full"
icon={<Play className="h-3.5 w-3.5" />}
disabled={
props.busy ||
Math.abs(Object.values(props.objectives).reduce((a, b) => a + b, 0) - 1) > 1e-6
}
onClick={() => void props.start()}
>
</Button>
</section>
<section className="rounded-xl border border-border bg-surface p-4">
<h2 className="text-xs font-semibold"> Sessions</h2>
<div className="mt-3 space-y-2">
{props.sessions.map((item) => (
<button
key={item.id}
className="w-full rounded border border-border bg-app p-2 text-left text-[10px] hover:bg-element-hover"
onClick={() => void props.open(item.id)}
>
<div className="flex justify-between">
<span>{item.config?.runName ?? item.id.slice(0, 8)}</span>
<Badge>{stateLabel(item.state)}</Badge>
</div>
<p className="mt-1 font-mono text-[9px] text-text-tertiary">{item.id}</p>
</button>
))}
</div>
</section>
</div>
);
}
function EvaluationCard({ trial }: { trial: TuningTrial }) {
return (
<div className="rounded-lg border border-border bg-surface p-3">
<h2 className="mb-2 text-xs font-semibold"></h2>
<div className="grid gap-2 sm:grid-cols-2 lg:grid-cols-4">
{Object.entries(trial.evaluation?.metrics ?? {}).map(([name, value]) => (
<div key={name} className="rounded bg-app p-2">
<p className="truncate text-[9px] text-text-tertiary" title={name}>
{name}
</p>
<p className="mt-1 font-mono text-sm">{value.toFixed(5)}</p>
</div>
))}
</div>
</div>
);
}
function Leaderboard({ trials, select }: { trials: TuningTrial[]; select(id: string): void }) {
const ranked = [...trials]
.filter((trial) => trial.score !== undefined && trial.score !== null)
.sort((a, b) => (b.score ?? -999) - (a.score ?? -999));
return (
<div className="rounded-lg border border-border bg-surface p-3">
<h2 className="mb-2 text-xs font-semibold"></h2>
<div className="overflow-auto">
<table className="w-full text-left text-[10px]">
<thead className="text-text-tertiary">
<tr>
<th></th>
<th>Trial</th>
<th>Rung</th>
<th></th>
<th></th>
</tr>
</thead>
<tbody>
{ranked.map((trial, index) => (
<tr
key={trial.id}
className="cursor-pointer border-t border-border hover:bg-element-hover"
onClick={() => select(trial.id)}
>
<td className="py-1.5">{index + 1}</td>
<td>T{trial.number}</td>
<td>{trial.rung}</td>
<td className="font-mono">{trial.score?.toFixed(5)}</td>
<td>{trial.eligible ? '通过' : '未通过'}</td>
</tr>
))}
</tbody>
</table>
</div>
</div>
);
}
function ApprovalCard(props: {
proposal: TuningProposal;
patchText: string;
setPatchText(value: string): void;
feedback: string;
setFeedback(value: string): void;
decide(proposal: TuningProposal, action: 'approve' | 'reject'): Promise<void>;
busy: boolean;
}) {
return (
<div className="rounded-lg border border-accent/50 bg-accent/5 p-3">
<div className="flex justify-between">
<h2 className="text-xs font-semibold"></h2>
<Badge tone="accent"> {Math.round(props.proposal.confidence * 100)}%</Badge>
</div>
<p className="mt-2 text-[10px] leading-4 text-text-secondary">{props.proposal.rationale}</p>
<label className="mt-2 block text-[10px] text-text-tertiary">
Patch
<textarea
className="field mt-1 h-36 w-full resize-y p-2 font-mono text-[10px]"
value={props.patchText}
onChange={(e) => props.setPatchText(e.target.value)}
/>
</label>
<label className="mt-2 block text-[10px] text-text-tertiary">
<input
className="field mt-1 h-7 w-full px-2"
value={props.feedback}
onChange={(e) => props.setFeedback(e.target.value)}
/>
</label>
<div className="mt-2 grid grid-cols-2 gap-2">
<Button
variant="primary"
icon={<Check className="h-3 w-3" />}
disabled={props.busy}
onClick={() => void props.decide(props.proposal, 'approve')}
>
/
</Button>
<Button
variant="danger"
icon={<X className="h-3 w-3" />}
disabled={props.busy}
onClick={() => void props.decide(props.proposal, 'reject')}
>
</Button>
</div>
</div>
);
}
function AgentTimeline({ session }: { session: TuningSession }) {
return (
<div className="rounded-lg border border-border bg-app p-3">
<h2 className="text-xs font-semibold">Agent 线</h2>
<div className="mt-2 max-h-72 space-y-2 overflow-auto panel-scroll">
{session.audit
.slice()
.reverse()
.map((event) => (
<div key={event.id} className="border-l border-accent/40 pl-2">
<p className="text-[10px] text-text-secondary">{event.type}</p>
<p className="text-[9px] text-text-tertiary">
{new Date(event.createdAt).toLocaleString()}
</p>
</div>
))}
</div>
</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 NumberInput({
label,
value,
min,
max,
change,
}: {
label: string;
value: number;
min: number;
max: number;
change(value: number): void;
}) {
return (
<Field label={label}>
<input
type="number"
min={min}
max={max}
className="field h-8 w-full px-2 text-xs"
value={value}
onChange={(e) => change(Number(e.target.value))}
/>
</Field>
);
}
+14
View File
@@ -0,0 +1,14 @@
import { StrictMode } from 'react';
import { createRoot } from 'react-dom/client';
import 'uplot/dist/uPlot.min.css';
import '../styles.css';
import { ErrorBoundary } from '../app/ErrorBoundary';
import { TuningApp } from './TuningApp';
createRoot(document.getElementById('root')!).render(
<StrictMode>
<ErrorBoundary>
<TuningApp />
</ErrorBoundary>
</StrictMode>,
);
+26
View File
@@ -0,0 +1,26 @@
<!doctype html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8" />
<meta name="viewport" content="width=device-width,initial-scale=1" />
<meta name="theme-color" content="#09111e" />
<meta name="description" content="Unitree Go2 奖励函数自调参 Agent 工作台" />
<link rel="icon" href="data:," />
<title>Go2 自调参 Agent</title>
<style>
html,
body,
#root {
height: 100%;
margin: 0;
}
body {
background: #09111e;
}
</style>
</head>
<body>
<div id="root"></div>
<script type="module" src="/src/tuning/main.tsx"></script>
</body>
</html>
+6
View File
@@ -86,6 +86,12 @@ export default defineConfig({
outDir: '../web-platform-dist',
emptyOutDir: true,
target: 'es2022',
rollupOptions: {
input: {
main: resolve(dirname(fileURLToPath(import.meta.url)), 'index.html'),
tuning: resolve(dirname(fileURLToPath(import.meta.url)), 'tuning.html'),
},
},
modulePreload: { polyfill: false },
},
worker: { format: 'es' },