chore(web-platform): release V0.6.1 工程质量优化
This commit is contained in:
@@ -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,请使用浮动基座模型');
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)),
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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(
|
||||
/非有限数/,
|
||||
);
|
||||
});
|
||||
});
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user