chore(web-platform): release V0.6.1 工程质量优化
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-08-28 15:38:10 +08:00
parent f4b415c54f
commit 60d3a6d68c
135 changed files with 13552 additions and 2882 deletions
+232 -68
View File
@@ -1,79 +1,243 @@
import type {MjData,MjModel} from '@mujoco/mujoco';
import {buildGo2wObservation,GO2W_VELOCITY_TASK} from '../tasks/go2wVelocity';
import type {JointBinding,RLCommand} from '../types';
import type {PolicyRuntimeBindings} from './OnnxPolicyRuntime';
import type { MjData, MjModel } from '@mujoco/mujoco';
import { buildGo2wObservation, GO2W_VELOCITY_TASK } from '../tasks/go2wVelocity';
import type { JointBinding, RLCommand } from '../types';
import type { PolicyRuntimeBindings } from './OnnxPolicyRuntime';
interface BoundJoint extends JointBinding {positionActuator:boolean;controlScale:number;}
interface BoundJoint extends JointBinding {
positionActuator: boolean;
controlScale: number;
}
function rotateInverse(quaternion:readonly number[],vector:readonly number[]):[number,number,number]{
const [w,x,y,z]=quaternion,[vx,vy,vz]=vector;
const tx=2*(y*vz-z*vy),ty=2*(z*vx-x*vz),tz=2*(x*vy-y*vx);
return [vx-w*tx+(y*tz-z*ty),vy-w*ty+(z*tx-x*tz),vz-w*tz+(x*ty-y*tx)];
function rotateInverse(
quaternion: readonly number[],
vector: readonly number[],
): [number, number, number] {
const [w, x, y, z] = quaternion,
[vx, vy, vz] = vector;
const tx = 2 * (y * vz - z * vy),
ty = 2 * (z * vx - x * vz),
tz = 2 * (x * vy - y * vx);
return [
vx - w * tx + (y * tz - z * ty),
vy - w * ty + (z * tx - x * tz),
vz - w * tz + (x * ty - y * tx),
];
}
/** 将 mjlab Go2 velocity 的 47 维 actor 观测和 12 维关节位置动作映射到 MuJoCo。 */
export class Go2wPolicyBindings implements PolicyRuntimeBindings {
private readonly joints:BoundJoint[];
private readonly baseBodyId:number;
private readonly baseFreeJointId:number;
private readonly gyroSensorId?:number;
private readonly wheelActuatorIds:number[];
private readonly joints: BoundJoint[];
private readonly baseBodyId: number;
private readonly baseFreeJointId: number;
private readonly gyroSensorId?: number;
private readonly wheelActuatorIds: number[];
constructor(private readonly model:MjModel,private readonly data:MjData,private readonly setActuator:(id:number,value:number)=>void){
const jointIds=new Map<string,number>(),actuatorIds=new Map<string,number>(),sensorIds=new Map<string,number>(),bodyIds=new Map<string,number>();
for(let id=0;id<model.njnt;id+=1){const item=model.jnt(id);try{if(item.name)jointIds.set(item.name,id);}finally{item.delete();}}
for(let id=0;id<model.nactuator;id+=1){const item=model.actuator(id);try{if(item.name)actuatorIds.set(item.name,id);}finally{item.delete();}}
for(let id=0;id<model.nsensor;id+=1){const item=model.sensor(id);try{if(item.name)sensorIds.set(item.name,id);}finally{item.delete();}}
for(let id=0;id<model.nbody;id+=1){const item=model.body(id);try{if(item.name)bodyIds.set(item.name,id);}finally{item.delete();}}
this.baseBodyId=bodyIds.get('base_link')??bodyIds.get('base')??this.findFloatingBaseBody();
this.baseFreeJointId=this.findFreeJoint(this.baseBodyId);
const gyroCandidate=sensorIds.get('imu_gyro')??sensorIds.get('__platform_imu_gyro__');
this.gyroSensorId=gyroCandidate!==undefined&&this.isBaseAlignedGyro(gyroCandidate)?gyroCandidate:undefined;
this.joints=GO2W_VELOCITY_TASK.jointNames.map((name,index)=>{
const jointId=jointIds.get(name);if(jointId===undefined)throw new Error(`Go2-W 策略找不到关节:${name}`);
const short=name.replace(/_joint$/,'');
const actuatorId=actuatorIds.get(short)??actuatorIds.get(`${name}_motor`);
if(actuatorId===undefined)throw new Error(`Go2-W 策略找不到驱动器:${short} 或 ${name}_motor`);
const joint=model.jnt(jointId),actuator=model.actuator(actuatorId);
try{
const address=Number(model.actuator_ctrladr[actuatorId]??actuatorId),nextAddress=actuatorId+1<model.nactuator?Number(model.actuator_ctrladr[actuatorId+1]):model.nu;
if(nextAddress-address!==1||Number(actuator.trntype)!==0||Number(actuator.trnid[0])!==jointId)throw new Error(`驱动器 ${actuator.name||actuatorId} 不是关节 ${name} 的标量 joint transmission`);
if(Number(actuator.gaintype)!==0||Number(actuator.dyntype)!==0)throw new Error(`驱动器 ${actuator.name||actuatorId} 必须使用 fixed gain 和无激活动力学`);
const gear=Number(actuator.gear[0]),gain=Number(actuator.gainprm[0]),positionActuator=Number(actuator.biastype)===1&&Math.abs(Number(actuator.biasprm[1])+gain)<1e-5;
const motorActuator=Number(actuator.biastype)===0;
if(!positionActuator&&!motorActuator)throw new Error(`驱动器 ${actuator.name||actuatorId} 不是受支持的 motor/position 类型`);
if(positionActuator&&(Math.abs(gear-1)>1e-5||Math.abs(gain-GO2W_VELOCITY_TASK.stiffness[index])>1e-4||Math.abs(Number(actuator.biasprm[2])+GO2W_VELOCITY_TASK.damping[index])>1e-4))throw new Error(`position 驱动器 ${actuator.name||actuatorId} 的 gear/kp/kd 与 mjlab deploy 配置不一致`);
const controlScale=gear*gain;
if(!Number.isFinite(controlScale)||Math.abs(controlScale)<1e-9)throw new Error(`驱动器 ${actuator.name||actuatorId} 的 gear × gain 无效`);
return {name,jointId,qposAddress:Number(joint.qposadr),qvelAddress:Number(joint.dofadr),actuatorId,positionActuator,controlScale};
}finally{actuator.delete();joint.delete();}
});
this.wheelActuatorIds=['FL','FR','RL','RR'].flatMap(prefix=>{
const id=actuatorIds.get(`${prefix}_wheel`)??actuatorIds.get(`${prefix}_wheel_joint_motor`)??actuatorIds.get(`${prefix}_foot_joint_motor`);
return id===undefined?[]:[id];
});
}
observe(time:number,lastAction:Float32Array,command:RLCommand):Float32Array{
const quaternion=Array.from(this.data.xquat.subarray(this.baseBodyId*4,this.baseBodyId*4+4),Number);
const projectedGravity=rotateInverse(quaternion,[0,0,-1]);
let angularVelocity:[number,number,number];
if(this.gyroSensorId!==undefined){const address=Number(this.model.sensor_adr[this.gyroSensorId]);angularVelocity=[Number(this.data.sensordata[address]),Number(this.data.sensordata[address+1]),Number(this.data.sensordata[address+2])];}
else {const joint=this.model.jnt(this.baseFreeJointId);try{const address=Number(joint.dofadr)+3;angularVelocity=[Number(this.data.qvel[address]),Number(this.data.qvel[address+1]),Number(this.data.qvel[address+2])];}finally{joint.delete();}}
return buildGo2wObservation({angularVelocity,projectedGravity,command,time,jointPosition:this.joints.map(item=>Number(this.data.qpos[item.qposAddress])),jointVelocity:this.joints.map(item=>Number(this.data.qvel[item.qvelAddress])),lastAction:Array.from(lastAction)});
}
apply(action:Float32Array):void{
for(let index=0;index<this.joints.length;index+=1){
const item=this.joints[index],target=GO2W_VELOCITY_TASK.defaultJointPosition[index]+GO2W_VELOCITY_TASK.actionScale[index]*action[index];
const torque=GO2W_VELOCITY_TASK.stiffness[index]*(target-Number(this.data.qpos[item.qposAddress]))-GO2W_VELOCITY_TASK.damping[index]*Number(this.data.qvel[item.qvelAddress]);
this.setActuator(item.actuatorId,item.positionActuator?target:torque/item.controlScale);
constructor(
private readonly model: MjModel,
private readonly data: MjData,
private readonly setActuator: (id: number, value: number) => void,
) {
const jointIds = new Map<string, number>(),
actuatorIds = new Map<string, number>(),
sensorIds = new Map<string, number>(),
bodyIds = new Map<string, number>();
for (let id = 0; id < model.njnt; id += 1) {
const item = model.jnt(id);
try {
if (item.name) jointIds.set(item.name, id);
} finally {
item.delete();
}
}
for(const id of this.wheelActuatorIds)this.setActuator(id,0);
for (let id = 0; id < model.nactuator; id += 1) {
const item = model.actuator(id);
try {
if (item.name) actuatorIds.set(item.name, id);
} finally {
item.delete();
}
}
for (let id = 0; id < model.nsensor; id += 1) {
const item = model.sensor(id);
try {
if (item.name) sensorIds.set(item.name, id);
} finally {
item.delete();
}
}
for (let id = 0; id < model.nbody; id += 1) {
const item = model.body(id);
try {
if (item.name) bodyIds.set(item.name, id);
} finally {
item.delete();
}
}
this.baseBodyId =
bodyIds.get('base_link') ?? bodyIds.get('base') ?? this.findFloatingBaseBody();
this.baseFreeJointId = this.findFreeJoint(this.baseBodyId);
const gyroCandidate = sensorIds.get('imu_gyro') ?? sensorIds.get('__platform_imu_gyro__');
this.gyroSensorId =
gyroCandidate !== undefined && this.isBaseAlignedGyro(gyroCandidate)
? gyroCandidate
: undefined;
this.joints = GO2W_VELOCITY_TASK.jointNames.map((name, index) => {
const jointId = jointIds.get(name);
if (jointId === undefined) throw new Error(`Go2-W 策略找不到关节:${name}`);
const short = name.replace(/_joint$/, '');
const actuatorId = actuatorIds.get(short) ?? actuatorIds.get(`${name}_motor`);
if (actuatorId === undefined)
throw new Error(`Go2-W 策略找不到驱动器:${short} 或 ${name}_motor`);
const joint = model.jnt(jointId),
actuator = model.actuator(actuatorId);
try {
const address = Number(model.actuator_ctrladr[actuatorId] ?? actuatorId),
nextAddress =
actuatorId + 1 < model.nactuator
? Number(model.actuator_ctrladr[actuatorId + 1])
: model.nu;
if (
nextAddress - address !== 1 ||
Number(actuator.trntype) !== 0 ||
Number(actuator.trnid[0]) !== jointId
)
throw new Error(
`驱动器 ${actuator.name || actuatorId} 不是关节 ${name} 的标量 joint transmission`,
);
if (Number(actuator.gaintype) !== 0 || Number(actuator.dyntype) !== 0)
throw new Error(
`驱动器 ${actuator.name || actuatorId} 必须使用 fixed gain 和无激活动力学`,
);
const gear = Number(actuator.gear[0]),
gain = Number(actuator.gainprm[0]),
positionActuator =
Number(actuator.biastype) === 1 && Math.abs(Number(actuator.biasprm[1]) + gain) < 1e-5;
const motorActuator = Number(actuator.biastype) === 0;
if (!positionActuator && !motorActuator)
throw new Error(`驱动器 ${actuator.name || actuatorId} 不是受支持的 motor/position 类型`);
if (
positionActuator &&
(Math.abs(gear - 1) > 1e-5 ||
Math.abs(gain - GO2W_VELOCITY_TASK.stiffness[index]) > 1e-4 ||
Math.abs(Number(actuator.biasprm[2]) + GO2W_VELOCITY_TASK.damping[index]) > 1e-4)
)
throw new Error(
`position 驱动器 ${actuator.name || actuatorId} 的 gear/kp/kd 与 mjlab deploy 配置不一致`,
);
const controlScale = gear * gain;
if (!Number.isFinite(controlScale) || Math.abs(controlScale) < 1e-9)
throw new Error(`驱动器 ${actuator.name || actuatorId} 的 gear × gain 无效`);
return {
name,
jointId,
qposAddress: Number(joint.qposadr),
qvelAddress: Number(joint.dofadr),
actuatorId,
positionActuator,
controlScale,
};
} finally {
actuator.delete();
joint.delete();
}
});
this.wheelActuatorIds = ['FL', 'FR', 'RL', 'RR'].flatMap((prefix) => {
const id =
actuatorIds.get(`${prefix}_wheel`) ??
actuatorIds.get(`${prefix}_wheel_joint_motor`) ??
actuatorIds.get(`${prefix}_foot_joint_motor`);
return id === undefined ? [] : [id];
});
}
clear():void{for(const item of this.joints)this.setActuator(item.actuatorId,0);for(const id of this.wheelActuatorIds)this.setActuator(id,0);}
private isBaseAlignedGyro(sensorId:number):boolean{const siteId=Number(this.model.sensor_objid[sensorId]);if(Number(this.model.sensor_dim[sensorId])!==3||siteId<0||siteId>=this.model.nsite||Number(this.model.site_bodyid[siteId])!==this.baseBodyId)return false;const offset=siteId*4;return Math.abs(Number(this.model.site_quat[offset])-1)<1e-5&&Math.abs(Number(this.model.site_quat[offset+1]))<1e-5&&Math.abs(Number(this.model.site_quat[offset+2]))<1e-5&&Math.abs(Number(this.model.site_quat[offset+3]))<1e-5;}
private findFloatingBaseBody():number{for(let jointId=0;jointId<this.model.njnt;jointId+=1)if(Number(this.model.jnt_type[jointId])===0)return Number(this.model.jnt_bodyid[jointId]);throw new Error('Go2-W 策略需要浮动基座(free joint)');}
private findFreeJoint(bodyId:number):number{for(let jointId=0;jointId<this.model.njnt;jointId+=1)if(Number(this.model.jnt_type[jointId])===0&&Number(this.model.jnt_bodyid[jointId])===bodyId)return jointId;throw new Error('Go2-W 基座没有 free joint,请使用浮动基座模型');}
observe(time: number, lastAction: Float32Array, command: RLCommand): Float32Array {
const quaternion = Array.from(
this.data.xquat.subarray(this.baseBodyId * 4, this.baseBodyId * 4 + 4),
Number,
);
const projectedGravity = rotateInverse(quaternion, [0, 0, -1]);
let angularVelocity: [number, number, number];
if (this.gyroSensorId !== undefined) {
const address = Number(this.model.sensor_adr[this.gyroSensorId]);
angularVelocity = [
Number(this.data.sensordata[address]),
Number(this.data.sensordata[address + 1]),
Number(this.data.sensordata[address + 2]),
];
} else {
const joint = this.model.jnt(this.baseFreeJointId);
try {
const address = Number(joint.dofadr) + 3;
angularVelocity = [
Number(this.data.qvel[address]),
Number(this.data.qvel[address + 1]),
Number(this.data.qvel[address + 2]),
];
} finally {
joint.delete();
}
}
return buildGo2wObservation({
angularVelocity,
projectedGravity,
command,
time,
jointPosition: this.joints.map((item) => Number(this.data.qpos[item.qposAddress])),
jointVelocity: this.joints.map((item) => Number(this.data.qvel[item.qvelAddress])),
lastAction: Array.from(lastAction),
});
}
apply(action: Float32Array): void {
for (let index = 0; index < this.joints.length; index += 1) {
const item = this.joints[index],
target =
GO2W_VELOCITY_TASK.defaultJointPosition[index] +
GO2W_VELOCITY_TASK.actionScale[index] * action[index];
const torque =
GO2W_VELOCITY_TASK.stiffness[index] * (target - Number(this.data.qpos[item.qposAddress])) -
GO2W_VELOCITY_TASK.damping[index] * Number(this.data.qvel[item.qvelAddress]);
this.setActuator(
item.actuatorId,
item.positionActuator ? target : torque / item.controlScale,
);
}
for (const id of this.wheelActuatorIds) this.setActuator(id, 0);
}
clear(): void {
for (const item of this.joints) this.setActuator(item.actuatorId, 0);
for (const id of this.wheelActuatorIds) this.setActuator(id, 0);
}
private isBaseAlignedGyro(sensorId: number): boolean {
const siteId = Number(this.model.sensor_objid[sensorId]);
if (
Number(this.model.sensor_dim[sensorId]) !== 3 ||
siteId < 0 ||
siteId >= this.model.nsite ||
Number(this.model.site_bodyid[siteId]) !== this.baseBodyId
)
return false;
const offset = siteId * 4;
return (
Math.abs(Number(this.model.site_quat[offset]) - 1) < 1e-5 &&
Math.abs(Number(this.model.site_quat[offset + 1])) < 1e-5 &&
Math.abs(Number(this.model.site_quat[offset + 2])) < 1e-5 &&
Math.abs(Number(this.model.site_quat[offset + 3])) < 1e-5
);
}
private findFloatingBaseBody(): number {
for (let jointId = 0; jointId < this.model.njnt; jointId += 1)
if (Number(this.model.jnt_type[jointId]) === 0) return Number(this.model.jnt_bodyid[jointId]);
throw new Error('Go2-W 策略需要浮动基座(free joint)');
}
private findFreeJoint(bodyId: number): number {
for (let jointId = 0; jointId < this.model.njnt; jointId += 1)
if (
Number(this.model.jnt_type[jointId]) === 0 &&
Number(this.model.jnt_bodyid[jointId]) === bodyId
)
return jointId;
throw new Error('Go2-W 基座没有 free joint,请使用浮动基座模型');
}
}
+190 -62
View File
@@ -1,83 +1,211 @@
import * as ort from 'onnxruntime-web/wasm';
import {GO2W_VELOCITY_TASK,clampGo2wCommand} from '../tasks/go2wVelocity';
import type {RLCommand,RLPolicyStatus} from '../types';
import { GO2W_VELOCITY_TASK, clampGo2wCommand } from '../tasks/go2wVelocity';
import type { RLCommand, RLPolicyStatus } from '../types';
ort.env.wasm.numThreads=1;
ort.env.wasm.proxy=false;
ort.env.wasm.numThreads = 1;
ort.env.wasm.proxy = false;
export interface PolicyRuntimeBindings {
observe(time:number,lastAction:Float32Array,command:RLCommand):Float32Array;
apply(action:Float32Array):void;
clear():void;
observe(time: number, lastAction: Float32Array, command: RLCommand): Float32Array;
apply(action: Float32Array): void;
clear(): void;
}
function message(error:unknown):string{return error instanceof Error?error.message:String(error);}
function message(error: unknown): string {
return error instanceof Error ? error.message : String(error);
}
/**
* ONNX Runtime Web 的 run() 是异步 API。物理循环会在每个 mj_step 前持续施加最近一次
* 完成的动作,并按控制频率启动下一次推理,避免阻塞 MuJoCo 的同步步进循环。
*/
export class OnnxPolicyRuntime {
private enabled=false;
private disposed=false;
private inFlight=false;
private nextInferenceTime=0;
private action=new Float32Array(GO2W_VELOCITY_TASK.actionSize);
private commandValue:RLCommand={linearX:0,linearY:0,angularZ:0};
private inferenceCount=0;
private lastInferenceMs=0;
private error?:string;
private epoch=0;
private runPromise?:Promise<void>;
private enabled = false;
private disposed = false;
private inFlight = false;
private nextInferenceTime = 0;
private action = new Float32Array(GO2W_VELOCITY_TASK.actionSize);
private commandValue: RLCommand = { linearX: 0, linearY: 0, angularZ: 0 };
private inferenceCount = 0;
private lastInferenceMs = 0;
private error?: string;
private epoch = 0;
private runPromise?: Promise<void>;
private constructor(private readonly session:ort.InferenceSession,private readonly bindings:PolicyRuntimeBindings,private readonly path:string,private readonly inputName:string,private readonly outputName:string){}
private constructor(
private readonly session: ort.InferenceSession,
private readonly bindings: PolicyRuntimeBindings,
private readonly path: string,
private readonly inputName: string,
private readonly outputName: string,
) {}
static async load(model:Uint8Array,path:string,bindings:PolicyRuntimeBindings):Promise<OnnxPolicyRuntime>{
const session=await ort.InferenceSession.create(model.slice(),{executionProviders:['wasm'],graphOptimizationLevel:'all'});
try{
if(session.inputNames.length!==1)throw new Error(`当前仅支持单输入策略,模型包含 ${session.inputNames.length} 个输入`);
if(session.outputNames.length<1)throw new Error('ONNX 策略没有输出');
const input=session.inputMetadata[0],output=session.outputMetadata[0];
if(!input?.isTensor||input.type!=='float32')throw new Error('策略输入必须是 float32 Tensor');
if(!output?.isTensor||output.type!=='float32')throw new Error('策略输出必须是 float32 Tensor');
if(input.shape.length!==2||output.shape.length!==2)throw new Error(`策略输入/输出必须是二维 [batch, features],实际为 [${input.shape}] / [${output.shape}]`);
const inputBatch=input.shape[0],outputBatch=output.shape[0],fixedInput=input.shape[1],fixedOutput=output.shape[1];
if(typeof inputBatch==='number'&&inputBatch!==-1&&inputBatch!==1)throw new Error(`策略输入 batch 必须为 1 或动态维度,实际为 ${inputBatch}`);
if(typeof outputBatch==='number'&&outputBatch!==-1&&outputBatch!==1)throw new Error(`策略输出 batch 必须为 1 或动态维度,实际为 ${outputBatch}`);
if(typeof fixedInput==='number'&&fixedInput>0&&fixedInput!==GO2W_VELOCITY_TASK.observationSize)throw new Error(`策略观测维度不匹配:模型 ${fixedInput},任务 ${GO2W_VELOCITY_TASK.observationSize}`);
if(typeof fixedOutput==='number'&&fixedOutput>0&&fixedOutput!==GO2W_VELOCITY_TASK.actionSize)throw new Error(`策略动作维度不匹配:模型 ${fixedOutput},任务 ${GO2W_VELOCITY_TASK.actionSize}`);
return new OnnxPolicyRuntime(session,bindings,path,session.inputNames[0],session.outputNames[0]);
}catch(error){await session.release();throw error;}
static async load(
model: Uint8Array,
path: string,
bindings: PolicyRuntimeBindings,
): Promise<OnnxPolicyRuntime> {
const session = await ort.InferenceSession.create(model.slice(), {
executionProviders: ['wasm'],
graphOptimizationLevel: 'all',
});
try {
if (session.inputNames.length !== 1)
throw new Error(`当前仅支持单输入策略,模型包含 ${session.inputNames.length} 个输入`);
if (session.outputNames.length < 1) throw new Error('ONNX 策略没有输出');
const input = session.inputMetadata[0],
output = session.outputMetadata[0];
if (!input?.isTensor || input.type !== 'float32')
throw new Error('策略输入必须是 float32 Tensor');
if (!output?.isTensor || output.type !== 'float32')
throw new Error('策略输出必须是 float32 Tensor');
if (input.shape.length !== 2 || output.shape.length !== 2)
throw new Error(
`策略输入/输出必须是二维 [batch, features],实际为 [${input.shape}] / [${output.shape}]`,
);
const inputBatch = input.shape[0],
outputBatch = output.shape[0],
fixedInput = input.shape[1],
fixedOutput = output.shape[1];
if (typeof inputBatch === 'number' && inputBatch !== -1 && inputBatch !== 1)
throw new Error(`策略输入 batch 必须为 1 或动态维度,实际为 ${inputBatch}`);
if (typeof outputBatch === 'number' && outputBatch !== -1 && outputBatch !== 1)
throw new Error(`策略输出 batch 必须为 1 或动态维度,实际为 ${outputBatch}`);
if (
typeof fixedInput === 'number' &&
fixedInput > 0 &&
fixedInput !== GO2W_VELOCITY_TASK.observationSize
)
throw new Error(
`策略观测维度不匹配:模型 ${fixedInput},任务 ${GO2W_VELOCITY_TASK.observationSize}`,
);
if (
typeof fixedOutput === 'number' &&
fixedOutput > 0 &&
fixedOutput !== GO2W_VELOCITY_TASK.actionSize
)
throw new Error(
`策略动作维度不匹配:模型 ${fixedOutput},任务 ${GO2W_VELOCITY_TASK.actionSize}`,
);
return new OnnxPolicyRuntime(
session,
bindings,
path,
session.inputNames[0],
session.outputNames[0],
);
} catch (error) {
await session.release();
throw error;
}
}
status():RLPolicyStatus{return {taskId:GO2W_VELOCITY_TASK.id,taskName:GO2W_VELOCITY_TASK.name,path:this.path,loaded:!this.disposed,enabled:this.enabled,controlHz:GO2W_VELOCITY_TASK.controlHz,observationSize:GO2W_VELOCITY_TASK.observationSize,actionSize:GO2W_VELOCITY_TASK.actionSize,inputName:this.inputName,outputName:this.outputName,command:{...this.commandValue},inferenceCount:this.inferenceCount,lastInferenceMs:this.lastInferenceMs,error:this.error};}
setCommand(command:RLCommand):void{this.commandValue=clampGo2wCommand(command);}
setEnabled(enabled:boolean,time:number):void{if(this.disposed)return;this.epoch+=1;this.enabled=enabled;this.error=undefined;this.nextInferenceTime=time;if(!enabled){this.action.fill(0);this.bindings.clear();}}
reset(time:number):void{this.epoch+=1;this.action.fill(0);this.nextInferenceTime=time;this.error=undefined;this.bindings.clear();}
status(): RLPolicyStatus {
return {
taskId: GO2W_VELOCITY_TASK.id,
taskName: GO2W_VELOCITY_TASK.name,
path: this.path,
loaded: !this.disposed,
enabled: this.enabled,
controlHz: GO2W_VELOCITY_TASK.controlHz,
observationSize: GO2W_VELOCITY_TASK.observationSize,
actionSize: GO2W_VELOCITY_TASK.actionSize,
inputName: this.inputName,
outputName: this.outputName,
command: { ...this.commandValue },
inferenceCount: this.inferenceCount,
lastInferenceMs: this.lastInferenceMs,
error: this.error,
};
}
setCommand(command: RLCommand): void {
this.commandValue = clampGo2wCommand(command);
}
setEnabled(enabled: boolean, time: number): void {
if (this.disposed) return;
this.epoch += 1;
this.enabled = enabled;
this.error = undefined;
this.nextInferenceTime = time;
if (!enabled) {
this.action.fill(0);
this.bindings.clear();
}
}
reset(time: number): void {
this.epoch += 1;
this.action.fill(0);
this.nextInferenceTime = time;
this.error = undefined;
this.bindings.clear();
}
step(time:number):void{
if(!this.enabled||this.disposed)return;
step(time: number): void {
if (!this.enabled || this.disposed) return;
this.bindings.apply(this.action);
if(this.inFlight||time+1e-9<this.nextInferenceTime)return;
let observation:Float32Array;
try{observation=this.bindings.observe(time,this.action,this.commandValue);}
catch(error){this.fail(error);return;}
this.inFlight=true;
this.nextInferenceTime=time+1/GO2W_VELOCITY_TASK.controlHz;
const started=performance.now(),epoch=this.epoch;
const input=new ort.Tensor('float32',observation,[1,observation.length]);
this.runPromise=this.session.run({[this.inputName]:input}).then(outputs=>{
try{
const output=outputs[this.outputName];
if(!output||output.type!=='float32')throw new Error(`找不到 float32 输出:${this.outputName}`);
if(output.data.length!==GO2W_VELOCITY_TASK.actionSize)throw new Error(`策略动作维度错误:期望 ${GO2W_VELOCITY_TASK.actionSize},实际 ${output.data.length}`);
const next=Float32Array.from(output.data as Float32Array,Number);
for(const value of next)if(!Number.isFinite(value))throw new Error('策略输出包含非有限数');
if(!this.disposed&&this.enabled&&epoch===this.epoch){this.action=next;this.inferenceCount+=1;this.lastInferenceMs=performance.now()-started;}
}finally{for(const value of Object.values(outputs))value.dispose();}
}).catch(error=>{if(epoch===this.epoch)this.fail(error);}).finally(()=>{input.dispose();this.inFlight=false;this.runPromise=undefined;});
if (this.inFlight || time + 1e-9 < this.nextInferenceTime) return;
let observation: Float32Array;
try {
observation = this.bindings.observe(time, this.action, this.commandValue);
} catch (error) {
this.fail(error);
return;
}
this.inFlight = true;
this.nextInferenceTime = time + 1 / GO2W_VELOCITY_TASK.controlHz;
const started = performance.now(),
epoch = this.epoch;
const input = new ort.Tensor('float32', observation, [1, observation.length]);
this.runPromise = this.session
.run({ [this.inputName]: input })
.then((outputs) => {
try {
const output = outputs[this.outputName];
if (!output || output.type !== 'float32')
throw new Error(`找不到 float32 输出:${this.outputName}`);
if (output.data.length !== GO2W_VELOCITY_TASK.actionSize)
throw new Error(
`策略动作维度错误:期望 ${GO2W_VELOCITY_TASK.actionSize},实际 ${output.data.length}`,
);
const next = Float32Array.from(output.data as Float32Array, Number);
for (const value of next)
if (!Number.isFinite(value)) throw new Error('策略输出包含非有限数');
if (!this.disposed && this.enabled && epoch === this.epoch) {
this.action = next;
this.inferenceCount += 1;
this.lastInferenceMs = performance.now() - started;
}
} finally {
for (const value of Object.values(outputs)) value.dispose();
}
})
.catch((error) => {
if (epoch === this.epoch) this.fail(error);
})
.finally(() => {
input.dispose();
this.inFlight = false;
this.runPromise = undefined;
});
}
private fail(error:unknown):void{if(this.disposed)return;this.error=message(error);this.enabled=false;this.bindings.clear();}
dispose():void{if(this.disposed)return;this.disposed=true;this.enabled=false;this.epoch+=1;this.bindings.clear();const pending=this.runPromise??Promise.resolve();void pending.catch(()=>{}).finally(()=>this.session.release().catch(error=>console.warn('[ONNX] 释放推理会话失败',error)));}
private fail(error: unknown): void {
if (this.disposed) return;
this.error = message(error);
this.enabled = false;
this.bindings.clear();
}
dispose(): void {
if (this.disposed) return;
this.disposed = true;
this.enabled = false;
this.epoch += 1;
this.bindings.clear();
const pending = this.runPromise ?? Promise.resolve();
void pending
.catch(() => {})
.finally(() =>
this.session.release().catch((error) => console.warn('[ONNX] 释放推理会话失败', error)),
);
}
}
+47 -18
View File
@@ -1,28 +1,57 @@
import {describe,expect,it} from 'vitest';
import {buildGo2wObservation,clampGo2wCommand,go2wGaitPhase,GO2W_VELOCITY_TASK} from './go2wVelocity';
import { describe, expect, it } from 'vitest';
import {
buildGo2wObservation,
clampGo2wCommand,
go2wGaitPhase,
GO2W_VELOCITY_TASK,
} from './go2wVelocity';
describe('Go2-W velocity task',()=>{
it('按 mjlab deploy 顺序构造 47 维 actor 观测',()=>{
const jointPosition=GO2W_VELOCITY_TASK.defaultJointPosition.map(value=>value+0.1);
const observation=buildGo2wObservation({angularVelocity:[1,2,3],projectedGravity:[0,0,-1],command:{linearX:0.5,linearY:-0.25,angularZ:0.2},time:0,jointPosition,jointVelocity:Array(12).fill(0.3),lastAction:Array(12).fill(-0.4)});
describe('Go2-W velocity task', () => {
it('按 mjlab deploy 顺序构造 47 维 actor 观测', () => {
const jointPosition = GO2W_VELOCITY_TASK.defaultJointPosition.map((value) => value + 0.1);
const observation = buildGo2wObservation({
angularVelocity: [1, 2, 3],
projectedGravity: [0, 0, -1],
command: { linearX: 0.5, linearY: -0.25, angularZ: 0.2 },
time: 0,
jointPosition,
jointVelocity: Array(12).fill(0.3),
lastAction: Array(12).fill(-0.4),
});
expect(observation).toHaveLength(47);
[1,2,3,0,0,-1,0.5,-0.25,0.2,0,1].forEach((value,index)=>expect(observation[index]).toBeCloseTo(value));
for(const value of observation.slice(11,23))expect(value).toBeCloseTo(0.1);
for(const value of observation.slice(23,35))expect(value).toBeCloseTo(0.3);
for(const value of observation.slice(35,47))expect(value).toBeCloseTo(-0.4);
[1, 2, 3, 0, 0, -1, 0.5, -0.25, 0.2, 0, 1].forEach((value, index) =>
expect(observation[index]).toBeCloseTo(value),
);
for (const value of observation.slice(11, 23)) expect(value).toBeCloseTo(0.1);
for (const value of observation.slice(23, 35)) expect(value).toBeCloseTo(0.3);
for (const value of observation.slice(35, 47)) expect(value).toBeCloseTo(-0.4);
});
it('静止时关闭步态相位,并限制速度命令范围',()=>{
expect(go2wGaitPhase(0.15,{linearX:0,linearY:0,angularZ:0})).toEqual([0,0]);
const moving=go2wGaitPhase(0.15,{linearX:1,linearY:0,angularZ:0});
it('静止时关闭步态相位,并限制速度命令范围', () => {
expect(go2wGaitPhase(0.15, { linearX: 0, linearY: 0, angularZ: 0 })).toEqual([0, 0]);
const moving = go2wGaitPhase(0.15, { linearX: 1, linearY: 0, angularZ: 0 });
expect(moving[0]).toBeCloseTo(1);
expect(moving[1]).toBeCloseTo(0);
expect(clampGo2wCommand({linearX:4,linearY:-4,angularZ:3})).toEqual({linearX:1,linearY:-0.5,angularZ:1});
expect(clampGo2wCommand({ linearX: 4, linearY: -4, angularZ: 3 })).toEqual({
linearX: 1,
linearY: -0.5,
angularZ: 1,
});
});
it('拒绝维度错误或非有限观测',()=>{
const valid={angularVelocity:[0,0,0],projectedGravity:[0,0,-1],command:{linearX:0,linearY:0,angularZ:0},time:0,jointPosition:Array(12).fill(0),jointVelocity:Array(12).fill(0),lastAction:Array(12).fill(0)};
expect(()=>buildGo2wObservation({...valid,lastAction:[0]})).toThrow(/观测维度/);
expect(()=>buildGo2wObservation({...valid,angularVelocity:[Number.NaN,0,0]})).toThrow(/非有限数/);
it('拒绝维度错误或非有限观测', () => {
const valid = {
angularVelocity: [0, 0, 0],
projectedGravity: [0, 0, -1],
command: { linearX: 0, linearY: 0, angularZ: 0 },
time: 0,
jointPosition: Array(12).fill(0),
jointVelocity: Array(12).fill(0),
lastAction: Array(12).fill(0),
};
expect(() => buildGo2wObservation({ ...valid, lastAction: [0] })).toThrow(/观测维度/);
expect(() => buildGo2wObservation({ ...valid, angularVelocity: [Number.NaN, 0, 0] })).toThrow(
/非有限数/,
);
});
});
+68 -43
View File
@@ -1,57 +1,82 @@
import type {RLCommand} from '../types';
import type { RLCommand } from '../types';
export const GO2W_VELOCITY_TASK={
id:'unitree-go2w-velocity' as const,
name:'Unitree Go2-W 平衡/速度控制',
controlHz:50,
gaitPeriod:0.6,
observationSize:47,
actionSize:12,
commandLimits:{linearX:[-0.5,1] as const,linearY:[-0.5,0.5] as const,angularZ:[-1,1] as const},
jointNames:[
'FL_hip_joint','FL_thigh_joint','FL_calf_joint',
'FR_hip_joint','FR_thigh_joint','FR_calf_joint',
'RL_hip_joint','RL_thigh_joint','RL_calf_joint',
'RR_hip_joint','RR_thigh_joint','RR_calf_joint',
export const GO2W_VELOCITY_TASK = {
id: 'unitree-go2w-velocity' as const,
name: 'Unitree Go2-W 平衡/速度控制',
controlHz: 50,
gaitPeriod: 0.6,
observationSize: 47,
actionSize: 12,
commandLimits: {
linearX: [-0.5, 1] as const,
linearY: [-0.5, 0.5] as const,
angularZ: [-1, 1] as const,
},
jointNames: [
'FL_hip_joint',
'FL_thigh_joint',
'FL_calf_joint',
'FR_hip_joint',
'FR_thigh_joint',
'FR_calf_joint',
'RL_hip_joint',
'RL_thigh_joint',
'RL_calf_joint',
'RR_hip_joint',
'RR_thigh_joint',
'RR_calf_joint',
] as const,
defaultJointPosition:[-0.1,0.9,-1.8,0.1,0.9,-1.8,-0.1,0.9,-1.8,0.1,0.9,-1.8] as const,
actionScale:[0.25,0.25,0.25,0.25,0.25,0.25,0.25,0.25,0.25,0.25,0.25,0.25] as const,
stiffness:[20,20,40,20,20,40,20,20,40,20,20,40] as const,
damping:[1,1,2,1,1,2,1,1,2,1,1,2] as const,
defaultJointPosition: [-0.1, 0.9, -1.8, 0.1, 0.9, -1.8, -0.1, 0.9, -1.8, 0.1, 0.9, -1.8] as const,
actionScale: [0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25, 0.25] as const,
stiffness: [20, 20, 40, 20, 20, 40, 20, 20, 40, 20, 20, 40] as const,
damping: [1, 1, 2, 1, 1, 2, 1, 1, 2, 1, 1, 2] as const,
};
export function clampGo2wCommand(command:RLCommand):RLCommand {
const limits=GO2W_VELOCITY_TASK.commandLimits;
const clamp=(value:number,range:readonly[number,number])=>Math.min(range[1],Math.max(range[0],Number.isFinite(value)?value:0));
return {linearX:clamp(command.linearX,limits.linearX),linearY:clamp(command.linearY,limits.linearY),angularZ:clamp(command.angularZ,limits.angularZ)};
export function clampGo2wCommand(command: RLCommand): RLCommand {
const limits = GO2W_VELOCITY_TASK.commandLimits;
const clamp = (value: number, range: readonly [number, number]) =>
Math.min(range[1], Math.max(range[0], Number.isFinite(value) ? value : 0));
return {
linearX: clamp(command.linearX, limits.linearX),
linearY: clamp(command.linearY, limits.linearY),
angularZ: clamp(command.angularZ, limits.angularZ),
};
}
export function go2wGaitPhase(time:number,command:RLCommand):[number,number] {
if(Math.hypot(command.linearX,command.linearY,command.angularZ)<0.1)return [0,0];
const phase=((time/GO2W_VELOCITY_TASK.gaitPeriod)%1+1)%1;
return [Math.sin(phase*2*Math.PI),Math.cos(phase*2*Math.PI)];
export function go2wGaitPhase(time: number, command: RLCommand): [number, number] {
if (Math.hypot(command.linearX, command.linearY, command.angularZ) < 0.1) return [0, 0];
const phase = (((time / GO2W_VELOCITY_TASK.gaitPeriod) % 1) + 1) % 1;
return [Math.sin(phase * 2 * Math.PI), Math.cos(phase * 2 * Math.PI)];
}
export function buildGo2wObservation(values:{
angularVelocity:readonly number[];
projectedGravity:readonly number[];
command:RLCommand;
time:number;
jointPosition:readonly number[];
jointVelocity:readonly number[];
lastAction:readonly number[];
}):Float32Array {
const phase=go2wGaitPhase(values.time,values.command);
const observation=new Float32Array([
...values.angularVelocity.slice(0,3),
...values.projectedGravity.slice(0,3),
values.command.linearX,values.command.linearY,values.command.angularZ,
export function buildGo2wObservation(values: {
angularVelocity: readonly number[];
projectedGravity: readonly number[];
command: RLCommand;
time: number;
jointPosition: readonly number[];
jointVelocity: readonly number[];
lastAction: readonly number[];
}): Float32Array {
const phase = go2wGaitPhase(values.time, values.command);
const observation = new Float32Array([
...values.angularVelocity.slice(0, 3),
...values.projectedGravity.slice(0, 3),
values.command.linearX,
values.command.linearY,
values.command.angularZ,
...phase,
...values.jointPosition.map((value,index)=>value-GO2W_VELOCITY_TASK.defaultJointPosition[index]),
...values.jointPosition.map(
(value, index) => value - GO2W_VELOCITY_TASK.defaultJointPosition[index],
),
...values.jointVelocity,
...values.lastAction,
]);
if(observation.length!==GO2W_VELOCITY_TASK.observationSize)throw new Error(`Go2-W 观测维度错误:期望 ${GO2W_VELOCITY_TASK.observationSize},实际 ${observation.length}`);
for(const value of observation)if(!Number.isFinite(value))throw new Error('Go2-W 观测包含非有限数');
if (observation.length !== GO2W_VELOCITY_TASK.observationSize)
throw new Error(
`Go2-W 观测维度错误:期望 ${GO2W_VELOCITY_TASK.observationSize},实际 ${observation.length}`,
);
for (const value of observation)
if (!Number.isFinite(value)) throw new Error('Go2-W 观测包含非有限数');
return observation;
}
+22 -22
View File
@@ -1,30 +1,30 @@
export interface RLCommand {
linearX:number;
linearY:number;
angularZ:number;
linearX: number;
linearY: number;
angularZ: number;
}
export interface RLPolicyStatus {
taskId:'unitree-go2w-velocity';
taskName:string;
path:string;
loaded:boolean;
enabled:boolean;
controlHz:number;
observationSize:number;
actionSize:number;
inputName:string;
outputName:string;
command:RLCommand;
inferenceCount:number;
lastInferenceMs:number;
error?:string;
taskId: 'unitree-go2w-velocity';
taskName: string;
path: string;
loaded: boolean;
enabled: boolean;
controlHz: number;
observationSize: number;
actionSize: number;
inputName: string;
outputName: string;
command: RLCommand;
inferenceCount: number;
lastInferenceMs: number;
error?: string;
}
export interface JointBinding {
name:string;
jointId:number;
qposAddress:number;
qvelAddress:number;
actuatorId:number;
name: string;
jointId: number;
qposAddress: number;
qvelAddress: number;
actuatorId: number;
}