feat(web-platform): release V0.5.1 前端 RL 接入
This commit is contained in:
@@ -0,0 +1,25 @@
|
||||
import {fireEvent,render,screen,waitFor} from '@testing-library/react';
|
||||
import {beforeEach,describe,expect,it,vi} from 'vitest';
|
||||
import {LocalTrainingPanel} from './LocalTrainingPanel';
|
||||
|
||||
beforeEach(()=>{localStorage.clear();vi.unstubAllGlobals();});
|
||||
|
||||
describe('LocalTrainingPanel',()=>{
|
||||
it('连接本地服务并从图形界面发起训练请求',async()=>{
|
||||
const health={version:'0.1.0',ready:true,trainerRoot:'/opt/unitree_rl_mjlab',python:'/env/bin/python',tasks:['Unitree-Go2-Flat']};
|
||||
const job={id:'a'.repeat(32),state:'queued',taskId:'Unitree-Go2-Flat',createdAt:'2025-01-01T00:00:00Z',iteration:0,maxIterations:2000,progress:0,message:'等待启动',logs:[],artifactReady:false};
|
||||
const fetchMock=vi.fn()
|
||||
.mockResolvedValueOnce(new Response(JSON.stringify(health),{status:200,headers:{'Content-Type':'application/json'}}))
|
||||
.mockResolvedValueOnce(new Response(JSON.stringify(job),{status:202,headers:{'Content-Type':'application/json'}}));
|
||||
vi.stubGlobal('fetch',fetchMock);
|
||||
render(<LocalTrainingPanel onPolicyReady={vi.fn()}/>);
|
||||
fireEvent.click(screen.getByRole('button',{name:'连接'}));
|
||||
expect(await screen.findByText('/opt/unitree_rl_mjlab')).toBeInTheDocument();
|
||||
fireEvent.change(screen.getByLabelText('并行环境'),{target:{value:'32'}});
|
||||
fireEvent.click(screen.getByRole('button',{name:'发起本地训练'}));
|
||||
await waitFor(()=>expect(fetchMock).toHaveBeenCalledTimes(2));
|
||||
const request=fetchMock.mock.calls[1][1] as RequestInit;
|
||||
expect(JSON.parse(String(request.body))).toMatchObject({taskId:'Unitree-Go2-Flat',numEnvs:32,device:'gpu',gpuIds:[0],wandbMode:'offline'});
|
||||
expect(await screen.findByText('排队中')).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,77 @@
|
||||
import {useEffect,useState,type ReactNode} from 'react';
|
||||
import {Download,Link,Play,Server,Square} from 'lucide-react';
|
||||
import {Badge,Button,ProgressBar,PropertyRow,Select} from '../../components/ui';
|
||||
import {LocalTrainingClient} from '../../training/LocalTrainingClient';
|
||||
import type {TrainingDevice,TrainingJob,TrainingServerInfo,WandbMode} from '../../training/types';
|
||||
|
||||
const ENDPOINT_KEY='mujoco-local-training-endpoint',JOB_KEY='mujoco-local-training-job';
|
||||
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 errorText(error:unknown):string{return error instanceof Error?error.message:String(error);}
|
||||
function stateLabel(state:TrainingJob['state']):string{return {queued:'排队中',running:'训练中',succeeded:'已完成',failed:'失败',cancelled:'已取消'}[state];}
|
||||
|
||||
export function LocalTrainingPanel({onPolicyReady}:{onPolicyReady(file:File):void}){
|
||||
const [endpoint,setEndpoint]=useState(()=>stored(ENDPOINT_KEY,DEFAULT_ENDPOINT));
|
||||
const [server,setServer]=useState<TrainingServerInfo>();
|
||||
const [job,setJob]=useState<TrainingJob>();
|
||||
const [busy,setBusy]=useState(false),[error,setError]=useState<string>();
|
||||
const [taskId,setTaskId]=useState('Unitree-Go2-Flat'),[numEnvs,setNumEnvs]=useState(4096),[maxIterations,setMaxIterations]=useState(2000),[seed,setSeed]=useState(42),[runName,setRunName]=useState('web'),[device,setDevice]=useState<TrainingDevice>('gpu'),[gpuIds,setGpuIds]=useState('0'),[wandbMode,setWandbMode]=useState<WandbMode>('offline');
|
||||
|
||||
const connect=async()=>{
|
||||
setBusy(true);setError(undefined);
|
||||
try{
|
||||
const client=new LocalTrainingClient(endpoint),info=await client.health();
|
||||
setServer(info);try{localStorage.setItem(ENDPOINT_KEY,client.endpoint);}catch{/* 当前会话仍可连接 */}
|
||||
if(info.tasks.length&&!info.tasks.includes(taskId))setTaskId(info.tasks[0]);
|
||||
const remembered=info.activeJobId??stored(JOB_KEY);
|
||||
if(remembered){try{setJob(await client.job(remembered));}catch{try{localStorage.removeItem(JOB_KEY);}catch{/* ignore */}}}
|
||||
if(!info.ready)setError(info.error??'训练服务尚未就绪');
|
||||
}catch(value){setServer(undefined);setError(errorText(value));}
|
||||
finally{setBusy(false);}
|
||||
};
|
||||
|
||||
const jobId=job?.id,jobState=job?.state;
|
||||
useEffect(()=>{
|
||||
if(!jobId||!jobState||!ACTIVE_STATES.has(jobState))return;
|
||||
let disposed=false;
|
||||
const refresh=async()=>{try{const next=await new LocalTrainingClient(endpoint).job(jobId);if(!disposed)setJob(next);}catch(value){if(!disposed)setError(errorText(value));}};
|
||||
const timer=window.setInterval(()=>void refresh(),1500);return()=>{disposed=true;window.clearInterval(timer);};
|
||||
},[endpoint,jobId,jobState]);
|
||||
|
||||
const start=async()=>{
|
||||
setBusy(true);setError(undefined);
|
||||
try{
|
||||
const ids=device==='gpu'?gpuIds.split(/[\s,]+/).filter(Boolean).map(Number):[];
|
||||
if(ids.some(id=>!Number.isInteger(id)||id<0))throw new Error('GPU 编号必须是非负整数');
|
||||
const next=await new LocalTrainingClient(endpoint).start({taskId,numEnvs,maxIterations,seed,runName,device,gpuIds:ids,wandbMode});
|
||||
setJob(next);try{localStorage.setItem(JOB_KEY,next.id);}catch{/* ignore */}
|
||||
}catch(value){setError(errorText(value));}finally{setBusy(false);}
|
||||
};
|
||||
const cancel=async()=>{if(!job)return;setBusy(true);setError(undefined);try{setJob(await new LocalTrainingClient(endpoint).cancel(job.id));}catch(value){setError(errorText(value));}finally{setBusy(false);}};
|
||||
const importResult=async()=>{if(!job)return;setBusy(true);setError(undefined);try{onPolicyReady(await new LocalTrainingClient(endpoint).downloadPolicy(job.id));}catch(value){setError(errorText(value));}finally{setBusy(false);}};
|
||||
const active=Boolean(job&&ACTIVE_STATES.has(job.state));
|
||||
|
||||
return <div>
|
||||
<label className="block text-xs text-text-secondary"><span className="mb-1 block">本地训练服务</span><div className="flex gap-2"><input aria-label="本地训练服务地址" className="field h-7 min-w-0 flex-1 px-2 text-xs text-text-primary" value={endpoint} disabled={active} onChange={event=>setEndpoint(event.target.value)}/><Button icon={<Link className="h-3.5 w-3.5"/>} disabled={busy||active} onClick={()=>void connect()}>连接</Button></div></label>
|
||||
<div className="mt-2 flex items-center justify-between rounded-md border border-border bg-surface px-2 py-1.5 text-[10px] text-text-tertiary"><span className="flex min-w-0 items-center gap-1.5 truncate"><Server className="h-3.5 w-3.5"/>{server?.trainerRoot??'请先启动本地训练服务'}</span><Badge tone={server?.ready?'success':'warning'}>{server?.ready?'可用':'离线'}</Badge></div>
|
||||
{server?.ready&&!job&&<div className="mt-3 space-y-2">
|
||||
<Field label="训练任务"><Select aria-label="训练任务" className="w-full" value={taskId} onChange={event=>setTaskId(event.target.value)}>{server.tasks.map(task=><option key={task} value={task}>{task}</option>)}</Select></Field>
|
||||
<div className="grid grid-cols-2 gap-2"><NumberField label="并行环境" value={numEnvs} min={1} max={16384} onChange={setNumEnvs}/><NumberField label="训练迭代" value={maxIterations} min={1} max={1000000} onChange={setMaxIterations}/><NumberField label="随机种子" value={seed} min={0} max={2147483647} onChange={setSeed}/><Field label="运行名称"><input aria-label="运行名称" className="field h-7 w-full px-2 text-xs text-text-primary" value={runName} onChange={event=>setRunName(event.target.value)}/></Field></div>
|
||||
<div className="grid grid-cols-2 gap-2"><Field label="计算设备"><Select aria-label="计算设备" className="w-full" value={device} onChange={event=>setDevice(event.target.value as TrainingDevice)}><option value="gpu">GPU</option><option value="cpu">CPU</option></Select></Field><Field label="GPU 编号"><input aria-label="GPU 编号" className="field h-7 w-full px-2 text-xs text-text-primary disabled:opacity-40" value={gpuIds} disabled={device==='cpu'} onChange={event=>setGpuIds(event.target.value)}/></Field></div>
|
||||
<Field label="实验记录"><Select aria-label="W&B 模式" className="w-full" value={wandbMode} onChange={event=>setWandbMode(event.target.value as WandbMode)}><option value="offline">本地离线(默认,无需登录)</option><option value="disabled">完全禁用 W&B</option><option value="online">在线 W&B(需要 API Key)</option></Select></Field>
|
||||
<Button variant="primary" className="w-full" icon={<Play className="h-3.5 w-3.5"/>} disabled={busy} onClick={()=>void start()}>发起本地训练</Button>
|
||||
<p className="text-[10px] leading-4 text-text-tertiary">训练使用本地 mjlab 任务资产,不会把浏览器中的模型上传到网络。服务一次只运行一个训练任务。</p>
|
||||
</div>}
|
||||
{job&&<div className="mt-3 rounded-lg border border-border bg-surface p-2.5">
|
||||
<div className="mb-2 flex items-center justify-between gap-2"><span className="truncate text-xs font-medium text-text-primary" title={job.id}>{job.taskId}</span><Badge tone={job.state==='succeeded'?'success':job.state==='failed'||job.state==='cancelled'?'warning':'accent'}>{stateLabel(job.state)}</Badge></div>
|
||||
<ProgressBar value={job.progress} label="训练进度"/><div className="mt-2"><PropertyRow label="迭代" value={`${job.iteration} / ${job.maxIterations}`}/><PropertyRow label="状态" value={job.message}/></div>
|
||||
{job.logs.length>0&&<details className="mt-2"><summary className="cursor-pointer text-[10px] text-text-secondary">最近日志</summary><pre className="mt-1 max-h-36 overflow-auto whitespace-pre-wrap break-all rounded bg-app p-2 text-[9px] leading-4 text-text-tertiary">{job.logs.slice(-40).join('\n')}</pre></details>}
|
||||
<div className="mt-3 grid grid-cols-2 gap-2">{active?<Button variant="danger" className="col-span-2" icon={<Square className="h-3.5 w-3.5"/>} disabled={busy} onClick={()=>void cancel()}>停止训练</Button>:<><Button disabled={busy||!job.artifactReady} icon={<Download className="h-3.5 w-3.5"/>} onClick={()=>void importResult()}>导入策略</Button><Button onClick={()=>{setJob(undefined);try{localStorage.removeItem(JOB_KEY);}catch{/* ignore */}}}>新建任务</Button></>}</div>
|
||||
</div>}
|
||||
{error&&<p role="alert" className="mt-2 break-words rounded bg-danger/10 p-2 text-[10px] leading-4 text-danger">{error}</p>}
|
||||
</div>;
|
||||
}
|
||||
|
||||
function Field({label,children}:{label:string;children:ReactNode}){return <label className="block text-[10px] text-text-tertiary"><span className="mb-1 block">{label}</span>{children}</label>;}
|
||||
function NumberField({label,value,min,max,onChange}:{label:string;value:number;min:number;max:number;onChange(value:number):void}){return <Field label={label}><input aria-label={label} type="number" className="field h-7 w-full px-2 text-xs text-text-primary" value={value} min={min} max={max} onChange={event=>onChange(Number(event.target.value))}/></Field>;}
|
||||
@@ -13,6 +13,7 @@ import {TreeSearchField} from './TreeSearchField';
|
||||
import {ProjectBreadcrumb} from './ProjectBreadcrumb';
|
||||
import {PythonControllerPanel} from './PythonControllerPanel';
|
||||
import {RLPolicyPanel} from './RLPolicyPanel';
|
||||
import {LocalTrainingPanel} from './LocalTrainingPanel';
|
||||
|
||||
export function SidebarPanel({title,side,children,visible=true}:{title:string;side:'left'|'right';children:ReactNode;visible?:boolean}){return <ResizablePanel side={side} storageKey={`mujoco-${side}-sidebar-width`} visible={visible}><aside className={`flex h-full w-full min-w-0 flex-col overflow-hidden bg-panel ${side==='left'?'border-r':'border-l'} border-border`}><h2 className="flex h-10 shrink-0 items-center gap-2 border-b border-border bg-panel px-3 text-sm font-semibold text-text-primary"><Settings2 aria-hidden="true" className="h-4 w-4 text-accent"/>{title}</h2>{children}</aside></ResizablePanel>;}
|
||||
|
||||
@@ -33,7 +34,7 @@ export function ModelControlsSidebar(props:ModelControlsProps){const [tab,setTab
|
||||
const properties=<><CollapsibleSection title="模型信息" defaultOpen badge={<Badge>{s.model.nbody} Body</Badge>}><div><PropertyRow label="Body" value={s.model.nbody}/><PropertyRow label="Joint" value={s.model.njnt}/><PropertyRow label="Geom" value={s.model.ngeom}/><PropertyRow label="Actuator" value={s.model.nactuator}/><PropertyRow label="qpos / qvel" value={`${s.model.nq} / ${s.model.nv}`}/></div></CollapsibleSection>
|
||||
{props.selectedFormat==='urdf'&&<CollapsibleSection title="URDF 处理方式" defaultOpen={false}><Select aria-label="URDF 处理方式" className="w-full" value={props.urdfMode} disabled={props.loading} onChange={event=>props.onUrdfMode(event.target.value as UrdfLoadMode)}><option value="mjcf">转换为 MJCF(推荐)</option><option value="native">MuJoCo 原生 URDF</option></Select><label className="mt-3 block text-xs text-text-secondary"><span className="mb-1 block">基座类型</span><Select aria-label="URDF 基座类型" className="w-full" value={props.baseMode} disabled={props.loading||props.urdfMode==='native'} onChange={event=>props.onBaseMode(event.target.value as UrdfBaseMode)}><option value="floating">浮动基座(Free Joint)</option><option value="fixed">固定基座(连接世界)</option></Select></label><p className="mt-2 text-xs text-text-tertiary">MJCF 模式保留 visual mesh、添加物理地面,并将模型最低点对齐到 z=0。</p><Check label="显示碰撞几何" checked={props.showCollision} onChange={props.onShowCollision}/></CollapsibleSection>}
|
||||
<CollapsibleSection title="当前选择" defaultOpen>{props.selection?<div className="text-xs"><PropertyRow label="Body" value={props.selection.bodyName} action={<CopyButton value={props.selection.bodyName} label="复制 Body 名称"/>}/><PropertyRow label="标识" value={`${props.selection.bodyId} / ${props.selection.geomId} / ${props.selection.geomType}`} action={<CopyButton value={`body ${props.selection.bodyId}, geom ${props.selection.geomId}, type ${props.selection.geomType}`} label="复制标识"/>}/><PropertyRow label="位置" value={props.selection.position.map(value=>value.toFixed(3)).join(', ')} action={<CopyButton value={props.selection.position.join(', ')} label="复制位置"/>}/></div>:<p className="flex items-center gap-2 text-xs text-text-tertiary"><Info className="h-3.5 w-3.5"/>在视口中单击物体</p>}</CollapsibleSection></>;
|
||||
const controls=<><CollapsibleSection title="ONNX 强化学习策略" defaultOpen badge={s.rlPolicy?<Badge>{s.rlPolicy.enabled?'推理':'停止'}</Badge>:undefined}><RLPolicyPanel paths={props.policyPaths} selectedPath={props.selectedPolicyPath} status={props.policyStatus??s.rlPolicy} loading={props.loading} onSelectPath={props.onSelectPolicyPath} onLoadPath={props.onLoadPolicyPath} onImport={props.onImportPolicy} onToggle={props.onTogglePolicy} onCommand={props.onPolicyCommand} onRemove={props.onRemovePolicy}/></CollapsibleSection><CollapsibleSection title="Python 控制器" defaultOpen badge={s.controller?<Badge>{s.controller.enabled?'运行':'停止'}</Badge>:undefined}><PythonControllerPanel paths={props.controllerPaths} selectedPath={props.selectedControllerPath} status={props.controllerStatus??s.controller} loading={props.loading} onSelectPath={props.onSelectControllerPath} onLoadPath={props.onLoadControllerPath} onImport={props.onImportController} onToggle={props.onToggleController} onCommand={props.onControllerCommand} onRemove={props.onRemoveController}/></CollapsibleSection><CollapsibleSection title="Actuator" defaultOpen={false} badge={<Badge>{s.actuators.length}</Badge>}>{s.actuators.length?s.actuators.map(actuator=><ActuatorControl key={actuator.id} actuator={actuator} onControl={value=>props.onActuator(actuator.id,value)} onParameters={parameters=>props.onActuatorParameters(actuator.id,parameters)}/>):<p className="text-xs text-text-tertiary">模型没有驱动器</p>}</CollapsibleSection>
|
||||
const controls=<><CollapsibleSection title="ONNX 强化学习策略" defaultOpen badge={s.rlPolicy?<Badge>{s.rlPolicy.enabled?'推理':'停止'}</Badge>:undefined}><RLPolicyPanel paths={props.policyPaths} selectedPath={props.selectedPolicyPath} status={props.policyStatus??s.rlPolicy} loading={props.loading} onSelectPath={props.onSelectPolicyPath} onLoadPath={props.onLoadPolicyPath} onImport={props.onImportPolicy} onToggle={props.onTogglePolicy} onCommand={props.onPolicyCommand} onRemove={props.onRemovePolicy}/></CollapsibleSection><CollapsibleSection title="本地强化学习训练" defaultOpen={false}><LocalTrainingPanel onPolicyReady={props.onImportPolicy}/></CollapsibleSection><CollapsibleSection title="Python 控制器" defaultOpen badge={s.controller?<Badge>{s.controller.enabled?'运行':'停止'}</Badge>:undefined}><PythonControllerPanel paths={props.controllerPaths} selectedPath={props.selectedControllerPath} status={props.controllerStatus??s.controller} loading={props.loading} onSelectPath={props.onSelectControllerPath} onLoadPath={props.onLoadControllerPath} onImport={props.onImportController} onToggle={props.onToggleController} onCommand={props.onControllerCommand} onRemove={props.onRemoveController}/></CollapsibleSection><CollapsibleSection title="Actuator" defaultOpen={false} badge={<Badge>{s.actuators.length}</Badge>}>{s.actuators.length?s.actuators.map(actuator=><ActuatorControl key={actuator.id} actuator={actuator} onControl={value=>props.onActuator(actuator.id,value)} onParameters={parameters=>props.onActuatorParameters(actuator.id,parameters)}/>):<p className="text-xs text-text-tertiary">模型没有驱动器</p>}</CollapsibleSection>
|
||||
<CollapsibleSection title="关节" defaultOpen badge={<Badge>{s.joints.length}</Badge>}><div className="mb-4 grid grid-cols-2 gap-2"><Button onClick={props.onResetJoints}>重置关节</Button><Button variant={props.ignoreJointLimits?'primary':'secondary'} aria-pressed={props.ignoreJointLimits} onClick={props.onToggleJointLimits}>忽略关节限位</Button><Button variant={props.jointAdvanced?'primary':'secondary'} aria-pressed={props.jointAdvanced} onClick={props.onToggleAdvanced}>高级</Button><Button variant={props.angleUnit==='deg'?'primary':'secondary'} aria-pressed={props.angleUnit==='deg'} onClick={props.onToggleAngleUnit}>{props.angleUnit==='rad'?'rad 弧度制':'° 角度制'}</Button></div>{s.joints.map(joint=>{const scale=joint.type===3&&props.angleUnit==='deg'?180/Math.PI:1,unit=joint.type===3?(props.angleUnit==='deg'?'°':' rad'):joint.type===2?' m':'';return <ControlSlider key={joint.id} label={`${joint.name}${joint.editable?'':'(只读)'}`} value={joint.value*scale} min={joint.min*scale} max={joint.max*scale} unit={unit} advanced={props.jointAdvanced} limited={joint.limited} limitsIgnored={joint.limitsIgnored} limitMin={joint.limitMin*scale} limitMax={joint.limitMax*scale} disabled={!joint.editable} onChange={value=>props.onJoint(joint.id,value/scale)}/>;})}</CollapsibleSection>
|
||||
<CollapsibleSection title="外力强度" defaultOpen={false}><ControlSlider label={`${props.forceScale.toFixed(0)} N/屏幕单位`} value={props.forceScale} min={5} max={200} onChange={props.onForceScale}/><p className="text-xs text-text-tertiary">选择“外力施加”,在动态物体上按住拖动,松开即清零。</p></CollapsibleSection></>;
|
||||
return <SidebarPanel title="模型与控制" side="right" visible={props.visible}><Tabs label="模型控制侧栏" value={tab} onValueChange={setTab} items={[{value:'properties',label:'属性',icon:<Info className="h-3.5 w-3.5"/>,content:properties},{value:'controls',label:'控制',icon:<SlidersHorizontal className="h-3.5 w-3.5"/>,content:controls}]}/></SidebarPanel>;
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
import {afterEach,describe,expect,it,vi} from 'vitest';
|
||||
import {LocalTrainingClient} from './LocalTrainingClient';
|
||||
|
||||
afterEach(()=>vi.unstubAllGlobals());
|
||||
|
||||
describe('LocalTrainingClient',()=>{
|
||||
it('规范化服务地址并提交受类型约束的 JSON 请求',async()=>{
|
||||
const fetchMock=vi.fn().mockResolvedValue(new Response(JSON.stringify({id:'a'.repeat(32),state:'queued'}),{status:202,headers:{'Content-Type':'application/json'}}));
|
||||
vi.stubGlobal('fetch',fetchMock);
|
||||
const client=new LocalTrainingClient('http://127.0.0.1:8765/');
|
||||
await client.start({taskId:'Unitree-Go2-Flat',numEnvs:16,maxIterations:2,seed:42,runName:'test',device:'cpu',gpuIds:[],wandbMode:'offline'});
|
||||
expect(fetchMock).toHaveBeenCalledWith('http://127.0.0.1:8765/api/training/jobs',expect.objectContaining({method:'POST'}));
|
||||
const options=fetchMock.mock.calls[0][1] as RequestInit;
|
||||
expect(JSON.parse(String(options.body))).toMatchObject({taskId:'Unitree-Go2-Flat',numEnvs:16,device:'cpu'});
|
||||
});
|
||||
|
||||
it('显示服务端返回的中文错误',async()=>{
|
||||
vi.stubGlobal('fetch',vi.fn().mockResolvedValue(new Response(JSON.stringify({error:'已有训练任务正在运行'}),{status:409,headers:{'Content-Type':'application/json'}})));
|
||||
await expect(new LocalTrainingClient('http://localhost:8765').health()).rejects.toThrow('已有训练任务正在运行');
|
||||
});
|
||||
|
||||
it('拒绝非 HTTP 地址',()=>{expect(()=>new LocalTrainingClient('file:///tmp/socket')).toThrow('http 或 https');});
|
||||
});
|
||||
@@ -0,0 +1,36 @@
|
||||
import type {TrainingJob,TrainingRequest,TrainingServerInfo} from './types';
|
||||
|
||||
function normalizeEndpoint(value:string):string{
|
||||
const endpoint=value.trim().replace(/\/+$/,'');
|
||||
let url:URL;
|
||||
try{url=new URL(endpoint);}catch{throw new Error('训练服务地址无效');}
|
||||
if(url.protocol!=='http:'&&url.protocol!=='https:')throw new Error('训练服务地址必须使用 http 或 https');
|
||||
return url.toString().replace(/\/$/,'');
|
||||
}
|
||||
|
||||
async function responseError(response:Response):Promise<Error>{
|
||||
try{const body=await response.json() as {error?:string};if(body.error)return new Error(body.error);}catch{/* 使用 HTTP 状态作为回退 */}
|
||||
return new Error(`本地训练服务请求失败(HTTP ${response.status})`);
|
||||
}
|
||||
|
||||
export class LocalTrainingClient {
|
||||
readonly endpoint:string;
|
||||
constructor(endpoint:string){this.endpoint=normalizeEndpoint(endpoint);}
|
||||
|
||||
private async json<T>(path:string,init?:RequestInit):Promise<T>{
|
||||
const response=await fetch(`${this.endpoint}${path}`,init);
|
||||
if(!response.ok)throw await responseError(response);
|
||||
return response.json() as Promise<T>;
|
||||
}
|
||||
|
||||
health():Promise<TrainingServerInfo>{return this.json('/api/training/health');}
|
||||
start(request:TrainingRequest):Promise<TrainingJob>{return this.json('/api/training/jobs',{method:'POST',headers:{'Content-Type':'application/json'},body:JSON.stringify(request)});}
|
||||
job(id:string):Promise<TrainingJob>{return this.json(`/api/training/jobs/${encodeURIComponent(id)}`);}
|
||||
cancel(id:string):Promise<TrainingJob>{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`);
|
||||
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'});
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
export type TrainingJobState='queued'|'running'|'succeeded'|'failed'|'cancelled';
|
||||
export type TrainingDevice='cpu'|'gpu';
|
||||
export type WandbMode='offline'|'online'|'disabled';
|
||||
|
||||
export interface TrainingServerInfo {
|
||||
version:string;
|
||||
ready:boolean;
|
||||
trainerRoot:string;
|
||||
python:string;
|
||||
tasks:string[];
|
||||
activeJobId?:string;
|
||||
error?:string;
|
||||
}
|
||||
|
||||
export interface TrainingRequest {
|
||||
taskId:string;
|
||||
numEnvs:number;
|
||||
maxIterations:number;
|
||||
seed:number;
|
||||
runName:string;
|
||||
device:TrainingDevice;
|
||||
gpuIds:number[];
|
||||
wandbMode:WandbMode;
|
||||
}
|
||||
|
||||
export interface TrainingJob {
|
||||
id:string;
|
||||
state:TrainingJobState;
|
||||
taskId:string;
|
||||
createdAt:string;
|
||||
startedAt?:string;
|
||||
endedAt?:string;
|
||||
iteration:number;
|
||||
maxIterations:number;
|
||||
progress:number;
|
||||
message:string;
|
||||
logs:string[];
|
||||
artifactReady:boolean;
|
||||
artifactName?:string;
|
||||
}
|
||||
@@ -4,6 +4,7 @@ import type {MjvGeom, MjvOption, MjvCamera, MjvScene} from '@mujoco/mujoco';
|
||||
import type {FrameResult, SimulationSession, SimulationSnapshot} from '../simulation/SimulationSession';
|
||||
import {meshIdFromSceneDataId} from '../simulation/geometry';
|
||||
import {OrientationGizmo} from './OrientationGizmo';
|
||||
import {cameraAlignedForce,closestRayAxisParameter,resolveHingeDragDelta,signedAngleAroundAxis} from './interactionMath';
|
||||
|
||||
export type InteractionMode = 'select' | 'joint' | 'force';
|
||||
export type ViewerTheme='light'|'dark';
|
||||
@@ -31,7 +32,8 @@ export class MuJoCoViewer {
|
||||
private mjScene: MjvScene | null = null;
|
||||
private frame = 0; private lastFpsAt=performance.now(); private fpsFrames=0; private lastSnapshotAt=0;
|
||||
private meshes: THREE.Mesh[]=[]; private geometries=new Map<string,THREE.BufferGeometry>(); private textures=new Map<number,THREE.DataTexture>();
|
||||
private raycaster=new THREE.Raycaster(); private pointer=new THREE.Vector2(); private selected:THREE.Mesh|null=null; private dragStart:THREE.Vector2|null=null; private dragAxis=new THREE.Vector2(1,0); private dragJointId=-1; private dragJointValue=0; private arrow:THREE.ArrowHelper|null=null;private highlightedJointId=-1;private highlightedBodyId=-1;private jointMarker:THREE.Mesh|null=null;
|
||||
private readonly sensorRotation=new THREE.Matrix4();private readonly renderSize=new THREE.Vector2();
|
||||
private raycaster=new THREE.Raycaster(); private pointer=new THREE.Vector2(); private selected:THREE.Mesh|null=null; private dragStart:THREE.Vector2|null=null; private dragJointId=-1; private dragJointType=-1;private dragJointValue=0;private dragHitDistance=0;private dragSlideParameter=0;private dragJointPivot=new THREE.Vector3();private dragJointAxisWorld=new THREE.Vector3();private dragJointStartWorld=new THREE.Vector3();private dragJointStartPlaneVector=new THREE.Vector3(); private arrow:THREE.ArrowHelper|null=null;private highlightedJointId=-1;private highlightedBodyId=-1;private jointMarker:THREE.Mesh|null=null;
|
||||
private resizeObserver:ResizeObserver;
|
||||
private orientationGizmo:OrientationGizmo;
|
||||
private grid:THREE.GridHelper;
|
||||
@@ -43,7 +45,7 @@ export class MuJoCoViewer {
|
||||
this.camera.up.set(0,0,1); this.camera.position.set(3,-3,2); this.controls=new OrbitControls(this.camera,this.renderer.domElement); this.controls.enableDamping=true;
|
||||
this.scene.background=new THREE.Color(0x0b1220);this.hemisphere=new THREE.HemisphereLight(0xffffff,0x223344,1.3);this.scene.add(this.hemisphere); const light=new THREE.DirectionalLight(0xffffff,2); light.position.set(4,-3,7); light.castShadow=true; this.scene.add(light);this.grid=new THREE.GridHelper(20,40,0x3b82f6,0x253047).rotateX(Math.PI/2);this.scene.add(this.grid);this.orientationGizmo=new OrientationGizmo(host);this.orientationGizmo.update(this.camera);
|
||||
this.resizeObserver=new ResizeObserver(()=>this.resize()); this.resizeObserver.observe(host); this.resize();
|
||||
this.renderer.domElement.addEventListener('pointerdown',this.onPointerDown); this.renderer.domElement.addEventListener('pointermove',this.onPointerMove); window.addEventListener('pointerup',this.onPointerUp);
|
||||
this.renderer.domElement.addEventListener('pointerdown',this.onPointerDown);this.renderer.domElement.addEventListener('pointermove',this.onPointerMove);this.renderer.domElement.addEventListener('lostpointercapture',this.onPointerUp);window.addEventListener('pointerup',this.onPointerUp);window.addEventListener('pointercancel',this.onPointerUp);window.addEventListener('blur',this.onPointerUp);
|
||||
this.frame=requestAnimationFrame(this.animate);
|
||||
}
|
||||
|
||||
@@ -62,28 +64,31 @@ export class MuJoCoViewer {
|
||||
private fitCamera(session:SimulationSession):void {const {extent,center}=session.geometryBounds();this.controls.target.set(center[0],center[1],center[2]);this.camera.position.set(center[0]+extent*1.5,center[1]-extent*1.5,center[2]+extent);this.camera.near=Math.max(.001,extent/1000);this.camera.far=Math.max(100,extent*100);this.camera.updateProjectionMatrix();this.controls.update();}
|
||||
private resize():void {const w=Math.max(1,this.host.clientWidth),h=Math.max(1,this.host.clientHeight); this.renderer.setSize(w,h,false); this.camera.aspect=w/h; this.camera.updateProjectionMatrix();}
|
||||
|
||||
private animate=(now:number):void=>{try {const result=this.session?.advance(now)??{steps:0,stepMs:0,overBudget:false};this.updateThemeTransition(now); this.controls.update();this.orientationGizmo.update(this.camera); if(this.session){this.updateMuJoCoScene();this.updateJointMarker();} this.renderer.setScissorTest(false);this.renderer.render(this.scene,this.camera);if(this.showSensorCamera&&this.updateSensorCamera())this.renderSensorCamera(); this.fpsFrames++; let fps=0;if(now-this.lastFpsAt>=500){fps=this.fpsFrames*1000/(now-this.lastFpsAt);this.fpsFrames=0;this.lastFpsAt=now;} const snapshot=this.session&&now-this.lastSnapshotAt>150?(this.lastSnapshotAt=now,this.session.snapshot()):undefined; this.callbacks.onFrame(result,fps,snapshot);}catch(error){this.callbacks.onError(error instanceof Error?error:new Error(String(error)));} this.frame=requestAnimationFrame(this.animate);};
|
||||
private animate=(now:number):void=>{try {const result=this.session?.advance(now)??{steps:0,stepMs:0,overBudget:false};this.updateThemeTransition(now); this.controls.update();this.orientationGizmo.update(this.camera); if(this.session){this.updateMuJoCoScene();this.updateJointMarker();} this.renderer.setScissorTest(false);this.renderer.render(this.scene,this.camera);if(this.showSensorCamera&&this.updateSensorCamera())this.renderSensorCamera(); this.fpsFrames++; let fps=0;if(now-this.lastFpsAt>=500){fps=this.fpsFrames*1000/(now-this.lastFpsAt);this.fpsFrames=0;this.lastFpsAt=now;} const snapshot=this.session&&now-this.lastSnapshotAt>150?(this.lastSnapshotAt=now,this.session.snapshot()):undefined; /* 避免每个 RAF 都触发 Zustand/React 全树重渲染。 */ if(snapshot||fps>0)this.callbacks.onFrame(result,fps,snapshot);}catch(error){this.callbacks.onError(error instanceof Error?error:new Error(String(error)));} this.frame=requestAnimationFrame(this.animate);};
|
||||
|
||||
private updateSensorCamera():boolean {if(!this.session||this.sensorCameraId<0||this.sensorCameraId>=this.session.model.ncam)return false;const id=this.sensorCameraId,p=id*3,m=id*9,data=this.session.data,model=this.session.model;this.sensorCamera.position.set(Number(data.cam_xpos[p]),Number(data.cam_xpos[p+1]),Number(data.cam_xpos[p+2]));const rotation=new THREE.Matrix4().set(Number(data.cam_xmat[m]),Number(data.cam_xmat[m+1]),Number(data.cam_xmat[m+2]),0,Number(data.cam_xmat[m+3]),Number(data.cam_xmat[m+4]),Number(data.cam_xmat[m+5]),0,Number(data.cam_xmat[m+6]),Number(data.cam_xmat[m+7]),Number(data.cam_xmat[m+8]),0,0,0,0,1);this.sensorCamera.quaternion.setFromRotationMatrix(rotation);this.sensorCamera.fov=Number(model.cam_fovy[id])||45;const extent=this.session.geometryBounds().extent;this.sensorCamera.near=Math.max(.001,extent/1000);this.sensorCamera.far=Math.max(100,extent*100);this.sensorCamera.updateProjectionMatrix();return true;}
|
||||
private renderSensorCamera():void {const size=this.renderer.getSize(new THREE.Vector2()),width=Math.max(120,Math.min(320,size.x*.32)),height=width*9/16,margin=16;this.sensorCamera.aspect=width/height;this.sensorCamera.updateProjectionMatrix();this.renderer.setViewport(margin,margin,width,height);this.renderer.setScissor(margin,margin,width,height);this.renderer.setScissorTest(true);this.renderer.render(this.scene,this.sensorCamera);this.renderer.setScissorTest(false);this.renderer.setViewport(0,0,size.x,size.y);}
|
||||
private updateSensorCamera():boolean {if(!this.session||this.sensorCameraId<0||this.sensorCameraId>=this.session.model.ncam)return false;const id=this.sensorCameraId,p=id*3,m=id*9,data=this.session.data,model=this.session.model;this.sensorCamera.position.set(Number(data.cam_xpos[p]),Number(data.cam_xpos[p+1]),Number(data.cam_xpos[p+2]));this.sensorRotation.set(Number(data.cam_xmat[m]),Number(data.cam_xmat[m+1]),Number(data.cam_xmat[m+2]),0,Number(data.cam_xmat[m+3]),Number(data.cam_xmat[m+4]),Number(data.cam_xmat[m+5]),0,Number(data.cam_xmat[m+6]),Number(data.cam_xmat[m+7]),Number(data.cam_xmat[m+8]),0,0,0,0,1);this.sensorCamera.quaternion.setFromRotationMatrix(this.sensorRotation);this.sensorCamera.fov=Number(model.cam_fovy[id])||45;const extent=this.session.geometryBounds().extent;this.sensorCamera.near=Math.max(.001,extent/1000);this.sensorCamera.far=Math.max(100,extent*100);this.sensorCamera.updateProjectionMatrix();return true;}
|
||||
private renderSensorCamera():void {const size=this.renderer.getSize(this.renderSize),width=Math.max(120,Math.min(320,size.x*.32)),height=width*9/16,margin=16;this.sensorCamera.aspect=width/height;this.sensorCamera.updateProjectionMatrix();this.renderer.setViewport(margin,margin,width,height);this.renderer.setScissor(margin,margin,width,height);this.renderer.setScissorTest(true);this.renderer.render(this.scene,this.sensorCamera);this.renderer.setScissorTest(false);this.renderer.setViewport(0,0,size.x,size.y);}
|
||||
|
||||
private updateMuJoCoScene():void {const s=this.session!; s.module.mjv_updateScene(s.model,s.data,this.option!,s.perturb,this.mjCamera!,s.module.mjtCatBit.mjCAT_ALL.value,this.mjScene!); const geoms=this.mjScene!.geoms; try {for(let i=0;i<geoms.size();i++){const geom=geoms.get(i);if(!geom)continue;try{let mesh=this.meshes[i];const key=this.geometryKey(geom);if(!mesh||mesh.userData.geometryKey!==key){if(mesh){this.scene.remove(mesh);this.disposeMesh(mesh);} mesh=this.createMesh(geom,key);this.meshes[i]=mesh;this.scene.add(mesh);} mesh.visible=true;this.updateMesh(mesh,geom);}finally{geom.delete();}} for(let i=geoms.size();i<this.meshes.length;i++)this.meshes[i].visible=false;}finally{geoms.delete();}}
|
||||
private geometryKey(g:MjvGeom):string {const dataId=g.type===this.session!.module.mjtGeom.mjGEOM_MESH.value?meshIdFromSceneDataId(g.dataid):g.dataid;return `${g.type}:${dataId}:${Array.from(g.size).join(',')}`;}
|
||||
private updateMuJoCoScene():void {const s=this.session!; s.module.mjv_updateScene(s.model,s.data,this.option!,s.perturb,this.mjCamera!,s.module.mjtCatBit.mjCAT_ALL.value,this.mjScene!); const geoms=this.mjScene!.geoms; try {const count=geoms.size();for(let i=0;i<count;i++){const geom=geoms.get(i);if(!geom)continue;try{let mesh=this.meshes[i];const key=this.geometryKey(geom);if(!mesh||mesh.userData.geometryKey!==key){if(mesh){this.scene.remove(mesh);this.disposeMesh(mesh);} mesh=this.createMesh(geom,key);this.meshes[i]=mesh;this.scene.add(mesh);} mesh.visible=true;this.updateMesh(mesh,geom);}finally{geom.delete();}} for(let i=count;i<this.meshes.length;i++)this.meshes[i].visible=false;}finally{geoms.delete();}}
|
||||
/** 模型 geom 的尺寸固定,以 objid 保持稳定;接触点等动态 geom 才包含尺寸。 */
|
||||
private geometryKey(g:MjvGeom):string {const m=this.session!.module,dataId=g.type===m.mjtGeom.mjGEOM_MESH.value?meshIdFromSceneDataId(g.dataid):g.dataid;if(g.objtype===m.mjtObj.mjOBJ_GEOM.value)return `model:${g.objid}:${g.type}:${dataId}`;return `dynamic:${g.type}:${dataId}:${Array.from(g.size).join(',')}`;}
|
||||
private primitive(g:MjvGeom):THREE.BufferGeometry {const m=this.session!.module,t=g.type,s=g.size;if(t===m.mjtGeom.mjGEOM_PLANE.value)return new THREE.PlaneGeometry(2*(s[0]||1e3),2*(s[1]||1e3));if(t===m.mjtGeom.mjGEOM_SPHERE.value)return new THREE.SphereGeometry(s[0],24,16);if(t===m.mjtGeom.mjGEOM_CAPSULE.value)return new CapsuleGeometry(s[0],2*s[2]);if(t===m.mjtGeom.mjGEOM_BOX.value)return new THREE.BoxGeometry(2*s[0],2*s[1],2*s[2]);if(t===m.mjtGeom.mjGEOM_CYLINDER.value){const x=new THREE.CylinderGeometry(s[0],s[0],2*s[2],24);x.rotateX(Math.PI/2);return x;}if(t===m.mjtGeom.mjGEOM_ELLIPSOID.value){const x=new THREE.SphereGeometry(1,24,16);x.scale(s[0],s[1],s[2]);return x;}if(t===m.mjtGeom.mjGEOM_MESH.value&&g.dataid>=0)return this.meshGeometry(meshIdFromSceneDataId(g.dataid));return new THREE.BufferGeometry();}
|
||||
private meshGeometry(id:number):THREE.BufferGeometry {const m=this.session!.model;const va=Number(m.mesh_vertadr[id]),vn=Number(m.mesh_vertnum[id]),fa=Number(m.mesh_faceadr[id]),fn=Number(m.mesh_facenum[id]);const positions=new Float32Array(vn*3);for(let i=0;i<positions.length;i++)positions[i]=m.mesh_vert[va*3+i];const indices=new Uint32Array(fn*3);for(let i=0;i<indices.length;i++)indices[i]=m.mesh_face[fa*3+i];const geometry=new THREE.BufferGeometry();geometry.setAttribute('position',new THREE.BufferAttribute(positions,3));geometry.setIndex(new THREE.BufferAttribute(indices,1));const na=Number(m.mesh_normaladr[id]),nn=Number(m.mesh_normalnum[id]);if(nn===vn){const normals=new Float32Array(nn*3);for(let i=0;i<normals.length;i++)normals[i]=m.mesh_normal[na*3+i];geometry.setAttribute('normal',new THREE.BufferAttribute(normals,3));}else geometry.computeVertexNormals();const ta=Number(m.mesh_texcoordadr[id]),tn=Number(m.mesh_texcoordnum[id]);if(tn===vn&&ta>=0){const uv=new Float32Array(tn*2);for(let i=0;i<uv.length;i++)uv[i]=m.mesh_texcoord[ta*2+i];geometry.setAttribute('uv',new THREE.BufferAttribute(uv,2));}geometry.computeBoundingSphere();return geometry;}
|
||||
private texture(id:number):THREE.DataTexture|undefined {if(id<0)return;let found=this.textures.get(id);if(found)return found;const m=this.session!.model,w=Number(m.tex_width[id]),h=Number(m.tex_height[id]),channels=Number(m.tex_nchannel[id]||3),adr=Number(m.tex_adr[id]);if(!w||!h)return;const data=new Uint8Array(w*h*channels);for(let i=0;i<data.length;i++)data[i]=m.tex_data[adr+i];found=new THREE.DataTexture(data,w,h,channels===4?THREE.RGBAFormat:THREE.RGBFormat);found.colorSpace=THREE.SRGBColorSpace;found.flipY=true;found.needsUpdate=true;this.textures.set(id,found);return found;}
|
||||
private createMesh(g:MjvGeom,key:string):THREE.Mesh {let geometry=this.geometries.get(key);if(!geometry){geometry=this.primitive(g);this.geometries.set(key,geometry);}const map=this.texture(g.texid);const material=new THREE.MeshStandardMaterial({color:new THREE.Color(g.rgba[0],g.rgba[1],g.rgba[2]),opacity:g.rgba[3],transparent:g.rgba[3]<1,...(map?{map}:{}),roughness:Math.max(.05,1-g.shininess),metalness:g.reflectance});const mesh=new THREE.Mesh(geometry,material);mesh.matrixAutoUpdate=false;mesh.castShadow=true;mesh.receiveShadow=true;mesh.userData.geometryKey=key;return mesh;}
|
||||
// 只缓存真正可复用的模型 mesh。动态接触几何会逐帧改变尺寸,必须由所属
|
||||
// THREE.Mesh 在替换时释放,否则 geometry key 会无限增长并最终耗尽标签页内存。
|
||||
private createMesh(g:MjvGeom,key:string):THREE.Mesh {const m=this.session!.module,isMesh=g.type===m.mjtGeom.mjGEOM_MESH.value&&g.dataid>=0,sharedKey=isMesh?`mesh:${meshIdFromSceneDataId(g.dataid)}`:undefined;let geometry=sharedKey?this.geometries.get(sharedKey):undefined;if(!geometry){geometry=this.primitive(g);if(sharedKey)this.geometries.set(sharedKey,geometry);}const map=this.texture(g.texid);const material=new THREE.MeshStandardMaterial({color:new THREE.Color(g.rgba[0],g.rgba[1],g.rgba[2]),opacity:g.rgba[3],transparent:g.rgba[3]<1,...(map?{map}:{}),roughness:Math.max(.05,1-g.shininess),metalness:g.reflectance});const mesh=new THREE.Mesh(geometry,material);mesh.matrixAutoUpdate=false;mesh.castShadow=true;mesh.receiveShadow=true;mesh.userData.geometryKey=key;mesh.userData.ownsGeometry=!sharedKey;return mesh;}
|
||||
private updateMesh(mesh:THREE.Mesh,g:MjvGeom):void {const mat=mesh.material as THREE.MeshStandardMaterial;mat.color.setRGB(g.rgba[0],g.rgba[1],g.rgba[2]);mat.opacity=g.rgba[3];mat.transparent=g.rgba[3]<1;mesh.matrix.set(g.mat[0],g.mat[1],g.mat[2],g.pos[0],g.mat[3],g.mat[4],g.mat[5],g.pos[1],g.mat[6],g.mat[7],g.mat[8],g.pos[2],0,0,0,1);mesh.matrixWorldNeedsUpdate=true;const geomId=g.objtype===this.session!.module.mjtObj.mjOBJ_GEOM.value?g.objid:-1;const bodyId=geomId>=0?Number(this.session!.model.geom_bodyid[geomId]):-1;mesh.userData.geomId=geomId;mesh.userData.bodyId=bodyId;mesh.userData.geomType=g.type;this.applyMeshHighlight(mesh);}
|
||||
|
||||
private eventPointer(event:PointerEvent):void {const r=this.renderer.domElement.getBoundingClientRect();this.pointer.set((event.clientX-r.left)/r.width*2-1,-((event.clientY-r.top)/r.height)*2+1);}
|
||||
private onPointerDown=(event:PointerEvent):void=>{this.eventPointer(event);this.raycaster.setFromCamera(this.pointer,this.camera);const hit=this.raycaster.intersectObjects(this.meshes.filter(m=>m.visible),false)[0];if(!hit)return;const mesh=hit.object as THREE.Mesh;this.select(mesh);const bodyId=Number(mesh.userData.bodyId);if(this.mode==='joint'){const joint=this.session?.snapshot().joints.find(j=>j.bodyId===bodyId&&j.editable);if(joint){this.dragStart=this.pointer.clone();this.dragJointId=joint.id;this.dragJointValue=joint.value;const p=joint.bodyId*3,xm=joint.bodyId*9;const origin=new THREE.Vector3(this.session!.data.xpos[p],this.session!.data.xpos[p+1],this.session!.data.xpos[p+2]);const local=new THREE.Vector3(...joint.axis);const axis=new THREE.Vector3(this.session!.data.xmat[xm]*local.x+this.session!.data.xmat[xm+1]*local.y+this.session!.data.xmat[xm+2]*local.z,this.session!.data.xmat[xm+3]*local.x+this.session!.data.xmat[xm+4]*local.y+this.session!.data.xmat[xm+5]*local.z,this.session!.data.xmat[xm+6]*local.x+this.session!.data.xmat[xm+7]*local.y+this.session!.data.xmat[xm+8]*local.z);const a=origin.clone().project(this.camera),b=origin.clone().add(axis).project(this.camera);this.dragAxis.set(b.x-a.x,b.y-a.y);if(this.dragAxis.lengthSq()<1e-6)this.dragAxis.set(1,0);else this.dragAxis.normalize();}}else if(this.mode==='force'&&bodyId>0){this.dragStart=this.pointer.clone();this.session?.initializePerturb(this.mjScene!,bodyId);this.showArrow(hit.point);this.renderer.domElement.setPointerCapture(event.pointerId);}};
|
||||
private onPointerMove=(event:PointerEvent):void=>{if(!this.dragStart||!this.session)return;this.eventPointer(event);const dx=this.pointer.x-this.dragStart.x,dy=this.pointer.y-this.dragStart.y;if(this.mode==='joint'&&this.dragJointId>=0)this.session.setJointPosition(this.dragJointId,this.dragJointValue+(dx*this.dragAxis.x+dy*this.dragAxis.y)*Math.PI);else if(this.mode==='force'&&this.selected){const bodyId=Number(this.selected.userData.bodyId),scale=this.forceScale;const force:[number,number,number]=[dx*scale,0,-dy*scale];this.session.setExternalForce(bodyId,force);this.updateArrow(force);}};
|
||||
private onPointerDown=(event:PointerEvent):void=>{this.eventPointer(event);this.raycaster.setFromCamera(this.pointer,this.camera);const hit=this.raycaster.intersectObjects(this.meshes.filter(m=>m.visible),false)[0];if(!hit)return;const mesh=hit.object as THREE.Mesh;this.select(mesh);const bodyId=Number(mesh.userData.bodyId);if(this.mode==='joint'){const joint=this.session?.snapshot().joints.find(j=>j.bodyId===bodyId&&j.editable);if(joint){this.dragStart=this.pointer.clone();this.dragJointId=joint.id;this.dragJointType=joint.type;this.dragJointValue=joint.value;this.dragHitDistance=hit.distance;const offset=joint.id*3;this.dragJointPivot.set(Number(this.session!.data.xanchor[offset]),Number(this.session!.data.xanchor[offset+1]),Number(this.session!.data.xanchor[offset+2]));this.dragJointAxisWorld.set(Number(this.session!.data.xaxis[offset]),Number(this.session!.data.xaxis[offset+1]),Number(this.session!.data.xaxis[offset+2])).normalize();this.raycaster.ray.at(hit.distance,this.dragJointStartWorld);const plane=new THREE.Plane().setFromNormalAndCoplanarPoint(this.dragJointAxisWorld,this.dragJointPivot),projected=plane.projectPoint(this.dragJointStartWorld,new THREE.Vector3());this.dragJointStartPlaneVector.copy(projected).sub(this.dragJointPivot);this.dragSlideParameter=closestRayAxisParameter(this.raycaster.ray,this.dragJointPivot,this.dragJointAxisWorld);this.renderer.domElement.setPointerCapture(event.pointerId);}}else if(this.mode==='force'&&bodyId>0){this.dragStart=this.pointer.clone();this.session?.initializePerturb(this.mjScene!,bodyId);this.showArrow(hit.point);this.renderer.domElement.setPointerCapture(event.pointerId);}};
|
||||
private onPointerMove=(event:PointerEvent):void=>{if(!this.dragStart||!this.session)return;this.eventPointer(event);this.raycaster.setFromCamera(this.pointer,this.camera);const delta=this.pointer.clone().sub(this.dragStart),rect=this.renderer.domElement.getBoundingClientRect(),aspect=rect.width/Math.max(1,rect.height),screenDelta=new THREE.Vector2(delta.x*aspect,delta.y);if(this.mode==='joint'&&this.dragJointId>=0){let value=this.dragJointValue,currentPlaneVector:THREE.Vector3|undefined,currentSlideParameter=Number.NaN;if(this.dragJointType===2){currentSlideParameter=closestRayAxisParameter(this.raycaster.ray,this.dragJointPivot,this.dragJointAxisWorld);if(Number.isFinite(currentSlideParameter)&&Number.isFinite(this.dragSlideParameter))value+=currentSlideParameter-this.dragSlideParameter;}else if(this.dragJointType===3){const currentWorld=this.raycaster.ray.at(this.dragHitDistance,new THREE.Vector3()),plane=new THREE.Plane().setFromNormalAndCoplanarPoint(this.dragJointAxisWorld,this.dragJointPivot);currentPlaneVector=plane.projectPoint(currentWorld,new THREE.Vector3()).sub(this.dragJointPivot);const worldDelta=signedAngleAroundAxis(this.dragJointStartPlaneVector,currentPlaneVector,this.dragJointAxisWorld),cameraForward=this.camera.getWorldDirection(new THREE.Vector3()),planeFacing=Math.abs(this.raycaster.ray.direction.dot(this.dragJointAxisWorld));const tangentWorld=cameraForward.clone().cross(this.dragJointAxisWorld).normalize(),a=this.dragJointPivot.clone().project(this.camera),b=this.dragJointPivot.clone().add(tangentWorld).project(this.camera),tangentScreen=new THREE.Vector2((b.x-a.x)*aspect,b.y-a.y);const tangentDelta=tangentScreen.lengthSq()>1e-10?screenDelta.dot(tangentScreen.normalize())*Math.PI:0;value+=resolveHingeDragDelta(worldDelta,tangentDelta,planeFacing);}if(this.session.setJointPosition(this.dragJointId,value)){this.dragJointValue=value;this.dragStart.copy(this.pointer);if(currentPlaneVector&¤tPlaneVector.lengthSq()>1e-12)this.dragJointStartPlaneVector.copy(currentPlaneVector);if(Number.isFinite(currentSlideParameter))this.dragSlideParameter=currentSlideParameter;}}else if(this.mode==='force'&&this.selected){const bodyId=Number(this.selected.userData.bodyId),right=new THREE.Vector3(1,0,0).applyQuaternion(this.camera.quaternion),vector=cameraAlignedForce(screenDelta.x,screenDelta.y,right,this.forceScale),force:[number,number,number]=[vector.x,vector.y,vector.z];this.session.setExternalForce(bodyId,force);this.updateArrow(force);}};
|
||||
private onPointerUp=():void=>{this.stopDrag();};
|
||||
private stopDrag():void {this.dragStart=null;this.dragJointId=-1;this.session?.clearExternalForce();if(this.arrow){this.scene.remove(this.arrow);this.arrow.dispose();this.arrow=null;}}
|
||||
private stopDrag():void {this.dragStart=null;this.dragJointId=-1;this.dragJointType=-1;this.dragHitDistance=0;this.dragSlideParameter=0;this.session?.clearExternalForce();if(this.arrow){this.scene.remove(this.arrow);this.arrow.dispose();this.arrow=null;}}
|
||||
private select(mesh:THREE.Mesh):void {const previous=this.selected;this.selected=mesh;if(previous)this.applyMeshHighlight(previous);this.applyMeshHighlight(mesh);const bodyId=Number(mesh.userData.bodyId),geomId=Number(mesh.userData.geomId);const name=this.session?.snapshot().bodies.find(b=>b.id===bodyId)?.name??`body_${bodyId}`;const e=mesh.matrix.elements;this.callbacks.onSelection({bodyId,geomId,bodyName:name,geomType:Number(mesh.userData.geomType),position:[e[12],e[13],e[14]]});}
|
||||
private showArrow(origin:THREE.Vector3):void {this.arrow=new THREE.ArrowHelper(new THREE.Vector3(1,0,0),origin,0.01,0xf97316);this.scene.add(this.arrow);}
|
||||
private updateArrow(force:[number,number,number]):void {if(!this.arrow)return;const v=new THREE.Vector3(...force),length=v.length()/25;if(length>1e-6){this.arrow.setDirection(v.normalize());this.arrow.setLength(length,Math.min(.2,length*.2),Math.min(.1,length*.1));}}
|
||||
private disposeMesh(mesh:THREE.Mesh):void {(mesh.material as THREE.Material).dispose();}
|
||||
private disposeMesh(mesh:THREE.Mesh):void {(mesh.material as THREE.Material).dispose();if(mesh.userData.ownsGeometry)mesh.geometry.dispose();}
|
||||
private releaseModel():void {this.stopDrag();for(const mesh of this.meshes){this.scene.remove(mesh);this.disposeMesh(mesh);}this.meshes=[];for(const g of this.geometries.values())g.dispose();this.geometries.clear();for(const t of this.textures.values())t.dispose();this.textures.clear();this.mjScene?.delete();this.mjCamera?.delete();this.option?.delete();this.mjScene=null;this.mjCamera=null;this.option=null;this.session=null;this.sensorCameraId=-1;this.selected=null;this.highlightedJointId=-1;this.highlightedBodyId=-1;if(this.jointMarker){this.scene.remove(this.jointMarker);this.jointMarker.geometry.dispose();(this.jointMarker.material as THREE.Material).dispose();this.jointMarker=null;}}
|
||||
dispose():void {cancelAnimationFrame(this.frame);this.releaseModel();this.resizeObserver.disconnect();this.renderer.domElement.removeEventListener('pointerdown',this.onPointerDown);this.renderer.domElement.removeEventListener('pointermove',this.onPointerMove);window.removeEventListener('pointerup',this.onPointerUp);this.controls.dispose();this.orientationGizmo.dispose();this.grid.geometry.dispose();const gridMaterials=Array.isArray(this.grid.material)?this.grid.material:[this.grid.material];for(const material of gridMaterials)material.dispose();this.renderer.dispose();this.renderer.domElement.remove();}
|
||||
dispose():void {cancelAnimationFrame(this.frame);this.releaseModel();this.resizeObserver.disconnect();this.renderer.domElement.removeEventListener('pointerdown',this.onPointerDown);this.renderer.domElement.removeEventListener('pointermove',this.onPointerMove);this.renderer.domElement.removeEventListener('lostpointercapture',this.onPointerUp);window.removeEventListener('pointerup',this.onPointerUp);window.removeEventListener('pointercancel',this.onPointerUp);window.removeEventListener('blur',this.onPointerUp);this.controls.dispose();this.orientationGizmo.dispose();this.grid.geometry.dispose();const gridMaterials=Array.isArray(this.grid.material)?this.grid.material:[this.grid.material];for(const material of gridMaterials)material.dispose();this.renderer.dispose();this.renderer.domElement.remove();}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
import * as THREE from 'three';
|
||||
import {cameraAlignedForce,closestRayAxisParameter,resolveHingeDragDelta,signedAngleAroundAxis} from './interactionMath';
|
||||
|
||||
describe('interactionMath',()=>{
|
||||
it('按右手定则计算绕关节轴的旋转方向',()=>{
|
||||
expect(signedAngleAroundAxis(new THREE.Vector3(1,0,0),new THREE.Vector3(0,1,0),new THREE.Vector3(0,0,1))).toBeCloseTo(Math.PI/2);
|
||||
expect(signedAngleAroundAxis(new THREE.Vector3(1,0,0),new THREE.Vector3(0,-1,0),new THREE.Vector3(0,0,1))).toBeCloseTo(-Math.PI/2);
|
||||
});
|
||||
|
||||
it('侧视关节平面时采用切线方向',()=>{
|
||||
expect(resolveHingeDragDelta(-.4,.25,.05)).toBe(.25);
|
||||
expect(resolveHingeDragDelta(-.4,.25,.8)).toBe(-.4);
|
||||
expect(resolveHingeDragDelta(0,.25,.8)).toBe(.25);
|
||||
});
|
||||
|
||||
it('通过指针射线和关节轴最近点稳定求解 slide 位移',()=>{
|
||||
const axisOrigin=new THREE.Vector3(0,0,0),axis=new THREE.Vector3(1,0,0);
|
||||
expect(closestRayAxisParameter(new THREE.Ray(new THREE.Vector3(2,0,3),new THREE.Vector3(0,0,-1)),axisOrigin,axis)).toBeCloseTo(2);
|
||||
expect(closestRayAxisParameter(new THREE.Ray(new THREE.Vector3(-1,0,3),new THREE.Vector3(0,0,-1)),axisOrigin,axis)).toBeCloseTo(-1);
|
||||
expect(closestRayAxisParameter(new THREE.Ray(new THREE.Vector3(),new THREE.Vector3(1,0,0)),axisOrigin,axis)).toBeNaN();
|
||||
});
|
||||
|
||||
it('外力的右拖和上拖分别映射到相机右方与世界上方',()=>{
|
||||
expect(cameraAlignedForce(.5,.25,new THREE.Vector3(0,-1,.3),100).toArray()).toEqual([0,-50,25]);
|
||||
expect(cameraAlignedForce(.5,0,new THREE.Vector3(-1,0,0),100).x).toBe(-50);
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,29 @@
|
||||
import * as THREE from 'three';
|
||||
|
||||
/** 计算绕世界轴从 start 到 end 的有符号角度,遵循右手定则。 */
|
||||
export function signedAngleAroundAxis(start:THREE.Vector3,end:THREE.Vector3,axis:THREE.Vector3):number {
|
||||
if(start.lengthSq()<=1e-12||end.lengthSq()<=1e-12||axis.lengthSq()<=1e-12)return Number.NaN;
|
||||
const a=start.clone().normalize(),b=end.clone().normalize(),normal=axis.clone().normalize();
|
||||
return Math.atan2(normal.dot(a.clone().cross(b)),THREE.MathUtils.clamp(a.dot(b),-1,1));
|
||||
}
|
||||
|
||||
/** 关节旋转平面接近侧视时,使用相机切线拖动,避免投影退化和方向跳变。 */
|
||||
export function resolveHingeDragDelta(worldDelta:number,tangentDelta:number,planeFacingRatio:number,threshold=.2):number {
|
||||
const tangentValid=Number.isFinite(tangentDelta),worldValid=Number.isFinite(worldDelta)&&(Math.abs(worldDelta)>1e-8||!tangentValid||Math.abs(tangentDelta)<=1e-8);
|
||||
if(planeFacingRatio<threshold&&tangentValid)return tangentDelta;
|
||||
if(worldValid)return worldDelta;
|
||||
return tangentValid?tangentDelta:0;
|
||||
}
|
||||
|
||||
/** 求指针射线和世界关节轴两条直线的最近点参数,参数单位与 slide qpos 一致。 */
|
||||
export function closestRayAxisParameter(ray:THREE.Ray,axisOrigin:THREE.Vector3,axisDirection:THREE.Vector3):number {
|
||||
const direction=ray.direction.clone().normalize(),axis=axisDirection.clone().normalize();if(direction.lengthSq()<=1e-12||axis.lengthSq()<=1e-12)return Number.NaN;
|
||||
const offset=ray.origin.clone().sub(axisOrigin),dot=direction.dot(axis),denominator=1-dot*dot;if(denominator<=1e-8)return Number.NaN;
|
||||
const value=(axis.dot(offset)-dot*direction.dot(offset))/denominator;return Number.isFinite(value)?value:Number.NaN;
|
||||
}
|
||||
|
||||
/** MuJoCo MOVE_V 风格:水平跟随相机右方向,垂直始终对应世界 +Z。 */
|
||||
export function cameraAlignedForce(dx:number,dy:number,cameraRight:THREE.Vector3,scale:number):THREE.Vector3 {
|
||||
const right=cameraRight.clone();right.z=0;if(right.lengthSq()<=1e-12)right.set(1,0,0);else right.normalize();
|
||||
return right.multiplyScalar(dx*scale).addScaledVector(new THREE.Vector3(0,0,1),dy*scale);
|
||||
}
|
||||
Reference in New Issue
Block a user