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
@@ -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;
}