v0.1.3
This commit is contained in:
@@ -1,7 +1,8 @@
|
||||
from robogauge.utils.task_register import task_register
|
||||
from robogauge.tasks.simulator.mujoco_config import MujocoConfig
|
||||
from robogauge.tasks.robots import RobotConfig
|
||||
from robogauge.tasks.robots import RobotConfig, Go2Config
|
||||
from robogauge.tasks.pipeline import BasePipeline
|
||||
from robogauge.tasks.gauge import BaseGaugeConfig
|
||||
|
||||
task_register.register('base', BasePipeline, MujocoConfig, BaseGaugeConfig, RobotConfig)
|
||||
task_register.register('go2', BasePipeline, MujocoConfig, BaseGaugeConfig, Go2Config)
|
||||
|
||||
@@ -9,6 +9,9 @@
|
||||
'''
|
||||
from robogauge.tasks.robots.base_robot_config import RobotConfig
|
||||
from robogauge.tasks.gauge.base_gauge_config import BaseGaugeConfig
|
||||
from robogauge.tasks.gauge.goal_data import GoalData, VelocityGoal, PositionGoal
|
||||
from robogauge.tasks.simulator.sim_data import SimData
|
||||
from robogauge.utils.logger import logger
|
||||
|
||||
class BaseGauge:
|
||||
def __init__(self, cfg: BaseGaugeConfig):
|
||||
@@ -20,10 +23,19 @@ class BaseGauge:
|
||||
def is_done(self) -> bool:
|
||||
return False
|
||||
|
||||
def get_goal(self) -> dict:
|
||||
goal = {}
|
||||
def get_goal(self) -> GoalData:
|
||||
goal = GoalData(
|
||||
goal_type='velocity',
|
||||
velocity_goal=VelocityGoal(
|
||||
lin_vel=[5.0, 0.0, 0.0],
|
||||
ang_vel=[0.0, 0.0, 0.0]
|
||||
)
|
||||
)
|
||||
return goal
|
||||
|
||||
def update_metrics(self, sim_info: dict):
|
||||
...
|
||||
def update_metrics(self, sim_data: SimData):
|
||||
if sim_data.n_step % int(0.1 / sim_data.sim_dt) != 0:
|
||||
return
|
||||
for i in range(len(sim_data.proprio.joint.force)):
|
||||
logger.log(sim_data.proprio.joint.force[i], f'dof/force_{i}', step=sim_data.n_step)
|
||||
|
||||
22
robogauge/tasks/gauge/goal_data.py
Normal file
22
robogauge/tasks/gauge/goal_data.py
Normal file
@@ -0,0 +1,22 @@
|
||||
from enum import Enum
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional, Literal
|
||||
|
||||
@dataclass
|
||||
class VelocityGoal:
|
||||
lin_vel: List[float] # x, y, z [m/s], z is ignored for ground robots
|
||||
ang_vel: List[float] # roll, pitch, yaw [rad/s], roll and pitch are ignored for ground robots
|
||||
|
||||
@dataclass
|
||||
class PositionGoal:
|
||||
# relative to robot's current position
|
||||
target_pos: List[float] # x, y, z [m], z is ignored for ground robots
|
||||
# reach target orientation
|
||||
target_quat: List[float] # x, y, z, w quaternion
|
||||
tolerance: float # [m] position tolerance to consider goal reached
|
||||
|
||||
@dataclass
|
||||
class GoalData:
|
||||
goal_type: Literal['velocity', 'position']
|
||||
velocity_goal: Optional[VelocityGoal] = None
|
||||
position_goal: Optional[PositionGoal] = None
|
||||
@@ -10,15 +10,17 @@
|
||||
import traceback
|
||||
from robogauge.utils.logger import logger
|
||||
from robogauge.tasks.simulator import MujocoSimulator, MujocoConfig
|
||||
from robogauge.tasks.robots import BaseRobot, RobotConfig
|
||||
from robogauge.tasks.robots import BaseRobot, RobotConfig, Go2Config, Go2
|
||||
from robogauge.tasks.gauge import BaseGauge, BaseGaugeConfig
|
||||
|
||||
class BasePipeline:
|
||||
def __init__(self,
|
||||
run_name: str,
|
||||
simulator_cfg: MujocoConfig,
|
||||
robot_cfg: RobotConfig,
|
||||
gauge_cfg: BaseGaugeConfig
|
||||
):
|
||||
self.run_name = run_name
|
||||
self.simulator_cfg = simulator_cfg
|
||||
self.robot_cfg = robot_cfg
|
||||
self.gauge_cfg = gauge_cfg
|
||||
@@ -28,31 +30,34 @@ class BasePipeline:
|
||||
self.gauge: BaseGauge = eval(gauge_cfg.gauge_class)(gauge_cfg)
|
||||
|
||||
def load(self):
|
||||
logger.create_tensorboard(self.run_name)
|
||||
self.sim.load(
|
||||
self.gauge_cfg.assets.terrain_xml,
|
||||
self.robot_cfg.assets.robot_xml,
|
||||
self.gauge_cfg.assets.terrain_spawn_xy,
|
||||
self.robot_cfg.assets.robot_spawn_height
|
||||
self.robot_cfg.assets.robot_spawn_height,
|
||||
self.robot_cfg.control.default_dof_pos
|
||||
)
|
||||
|
||||
def run(self):
|
||||
try:
|
||||
self.load()
|
||||
info = self.sim.step()
|
||||
sim_data = self.sim.step()
|
||||
frame_skip = int(self.robot_cfg.control.control_dt / self.simulator_cfg.physics.simulation_dt)
|
||||
logger.info(f"Sim FPS: {1.0 / self.simulator_cfg.physics.simulation_dt:.2f}, Control FPS: {1.0 / self.robot_cfg.control.control_dt:.2f}, Frame Skip: {frame_skip:d}")
|
||||
logger.info("Starting pipeline...")
|
||||
logger.info("Running pipeline...")
|
||||
while not self.gauge.is_done():
|
||||
goal = self.gauge.get_goal()
|
||||
obs = self.robot.build_observation(info, goal)
|
||||
action = self.robot.get_action(obs)
|
||||
obs = self.robot.build_observation(sim_data, goal)
|
||||
action, p_gains, d_gains, control_type = self.robot.get_action(obs)
|
||||
self.sim.setup_action(action, p_gains, d_gains, control_type)
|
||||
for _ in range(frame_skip):
|
||||
self.sim.apply_action(action)
|
||||
info = self.sim.step()
|
||||
self.gauge.update_metrics(info)
|
||||
sim_data = self.sim.step()
|
||||
self.gauge.update_metrics(sim_data)
|
||||
if self.gauge.is_reset():
|
||||
self.sim.reset()
|
||||
info = self.sim.step()
|
||||
sim_data = self.sim.step()
|
||||
finally:
|
||||
self.sim.close_viewer()
|
||||
logger.info("Pipeline execution finished.")
|
||||
logger.info(f"Logging saved at: {logger.log_dir}")
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from .base_robot_config import RobotConfig
|
||||
from .base_robot import BaseRobot
|
||||
# from .go2.go2_config import Go2Config
|
||||
# from .go2.go2_controller import Go2Controller
|
||||
from .go2.go2_config import Go2Config
|
||||
from .go2.go2 import Go2
|
||||
|
||||
@@ -9,21 +9,55 @@
|
||||
'''
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
from robogauge.utils.helpers import parse_path
|
||||
from robogauge.utils.logger import logger
|
||||
from robogauge.tasks.robots.base_robot_config import RobotConfig
|
||||
from robogauge.tasks.simulator.sim_data import SimData
|
||||
from robogauge.tasks.gauge.goal_data import GoalData
|
||||
|
||||
class BaseRobot:
|
||||
def __init__(self, cfg: RobotConfig):
|
||||
self.num_act = cfg.mdp.num_actions
|
||||
self.num_obs = cfg.mdp.num_observations
|
||||
self.model = None
|
||||
self.cfg = cfg
|
||||
self.device = self.cfg.control.device
|
||||
self.num_obs = cfg.control.num_observations
|
||||
self.num_action = cfg.control.num_actions
|
||||
self.control_type = cfg.control.control_type
|
||||
self.p_gains = np.array(cfg.control.p_gains)
|
||||
self.d_gains = np.array(cfg.control.d_gains)
|
||||
script_model_path = parse_path(cfg.control.torch_script_model_path)
|
||||
logger.info(f"Loading robot model from '{script_model_path}'")
|
||||
self.model = torch.jit.load(script_model_path).to(self.device)
|
||||
self.model.eval()
|
||||
|
||||
def load_model(self):
|
||||
...
|
||||
|
||||
def build_observation(self, sim_info: dict, goal_info: dict) -> np.ndarray:
|
||||
obs = np.zeros(self.num_obs)
|
||||
def build_observation(self, sim_data: SimData, goal_data: GoalData) -> np.ndarray:
|
||||
obs = np.zeros(self.num_obs, dtype=np.float32)
|
||||
return obs
|
||||
|
||||
def get_action(self, obs) -> np.ndarray:
|
||||
action = np.zeros_like(self.num_act)
|
||||
return action
|
||||
def get_action(self, obs: np.ndarray):
|
||||
"""
|
||||
Returns:
|
||||
action: (num_action,) target joint positions/velocities/torques
|
||||
p_gains: (num_action,) proportional gains for Mujoco PD controller
|
||||
d_gains: (num_action,) derivative gains for Mujoco PD controller
|
||||
control_type: 'P', 'V', or 'T' for position/velocity/torque control
|
||||
"""
|
||||
action = np.zeros(self.num_action, dtype=np.float32)
|
||||
return action, self.p_gains, self.d_gains, self.control_type
|
||||
|
||||
def get_projected_gravity(quat):
|
||||
""" Compute world frame gravity (0, 0, -1) projected into robot base frame.
|
||||
Args:
|
||||
quat: (4,) quaternion (w, x, y, z) from robot base to world frame
|
||||
Returns:
|
||||
projected_gravity: (3,) projected gravity vector in robot base frame
|
||||
"""
|
||||
qw, qx, qy, qz = quat
|
||||
|
||||
gravity_orientation = np.zeros(3)
|
||||
|
||||
gravity_orientation[0] = 2 * (-qz * qx + qw * qy)
|
||||
gravity_orientation[1] = -2 * (qz * qy + qw * qx)
|
||||
gravity_orientation[2] = 1 - 2 * (qw * qw + qz * qz)
|
||||
|
||||
return gravity_orientation
|
||||
|
||||
@@ -7,6 +7,7 @@
|
||||
@Blog : https://wty-yy.github.io/
|
||||
@Desc : Base Robot Configuration
|
||||
'''
|
||||
from typing_extensions import Literal
|
||||
from robogauge.utils.config import Config
|
||||
|
||||
class RobotConfig(Config):
|
||||
@@ -17,16 +18,32 @@ class RobotConfig(Config):
|
||||
robot_spawn_height = 0.1 # z [m]
|
||||
|
||||
class control:
|
||||
torch_script_model_path = "{ROBOGAUGE_ROOT_DIR}/resources/models/go2/go2_cts_61500.pt"
|
||||
device = 'cpu'
|
||||
torch_script_model_path = "{ROBOGAUGE_ROOT_DIR}/resources/models/go2/go2_cts_83501.pt"
|
||||
control_dt = 0.02 # 50 Hz
|
||||
action_scale = 0.25 # target pos = action_scale * action * default_pos
|
||||
stiffness = 20.0 # [N*m/rad]
|
||||
damping = 0.5 # [N*m*s/rad]
|
||||
control_type = 'P' # Position control
|
||||
|
||||
# Mujoco joint PD gains
|
||||
p_gains = [20.0, 20.0, 20.0, 20.0, 20.0, 20.0, 20.0, 20.0, 20.0, 20.0, 20.0, 20.0] # [N*m/rad]
|
||||
d_gains = [0.5, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5] # [N*m*s/rad]
|
||||
|
||||
class mdp:
|
||||
num_observations = 46
|
||||
num_observations = 45
|
||||
num_actions = 12
|
||||
|
||||
|
||||
max_velocity_cmd = [1.5, 1.0, 2.0]
|
||||
default_dof_pos = [0.1, 0.8, -1.5, -0.1, 0.8, -1.5,
|
||||
0.1, 1.0, -1.5, -0.1, 1.0, -1.5]
|
||||
|
||||
mj2model_dof_indices = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11]
|
||||
|
||||
class scales:
|
||||
lin_vel = 2.0
|
||||
ang_vel = 0.25
|
||||
dof_pos = 1.0
|
||||
dof_vel = 0.05
|
||||
action = 0.25 # target pos = action_scale * action * default_pos
|
||||
cmd = [2.0, 2.0, 0.25]
|
||||
|
||||
class commands:
|
||||
lin_vel_x = [-1, 1] # min max [m/s]
|
||||
lin_vel_y = [-1, 1] # min max [m/s]
|
||||
|
||||
@@ -0,0 +1,57 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
'''
|
||||
@File : go2.py
|
||||
@Time : 2025/11/28 15:27:07
|
||||
@Author : wty-yy
|
||||
@Version : 1.0
|
||||
@Blog : https://wty-yy.github.io/
|
||||
@Desc : None
|
||||
'''
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
from robogauge.tasks.robots.base_robot import BaseRobot, get_projected_gravity
|
||||
from robogauge.tasks.robots.go2.go2_config import Go2Config
|
||||
from robogauge.tasks.simulator.sim_data import SimData
|
||||
from robogauge.tasks.gauge.goal_data import GoalData
|
||||
from robogauge.utils.logger import logger
|
||||
|
||||
class Go2(BaseRobot):
|
||||
def __init__(self, cfg: Go2Config):
|
||||
super().__init__(cfg)
|
||||
self.max_velocity_cmd = np.array(cfg.control.max_velocity_cmd, dtype=np.float32)
|
||||
self.default_dof_pos = np.array(cfg.control.default_dof_pos, dtype=np.float32)
|
||||
self.last_action = np.zeros(self.num_action, dtype=np.float32)
|
||||
self.action_scale = cfg.control.scales.action
|
||||
self.mj2model_idx = self.cfg.control.mj2model_dof_indices
|
||||
self.model2mj_idx = [self.mj2model_idx.index(i) for i in range(len(self.mj2model_idx))]
|
||||
|
||||
def build_observation(self, sim_data: SimData, goal_data: GoalData) -> np.ndarray:
|
||||
sim_proprio = sim_data.proprio
|
||||
obs = np.zeros(self.num_obs)
|
||||
if goal_data.goal_type == 'velocity':
|
||||
ang_vel = sim_proprio.imu.ang_vel * self.cfg.control.scales.ang_vel
|
||||
projected_gravity = get_projected_gravity(sim_proprio.imu.quat)
|
||||
dof_pos = (sim_proprio.joint.pos - self.default_dof_pos) * self.cfg.control.scales.dof_pos
|
||||
dof_vel = sim_proprio.joint.vel * self.cfg.control.scales.dof_vel
|
||||
|
||||
cmd = np.array(goal_data.velocity_goal.lin_vel[:2] + goal_data.velocity_goal.ang_vel[2:3], np.float32)
|
||||
cmd = np.minimum(np.maximum(cmd, -self.max_velocity_cmd), self.max_velocity_cmd)
|
||||
cmd *= self.cfg.control.scales.cmd
|
||||
|
||||
obs[:3] = ang_vel
|
||||
obs[3:6] = projected_gravity
|
||||
obs[6:9] = cmd
|
||||
obs[9:9+self.num_action] = dof_pos[self.mj2model_idx]
|
||||
obs[9+self.num_action:9+2*self.num_action] = dof_vel[self.mj2model_idx]
|
||||
obs[9+2*self.num_action:9+3*self.num_action] = self.last_action[self.mj2model_idx]
|
||||
else:
|
||||
raise NotImplementedError(f"Goal type '{goal_data.goal_type}' not implemented in Go2 robot.")
|
||||
return obs
|
||||
|
||||
def get_action(self, obs: np.ndarray):
|
||||
obs_tensor = torch.tensor(obs, dtype=torch.float32).unsqueeze(0).to(self.device)
|
||||
action = self.model(obs_tensor).detach().cpu().numpy().squeeze(0)[self.model2mj_idx]
|
||||
self.last_action = action
|
||||
target_dof_pos = action * self.action_scale + self.default_dof_pos
|
||||
return target_dof_pos, self.p_gains, self.d_gains, self.control_type
|
||||
|
||||
@@ -7,14 +7,44 @@
|
||||
@Blog : https://wty-yy.github.io/
|
||||
@Desc : Go2 Robot Configuration
|
||||
'''
|
||||
from typing_extensions import Literal
|
||||
from robogauge.tasks.robots import RobotConfig
|
||||
|
||||
class Go2Config(RobotConfig):
|
||||
robot_class = 'Go2'
|
||||
|
||||
class assets:
|
||||
robot_xml = "{ROBOGAUGE_ROOT_DIR}/resources/robots/go2/go2.xml"
|
||||
robot_spawn_height = 0.1 # z [m]
|
||||
|
||||
class control:
|
||||
class control(RobotConfig.control):
|
||||
device = 'cpu'
|
||||
torch_script_model_path = "{ROBOGAUGE_ROOT_DIR}/resources/models/go2/go2_cts_83501.pt"
|
||||
control_dt = 0.02 # 50 Hz
|
||||
action_scale = 0.25 # scale for normalized actions
|
||||
control_type = 'P' # Position control
|
||||
|
||||
# Mujoco joint PD gains
|
||||
p_gains = [20.0, 20.0, 20.0, 20.0, 20.0, 20.0, 20.0, 20.0, 20.0, 20.0, 20.0, 20.0] # [N*m/rad]
|
||||
d_gains = [0.5, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5] # [N*m*s/rad]
|
||||
|
||||
num_observations = 45
|
||||
num_actions = 12
|
||||
|
||||
max_velocity_cmd = [1.5, 1.0, 2.0]
|
||||
default_dof_pos = [0.1, 0.8, -1.5, -0.1, 0.8, -1.5,
|
||||
0.1, 1.0, -1.5, -0.1, 1.0, -1.5]
|
||||
|
||||
mj2model_dof_indices = [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11]
|
||||
|
||||
class scales(RobotConfig.control.scales):
|
||||
lin_vel = 2.0
|
||||
ang_vel = 0.25
|
||||
dof_pos = 1.0
|
||||
dof_vel = 0.05
|
||||
action = 0.25 # target pos = action_scale * action * default_pos
|
||||
cmd = [2.0, 2.0, 0.25]
|
||||
|
||||
class commands(RobotConfig.commands):
|
||||
lin_vel_x = [-1, 1] # min max [m/s]
|
||||
lin_vel_y = [-1, 1] # min max [m/s]
|
||||
ang_vel_yaw = [-1, 1] # min max [rad/s]
|
||||
|
||||
@@ -13,7 +13,7 @@ class MujocoConfig(Config):
|
||||
simulator_class = 'MujocoSimulator'
|
||||
|
||||
class physics:
|
||||
simulation_dt = 0.005 # 200 Hz
|
||||
simulation_dt = 0.002 # 500 Hz
|
||||
|
||||
class viewer:
|
||||
headless = False
|
||||
@@ -21,5 +21,6 @@ class MujocoConfig(Config):
|
||||
|
||||
class render:
|
||||
save_video = False
|
||||
video_fps = 30
|
||||
height = 480
|
||||
width = 640
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
@Author : wty-yy
|
||||
@Version : 1.0
|
||||
@Blog : https://wty-yy.github.io/
|
||||
@Desc : None
|
||||
@Desc : Mujoco Simulator for Robogauge
|
||||
'''
|
||||
import mujoco
|
||||
import mujoco.viewer
|
||||
@@ -15,11 +15,15 @@ import re
|
||||
import time
|
||||
import imageio
|
||||
import numpy as np
|
||||
from typing import Literal
|
||||
|
||||
from robogauge.utils.logger import logger
|
||||
from robogauge.utils.helpers import pares_path
|
||||
from robogauge.utils.helpers import parse_path
|
||||
from robogauge.tasks.simulator.mujoco_config import MujocoConfig
|
||||
from robogauge.tasks.simulator.sim_data import RobotProprioception, JointState, BaseState, IMUState
|
||||
from robogauge.tasks.simulator.sim_data import (
|
||||
SimData,
|
||||
RobotProprioception, JointState, BaseState, IMUState
|
||||
)
|
||||
|
||||
class MujocoSimulator:
|
||||
def __init__(self, sim_cfg: MujocoConfig):
|
||||
@@ -28,12 +32,14 @@ class MujocoSimulator:
|
||||
self.robot_xml = None
|
||||
self.terrain_spawn_xy = None
|
||||
self.robot_spawn_height = None
|
||||
self.default_dof_pos = None
|
||||
self.viewer = None
|
||||
self.renderer = None
|
||||
self.vid_writer = None
|
||||
self.vid_count = 0
|
||||
self._pause = True
|
||||
self.n_step = 0
|
||||
self.sim_time = 0.0
|
||||
|
||||
def load(
|
||||
self,
|
||||
@@ -41,16 +47,19 @@ class MujocoSimulator:
|
||||
robot_xml: str = None,
|
||||
terrain_spawn_xy: list = None,
|
||||
robot_spawn_height: float = None,
|
||||
default_dof_pos: list = None,
|
||||
):
|
||||
""" Load terrain and robot into the simulator, support re-loading. """
|
||||
if terrain_xml is not None:
|
||||
self.terrain_xml = pares_path(terrain_xml)
|
||||
self.terrain_xml = parse_path(terrain_xml)
|
||||
if robot_xml is not None:
|
||||
self.robot_xml = pares_path(robot_xml)
|
||||
self.robot_xml = parse_path(robot_xml)
|
||||
if terrain_spawn_xy is not None:
|
||||
self.terrain_spawn_xy = terrain_spawn_xy
|
||||
if robot_spawn_height is not None:
|
||||
self.robot_spawn_height = robot_spawn_height
|
||||
if default_dof_pos is not None:
|
||||
self.default_dof_pos = default_dof_pos
|
||||
|
||||
terrain_xml = self.terrain_xml
|
||||
robot_xml = self.robot_xml
|
||||
@@ -58,6 +67,8 @@ class MujocoSimulator:
|
||||
robot_spawn_height = self.robot_spawn_height
|
||||
if terrain_xml is None or robot_xml is None:
|
||||
raise ValueError("Terrain and robot XML paths must be provided.")
|
||||
if default_dof_pos is None:
|
||||
raise ValueError("Default DOF positions must be provided.")
|
||||
|
||||
robot_mjcf = mjcf.from_path(robot_xml)
|
||||
terrain_mjcf = mjcf.from_path(terrain_xml)
|
||||
@@ -74,6 +85,10 @@ class MujocoSimulator:
|
||||
self.mj_model = self.mj_physics.model.ptr
|
||||
self.mj_data = self.mj_physics.data.ptr
|
||||
self.mj_model.opt.timestep = self.cfg.physics.simulation_dt
|
||||
self.sim_dt = self.cfg.physics.simulation_dt
|
||||
self.mj_data.qpos[7:] = default_dof_pos
|
||||
mujoco.mj_forward(self.mj_model, self.mj_data)
|
||||
|
||||
self.headless = self.cfg.viewer.headless
|
||||
if self.cfg.render.save_video and self.headless:
|
||||
logger.warning("Cannot save video in headless mode, disabling video saving.")
|
||||
@@ -94,23 +109,33 @@ class MujocoSimulator:
|
||||
vid_path = str(vid_dir / f"sim_video_{self.vid_count:03d}.mp4")
|
||||
self.vid_writer = imageio.get_writer(
|
||||
vid_path,
|
||||
fps=int(1 / self.cfg.physics.simulation_dt),
|
||||
fps=self.cfg.render.video_fps,
|
||||
)
|
||||
self.vid_frame_skip = int(1 / (self.cfg.render.video_fps * self.sim_dt * 2))
|
||||
logger.info(f"Saving simulation video to: {vid_path}")
|
||||
self.vid_count += 1
|
||||
self._pause = False
|
||||
self.n_step = 0
|
||||
self.sim_time = 0.0
|
||||
self.preload_sensors()
|
||||
|
||||
|
||||
# Robot controller placeholders
|
||||
self.action = None
|
||||
self.p_gains = None
|
||||
self.d_gains = None
|
||||
self.control_type = None
|
||||
|
||||
def key_callback(self, keycode):
|
||||
if keycode == 32:
|
||||
self._pause = not self._pause
|
||||
logger.info(f"Pause toggled: {self._pause}")
|
||||
|
||||
def step(self) -> dict:
|
||||
def step(self) -> SimData:
|
||||
""" Simulation step, pause will block thread. """
|
||||
while self._pause:
|
||||
time.sleep(0.1)
|
||||
self.update_torque()
|
||||
self.mj_physics.step()
|
||||
if self.viewer is not None:
|
||||
if self.viewer.is_running():
|
||||
@@ -120,7 +145,7 @@ class MujocoSimulator:
|
||||
if time_untile_next_render > 0:
|
||||
time.sleep(time_untile_next_render)
|
||||
self.viewer.sync()
|
||||
if self.vid_writer is not None:
|
||||
if self.vid_writer is not None and self.n_step % self.vid_frame_skip == 0:
|
||||
self.renderer.update_scene(self.mj_data, camera=self.viewer.cam)
|
||||
frame = self.renderer.render()
|
||||
self.vid_writer.append_data(frame)
|
||||
@@ -129,43 +154,75 @@ class MujocoSimulator:
|
||||
logger.warning("Viewer closed by user, stop video recording.")
|
||||
self.close_viewer()
|
||||
|
||||
n_sensor = self.mj_model.nsensor
|
||||
|
||||
proprio = RobotProprioception(
|
||||
self.proprio = proprio = RobotProprioception(
|
||||
joint=JointState(
|
||||
pos=self.get_sensor_data('joint_pos'),
|
||||
vel=self.get_sensor_data('joint_vel'),
|
||||
force=self.get_sensor_data('joint_eff'),
|
||||
),
|
||||
imu=IMUState(
|
||||
quat=self.get_sensor_data('imu_quat'),
|
||||
ang_vel=self.get_sensor_data('imu_ang_vel'),
|
||||
acc=self.get_sensor_data('imu_acc'),
|
||||
pos=self.get_sensor_data('imu_pos'),
|
||||
quat=self.get_sensor_data('imu_quat'),
|
||||
acc=self.get_sensor_data('imu_acc'),
|
||||
lin_vel=self.get_sensor_data('imu_lin_vel'),
|
||||
ang_vel=self.get_sensor_data('imu_ang_vel'),
|
||||
),
|
||||
base=BaseState(
|
||||
pos=self.mj_data.qpos[:3], # world frame
|
||||
quat=self.mj_data.qpos[3:7], # world frame
|
||||
vel=self.mj_data.qvel[:3], # body frame
|
||||
lin_vel=self.mj_data.qvel[:3], # body frame
|
||||
ang_vel=self.mj_data.qvel[3:6], # body frame
|
||||
)
|
||||
)
|
||||
if self.n_step % int(0.1 / self.sim_dt) == 0:
|
||||
logger.log(value=np.mean(proprio.imu.quat - proprio.base.quat), tag="sim/delta_quat", step=self.n_step)
|
||||
logger.log(value=np.mean(proprio.imu.ang_vel - proprio.base.ang_vel), tag="sim/delta_ang_vel", step=self.n_step)
|
||||
logger.log(value=np.mean(proprio.imu.lin_vel - proprio.base.lin_vel), tag="sim/delta_lin_vel", step=self.n_step)
|
||||
if self.n_step == 0:
|
||||
self.debug_print_proprio_shapes(proprio)
|
||||
self.n_step += 1
|
||||
self.debug_print_proprio_shapes()
|
||||
|
||||
return proprio
|
||||
sim_data = SimData(
|
||||
n_step=self.n_step,
|
||||
sim_dt=self.sim_dt,
|
||||
proprio=proprio
|
||||
)
|
||||
|
||||
self.n_step += 1
|
||||
self.sim_time = self.n_step * self.sim_dt
|
||||
return sim_data
|
||||
|
||||
def reset(self):
|
||||
""" Reset the simulator to initial state. """
|
||||
self.mj_physics.reset()
|
||||
self.mj_data.qpos[7:] = self.default_dof_pos
|
||||
mujoco.mj_forward(self.mj_model, self.mj_data)
|
||||
|
||||
self.action = None
|
||||
if self.viewer is not None:
|
||||
self.viewer.sync()
|
||||
|
||||
def apply_action(self, action: np.ndarray):
|
||||
""" Apply action to the simulator. """
|
||||
self.mj_data.ctrl[:] = action
|
||||
def setup_action(self,
|
||||
action: np.ndarray,
|
||||
p_gains: np.ndarray = None,
|
||||
d_gains: np.ndarray = None,
|
||||
control_type: Literal['P'] = 'P'
|
||||
):
|
||||
""" Setup action to the simulator. """
|
||||
self.action = action
|
||||
self.p_gains = p_gains
|
||||
self.d_gains = d_gains
|
||||
self.control_type = control_type
|
||||
|
||||
def update_torque(self):
|
||||
if self.action is None:
|
||||
return
|
||||
dof_pos = self.proprio.joint.pos
|
||||
dof_vel = self.proprio.joint.vel
|
||||
if self.control_type == 'P':
|
||||
torques = self.p_gains * (self.action - dof_pos) - self.d_gains * dof_vel
|
||||
else:
|
||||
raise NotImplementedError(f"Control type '{self.control_type}' not implemented.")
|
||||
self.mj_data.ctrl[:] = torques
|
||||
|
||||
def close_viewer(self):
|
||||
""" Close the viewer and video writer. """
|
||||
@@ -189,12 +246,16 @@ class MujocoSimulator:
|
||||
self.imu_acc = self.find_sensors(tag_name="accelerometer")
|
||||
self.imu_pos = self.find_sensors(tag_name="framepos")
|
||||
self.imu_lin_vel = self.find_sensors(tag_name="framelinvel")
|
||||
actuator_names = [mujoco.mj_id2name(self.mj_model, mujoco.mjtObj.mjOBJ_ACTUATOR, i) for i in range(self.mj_model.nu)]
|
||||
logger.info(
|
||||
f"\n{'='*20} XML SENSOR NAMES {'='*20}\n"
|
||||
f"""Joint Position Sensors [{len(self.joint_pos_sensor_names)}]: {self.joint_pos_sensor_names}\n"""
|
||||
f"""Joint Velocity Sensors [{len(self.joint_vel_sensor_names)}]: {self.joint_vel_sensor_names}\n"""
|
||||
f"""Joint Effort Sensors [{len(self.joint_eff_sensor_names)}]: {self.joint_eff_sensor_names}\n"""
|
||||
f"""IMU Sensors: Quat{self.imu_quat}, AngVel{self.imu_ang_vel}, Acc{self.imu_acc}, Pos{self.imu_pos}, LinVel{self.imu_lin_vel}"""
|
||||
f"""Joint Position Sensors [{len(self.joint_pos_sensor_names)}]: {[x.rsplit('/')[-1] for x in self.joint_pos_sensor_names]}\n"""
|
||||
f"""Joint Velocity Sensors [{len(self.joint_vel_sensor_names)}]: {[x.rsplit('/')[-1] for x in self.joint_vel_sensor_names]}\n"""
|
||||
f"""Joint Effort Sensors [{len(self.joint_eff_sensor_names)}]: {[x.rsplit('/')[-1] for x in self.joint_eff_sensor_names]}\n"""
|
||||
f"""Actuators [{len(actuator_names)}]: {[x.rsplit('/')[-1] for x in actuator_names]}\n"""
|
||||
f"""IMU Sensors: Quat{self.imu_quat}, AngVel{self.imu_ang_vel}, Acc{self.imu_acc}, Pos{self.imu_pos}, LinVel{self.imu_lin_vel}\n"""
|
||||
f"""!!!Checkout actuators order is consistent with joint sensors!!!\n"""
|
||||
f"{'='*58}"
|
||||
)
|
||||
|
||||
# Cache sensor indices
|
||||
@@ -259,7 +320,7 @@ class MujocoSimulator:
|
||||
data_list.append(self.mj_data.sensordata[adr:adr+dim])
|
||||
return np.concatenate(data_list)
|
||||
|
||||
def debug_print_proprio_shapes(self, proprio: RobotProprioception):
|
||||
def debug_print_proprio_shapes(self):
|
||||
"""Log shapes (or lengths) of each numpy vector inside a RobotProprioception.
|
||||
|
||||
This helps debug mismatched sensor sizes between robots.
|
||||
@@ -271,9 +332,9 @@ class MujocoSimulator:
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
jp = proprio.joint
|
||||
bs = proprio.base
|
||||
imu = proprio.imu
|
||||
jp = self.proprio.joint
|
||||
bs = self.proprio.base
|
||||
imu = self.proprio.imu
|
||||
|
||||
logger.info("Proprioception shapes:")
|
||||
logger.info(f" joint.pos: { _shape(jp.pos) }")
|
||||
@@ -282,7 +343,7 @@ class MujocoSimulator:
|
||||
|
||||
logger.info(f" base.pos: { _shape(bs.pos) }")
|
||||
logger.info(f" base.quat: { _shape(bs.quat) }")
|
||||
logger.info(f" base.vel: { _shape(bs.vel) }")
|
||||
logger.info(f" base.vel: { _shape(bs.lin_vel) }")
|
||||
logger.info(f" base.ang_vel: { _shape(bs.ang_vel) }")
|
||||
|
||||
logger.info(f" imu.quat: { _shape(imu.quat) }")
|
||||
|
||||
@@ -11,19 +11,25 @@ class JointState:
|
||||
class BaseState:
|
||||
pos: np.ndarray
|
||||
quat: np.ndarray
|
||||
vel: np.ndarray
|
||||
lin_vel: np.ndarray
|
||||
ang_vel: np.ndarray
|
||||
|
||||
@dataclass
|
||||
class IMUState:
|
||||
quat: np.ndarray
|
||||
ang_vel: np.ndarray
|
||||
acc: np.ndarray
|
||||
pos: np.ndarray
|
||||
quat: np.ndarray
|
||||
acc: np.ndarray
|
||||
lin_vel: np.ndarray
|
||||
ang_vel: np.ndarray
|
||||
|
||||
@dataclass
|
||||
class RobotProprioception:
|
||||
joint: JointState
|
||||
base: BaseState
|
||||
imu: IMUState
|
||||
|
||||
@dataclass
|
||||
class SimData:
|
||||
n_step: int
|
||||
sim_dt: float
|
||||
proprio: RobotProprioception
|
||||
|
||||
Reference in New Issue
Block a user