feat(training): release V0.8 自调参 Agent
This commit is contained in:
@@ -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/`。
|
||||
|
||||
@@ -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)。
|
||||
|
||||
|
||||
@@ -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}`));
|
||||
|
||||
@@ -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 */
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
/* 当前内存会话仍可继续 */
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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');
|
||||
}
|
||||
});
|
||||
});
|
||||
@@ -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>
|
||||
);
|
||||
}
|
||||
@@ -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>,
|
||||
);
|
||||
@@ -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>
|
||||
@@ -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' },
|
||||
|
||||
Reference in New Issue
Block a user