feat(tuning): release V0.8.2 Agent 界面重构
This commit is contained in:
@@ -62,13 +62,14 @@ describe('LocalTrainingClient', () => {
|
||||
);
|
||||
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.tuningMetrics('a'.repeat(32), 'b'.repeat(32), ['Evaluation/score'], 500, 120);
|
||||
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]).toContain('afterStep=120');
|
||||
expect(fetchMock.mock.calls[0][0]).not.toContain('deep-secret');
|
||||
const approval = fetchMock.mock.calls[1][1] as RequestInit;
|
||||
expect(approval.method).toBe('POST');
|
||||
@@ -119,11 +120,25 @@ describe('LocalTrainingClient', () => {
|
||||
});
|
||||
await client.tuningSession('a'.repeat(32));
|
||||
await client.tuningAction('a'.repeat(32), 'pause');
|
||||
await client.setTuningMode('a'.repeat(32), 'approval');
|
||||
await client.setTuningConstraints('a'.repeat(32), 2, {
|
||||
'weights.pose': { kind: 'range', min: 0.5, max: 1.5 },
|
||||
});
|
||||
await client.stepTuning('a'.repeat(32));
|
||||
await client.rollbackTuning('a'.repeat(32), 'b'.repeat(32), true);
|
||||
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);
|
||||
expect(fetchMock).toHaveBeenCalledTimes(13);
|
||||
const modeRequest = fetchMock.mock.calls.find(([url]) => String(url).endsWith('/mode'))?.[1] as
|
||||
RequestInit | undefined;
|
||||
expect(JSON.parse(String(modeRequest?.body))).toEqual({ mode: 'approval' });
|
||||
const constraintsRequest = fetchMock.mock.calls.find(([url]) =>
|
||||
String(url).endsWith('/constraints'),
|
||||
)?.[1] as RequestInit | undefined;
|
||||
expect(constraintsRequest?.method).toBe('PUT');
|
||||
expect(JSON.parse(String(constraintsRequest?.body))).toMatchObject({ revision: 2 });
|
||||
});
|
||||
|
||||
it('拒绝非 HTTP 地址和空访问令牌', () => {
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import type {
|
||||
ParameterConstraint,
|
||||
RewardPreset,
|
||||
TuningCapability,
|
||||
TuningCreateRequest,
|
||||
@@ -101,9 +102,41 @@ export class LocalTrainingClient {
|
||||
method: 'POST',
|
||||
});
|
||||
}
|
||||
setTuningMode(id: string, mode: TuningSession['mode']): Promise<TuningSession> {
|
||||
return this.json(`/api/tuning/sessions/${encodeURIComponent(id)}/mode`, {
|
||||
method: 'POST',
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify({ mode }),
|
||||
});
|
||||
}
|
||||
cancelTuning(id: string): Promise<TuningSession> {
|
||||
return this.json(`/api/tuning/sessions/${encodeURIComponent(id)}`, { method: 'DELETE' });
|
||||
}
|
||||
setTuningConstraints(
|
||||
id: string,
|
||||
revision: number,
|
||||
constraints: Record<string, ParameterConstraint>,
|
||||
): Promise<TuningSession> {
|
||||
return this.json(`/api/tuning/sessions/${encodeURIComponent(id)}/constraints`, {
|
||||
method: 'PUT',
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify({ revision, constraints }),
|
||||
});
|
||||
}
|
||||
stepTuning(id: string): Promise<TuningSession> {
|
||||
return this.json(`/api/tuning/sessions/${encodeURIComponent(id)}/step`, {
|
||||
method: 'POST',
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify({ count: 1 }),
|
||||
});
|
||||
}
|
||||
rollbackTuning(id: string, trialId: string, checkpoint = false): Promise<TuningSession> {
|
||||
return this.json(`/api/tuning/sessions/${encodeURIComponent(id)}/rollback`, {
|
||||
method: 'POST',
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify({ trialId, checkpoint }),
|
||||
});
|
||||
}
|
||||
decideProposal(
|
||||
sessionId: string,
|
||||
proposalId: string,
|
||||
@@ -127,11 +160,16 @@ export class LocalTrainingClient {
|
||||
trialId: string,
|
||||
tags: string[] = [],
|
||||
maxPoints = 1000,
|
||||
afterStep?: number,
|
||||
signal?: AbortSignal,
|
||||
): Promise<TuningMetricsResponse> {
|
||||
const query = new URLSearchParams({ maxPoints: String(maxPoints) });
|
||||
if (tags.length) query.set('tags', tags.join(','));
|
||||
if (afterStep !== undefined && Number.isFinite(afterStep))
|
||||
query.set('afterStep', String(afterStep));
|
||||
return this.json(
|
||||
`/api/tuning/sessions/${encodeURIComponent(sessionId)}/trials/${encodeURIComponent(trialId)}/metrics?${query}`,
|
||||
{ signal },
|
||||
);
|
||||
}
|
||||
presets(): Promise<RewardPreset[]> {
|
||||
|
||||
@@ -3,6 +3,7 @@ import { beforeEach, describe, expect, it, vi } from 'vitest';
|
||||
import { LocalTrainingPanel } from './LocalTrainingPanel';
|
||||
|
||||
beforeEach(() => {
|
||||
vi.restoreAllMocks();
|
||||
localStorage.clear();
|
||||
sessionStorage.clear();
|
||||
vi.unstubAllGlobals();
|
||||
@@ -101,4 +102,60 @@ describe('LocalTrainingPanel', () => {
|
||||
expect(await screen.findByRole('button', { name: '发起本地训练' })).toBeInTheDocument();
|
||||
expect(sessionStorage.getItem('mujoco-local-training-token')).toBe('new-secret-token');
|
||||
});
|
||||
|
||||
it('接收调参窗口传回的策略文件,无需主工作台重复持有令牌', async () => {
|
||||
const onPolicyReady = vi.fn();
|
||||
const reply = vi.spyOn(window, 'postMessage').mockImplementation(() => undefined);
|
||||
render(<LocalTrainingPanel onPolicyReady={onPolicyReady} />);
|
||||
const policy = new File([new Uint8Array([1, 2, 3])], 'best-policy.onnx', {
|
||||
type: 'application/octet-stream',
|
||||
});
|
||||
|
||||
window.dispatchEvent(
|
||||
new MessageEvent('message', {
|
||||
origin: window.location.origin,
|
||||
source: window,
|
||||
data: {
|
||||
type: 'mujoco-tuning-import-policy',
|
||||
sessionId: 'a'.repeat(32),
|
||||
policy,
|
||||
},
|
||||
}),
|
||||
);
|
||||
|
||||
await waitFor(() => expect(onPolicyReady).toHaveBeenCalledWith(policy));
|
||||
expect(reply).toHaveBeenCalledWith(
|
||||
expect.objectContaining({
|
||||
type: 'mujoco-tuning-import-policy-result',
|
||||
ok: true,
|
||||
}),
|
||||
window.location.origin,
|
||||
);
|
||||
});
|
||||
|
||||
it('旧调参消息缺少主工作台令牌时显示错误而不是抛出未处理异常', async () => {
|
||||
const reply = vi.spyOn(window, 'postMessage').mockImplementation(() => undefined);
|
||||
render(<LocalTrainingPanel onPolicyReady={vi.fn()} />);
|
||||
|
||||
window.dispatchEvent(
|
||||
new MessageEvent('message', {
|
||||
origin: window.location.origin,
|
||||
source: window,
|
||||
data: {
|
||||
type: 'mujoco-tuning-import-policy',
|
||||
sessionId: 'a'.repeat(32),
|
||||
},
|
||||
}),
|
||||
);
|
||||
|
||||
expect(await screen.findByRole('alert')).toHaveTextContent('请输入训练服务访问令牌');
|
||||
expect(reply).toHaveBeenCalledWith(
|
||||
expect.objectContaining({
|
||||
type: 'mujoco-tuning-import-policy-result',
|
||||
ok: false,
|
||||
error: '请输入训练服务访问令牌',
|
||||
}),
|
||||
window.location.origin,
|
||||
);
|
||||
});
|
||||
});
|
||||
|
||||
@@ -129,18 +129,44 @@ export function LocalTrainingPanel({ onPolicyReady }: { onPolicyReady(file: File
|
||||
typeof event.data !== 'object'
|
||||
)
|
||||
return;
|
||||
const data = event.data as { type?: string; sessionId?: string };
|
||||
const data = event.data as { type?: string; sessionId?: string; policy?: unknown };
|
||||
const source = event.source as Window;
|
||||
if (data.type === 'mujoco-tuning-ready') {
|
||||
(event.source as Window).postMessage(
|
||||
{ type: 'mujoco-tuning-credentials', endpoint, token },
|
||||
event.origin,
|
||||
);
|
||||
source.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)));
|
||||
const reply = (ok: boolean, message?: string) => {
|
||||
try {
|
||||
source.postMessage(
|
||||
{
|
||||
type: 'mujoco-tuning-import-policy-result',
|
||||
sessionId: data.sessionId,
|
||||
ok,
|
||||
error: message,
|
||||
},
|
||||
event.origin,
|
||||
);
|
||||
} catch {
|
||||
/* 调参窗口可能已关闭;不影响主工作台继续导入 */
|
||||
}
|
||||
};
|
||||
void (async () => {
|
||||
try {
|
||||
const policy =
|
||||
data.policy === undefined
|
||||
? await new LocalTrainingClient(endpoint, token).downloadBestPolicy(data.sessionId!)
|
||||
: data.policy;
|
||||
if (!(policy instanceof File) || !/\.onnx$/i.test(policy.name))
|
||||
throw new Error('调参工作台返回的 ONNX 策略无效');
|
||||
if (policy.size > 64 * 1024 * 1024) throw new Error('ONNX 策略不能超过 64 MiB');
|
||||
onPolicyReady(policy);
|
||||
reply(true);
|
||||
} catch (value) {
|
||||
const message = errorText(value);
|
||||
setError(message);
|
||||
reply(false, message);
|
||||
}
|
||||
})();
|
||||
}
|
||||
};
|
||||
window.addEventListener('message', receive);
|
||||
|
||||
@@ -43,6 +43,7 @@ export interface TrainingJob {
|
||||
}
|
||||
|
||||
export type TuningMode = 'automatic' | 'approval';
|
||||
export type TuningRunPolicy = 'continuous' | 'step';
|
||||
export type TuningSessionState =
|
||||
| 'queued'
|
||||
| 'running'
|
||||
@@ -94,11 +95,14 @@ export interface TuningCreateRequest {
|
||||
fallbackEnabled: boolean;
|
||||
}
|
||||
|
||||
export type TuningTrialState =
|
||||
'queued' | 'training' | 'evaluating' | 'completed' | 'interrupted' | 'failed' | 'cancelled';
|
||||
|
||||
export interface TuningTrial {
|
||||
id: string;
|
||||
sessionId: string;
|
||||
number: number;
|
||||
state: string;
|
||||
state: TuningTrialState;
|
||||
rung: number;
|
||||
targetIterations: number;
|
||||
rewardConfig: RewardConfiguration;
|
||||
@@ -138,6 +142,18 @@ export interface TuningAuditEvent {
|
||||
createdAt: string;
|
||||
}
|
||||
|
||||
export type ParameterConstraint =
|
||||
{ kind: 'range'; min: number; max: number } | { kind: 'fixed'; value: number };
|
||||
|
||||
export interface TuningControlState {
|
||||
runPolicy: TuningRunPolicy;
|
||||
dispatchTokens: number;
|
||||
constraintsRevision: number;
|
||||
constraints: Record<string, ParameterConstraint>;
|
||||
activeBaseTrialId?: string;
|
||||
effectiveAfterCurrent: boolean;
|
||||
}
|
||||
|
||||
export interface TuningSession {
|
||||
id: string;
|
||||
state: TuningSessionState;
|
||||
@@ -154,6 +170,8 @@ export interface TuningSession {
|
||||
trials: TuningTrial[];
|
||||
proposals: TuningProposal[];
|
||||
audit: TuningAuditEvent[];
|
||||
/** 旧服务响应可能缺失;前端会回退到 continuous + 空约束。 */
|
||||
control?: TuningControlState;
|
||||
}
|
||||
|
||||
export interface ScalarPoint {
|
||||
@@ -170,6 +188,8 @@ export interface ScalarSeries {
|
||||
export interface TuningMetricsResponse {
|
||||
trialId: string;
|
||||
series: ScalarSeries[];
|
||||
/** 本响应中最大的 step;用于下一次增量请求。 */
|
||||
nextStep?: number;
|
||||
}
|
||||
|
||||
export interface RewardPreset {
|
||||
|
||||
Reference in New Issue
Block a user