This commit is contained in:
wty-yy
2025-11-27 17:58:37 +08:00
parent fc7e184741
commit 4cc0687114
5 changed files with 194 additions and 13 deletions

View File

@@ -3,6 +3,7 @@
## 20251127
### v0.1.2
1. 修改logger
2. 完成mujoco的sensor数据, sim_data数据类处理
## 20251124-25
### v0.1.1

View File

@@ -15,5 +15,6 @@ from robogauge.utils.logger import logger
if __name__ == '__main__':
args = parse_args()
logger.create(args.experiment_name)
logger.info(f"Starting experiment: {args.experiment_name}")
pipeline: BasePipeline = task_register.make_pipeline(args.task_name, args=args)
pipeline.run()

View File

@@ -11,6 +11,7 @@ import mujoco
import mujoco.viewer
from dm_control import mjcf
import re
import time
import imageio
import numpy as np
@@ -18,6 +19,7 @@ import numpy as np
from robogauge.utils.logger import logger
from robogauge.utils.helpers import pares_path
from robogauge.tasks.simulator.mujoco_config import MujocoConfig
from robogauge.tasks.simulator.sim_data import RobotProprioception, JointState, BaseState, IMUState
class MujocoSimulator:
def __init__(self, sim_cfg: MujocoConfig):
@@ -31,6 +33,7 @@ class MujocoSimulator:
self.vid_writer = None
self.vid_count = 0
self._pause = True
self.n_step = 0
def load(
self,
@@ -76,10 +79,15 @@ class MujocoSimulator:
logger.warning("Cannot save video in headless mode, disabling video saving.")
self.cfg.render.save_video = False
if not self.headless:
self.viewer = mujoco.viewer.launch_passive(self.mj_model, self.mj_data, key_callback=self.key_callback)
self.viewer = mujoco.viewer.launch_passive(
self.mj_model, self.mj_data, key_callback=self.key_callback
)
self.last_render_time = time.time()
if self.cfg.render.save_video:
self.renderer = mujoco.Renderer(self.mj_model, height=self.cfg.render.height, width=self.cfg.render.width)
self.renderer = mujoco.Renderer(
self.mj_model, height=self.cfg.render.height,
width=self.cfg.render.width
)
vid_dir = logger.log_dir / "videos"
vid_dir.mkdir(parents=True, exist_ok=True)
@@ -91,6 +99,8 @@ class MujocoSimulator:
logger.info(f"Saving simulation video to: {vid_path}")
self.vid_count += 1
self._pause = False
self.n_step = 0
self.preload_sensors()
def key_callback(self, keycode):
if keycode == 32:
@@ -104,7 +114,9 @@ class MujocoSimulator:
self.mj_physics.step()
if self.viewer is not None:
if self.viewer.is_running():
time_untile_next_render = self.cfg.physics.simulation_dt - (time.time() - self.last_render_time)
time_untile_next_render = self.cfg.physics.simulation_dt - (
time.time() - self.last_render_time
)
if time_untile_next_render > 0:
time.sleep(time_untile_next_render)
self.viewer.sync()
@@ -117,8 +129,33 @@ class MujocoSimulator:
logger.warning("Viewer closed by user, stop video recording.")
self.close_viewer()
info = {}
return info
n_sensor = self.mj_model.nsensor
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'),
lin_vel=self.get_sensor_data('imu_lin_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
ang_vel=self.mj_data.qvel[3:6], # body frame
)
)
if self.n_step == 0:
self.debug_print_proprio_shapes(proprio)
self.n_step += 1
return proprio
def reset(self):
""" Reset the simulator to initial state. """
@@ -141,3 +178,115 @@ class MujocoSimulator:
self.vid_writer = None
self.renderer = None
logger.info("Closing video writer.")
def preload_sensors(self):
# Preload sensor names
self.joint_pos_sensor_names = self.find_sensors(tag_name="jointpos")
self.joint_vel_sensor_names = self.find_sensors(tag_name="jointvel")
self.joint_eff_sensor_names = self.find_sensors(tag_name="jointactuatorfrc")
self.imu_quat = self.find_sensors(tag_name="framequat")
self.imu_ang_vel = self.find_sensors(tag_name="gyro")
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")
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}"""
)
# Cache sensor indices
self.sensor_cache = {}
all_lists = {
'joint_pos': self.joint_pos_sensor_names,
'joint_vel': self.joint_vel_sensor_names,
'joint_eff': self.joint_eff_sensor_names,
'imu_quat': self.imu_quat,
'imu_ang_vel': self.imu_ang_vel,
'imu_acc': self.imu_acc,
'imu_pos': self.imu_pos,
'imu_lin_vel': self.imu_lin_vel
}
for key, name_list in all_lists.items():
indices = []
for name in name_list:
sid = mujoco.mj_name2id(self.mj_model, mujoco.mjtObj.mjOBJ_SENSOR, name)
if sid == -1: continue
adr = int(self.mj_model.sensor_adr[sid])
dim = int(self.mj_model.sensor_dim[sid])
indices.append((adr, dim))
self.sensor_cache[key] = indices
def find_sensors(self, *, pattern: re.Pattern = None, tag_name: str = None) -> list:
model = self.mj_model
found = []
tag_map = {
"jointpos": mujoco.mjtSensor.mjSENS_JOINTPOS,
"jointvel": mujoco.mjtSensor.mjSENS_JOINTVEL,
"jointactuatorfrc": mujoco.mjtSensor.mjSENS_JOINTACTFRC,
"accelerometer": mujoco.mjtSensor.mjSENS_ACCELEROMETER,
"gyro": mujoco.mjtSensor.mjSENS_GYRO,
"framepos": mujoco.mjtSensor.mjSENS_FRAMEPOS,
"framequat": mujoco.mjtSensor.mjSENS_FRAMEQUAT,
"framelinvel": mujoco.mjtSensor.mjSENS_FRAMELINVEL,
"frameangvel": mujoco.mjtSensor.mjSENS_FRAMEANGVEL,
}
tag_type_id = None
if tag_name:
if tag_name not in tag_map:
logger.warning(f"Unknown tag_name '{tag_name}', ignoring tag filter.")
return []
tag_type_id = tag_map[tag_name]
for i in range(model.nsensor):
if tag_type_id is not None and model.sensor_type[i] != tag_type_id:
continue
name = mujoco.mj_id2name(model, mujoco.mjtObj.mjOBJ_SENSOR, i)
if pattern is None or name and pattern.search(name):
found.append(name)
return found
def get_sensor_data(self, cache_key: str) -> np.ndarray:
ids = self.sensor_cache.get(cache_key, [])
if not ids:
return np.array([])
data_list = []
for adr, dim in ids:
data_list.append(self.mj_data.sensordata[adr:adr+dim])
return np.concatenate(data_list)
def debug_print_proprio_shapes(self, proprio: RobotProprioception):
"""Log shapes (or lengths) of each numpy vector inside a RobotProprioception.
This helps debug mismatched sensor sizes between robots.
"""
def _shape(x):
try:
arr = np.asarray(x)
return arr.shape
except Exception:
return None
jp = proprio.joint
bs = proprio.base
imu = proprio.imu
logger.info("Proprioception shapes:")
logger.info(f" joint.pos: { _shape(jp.pos) }")
logger.info(f" joint.vel: { _shape(jp.vel) }")
logger.info(f" joint.force: { _shape(jp.force) }")
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.ang_vel: { _shape(bs.ang_vel) }")
logger.info(f" imu.quat: { _shape(imu.quat) }")
logger.info(f" imu.ang_vel: { _shape(imu.ang_vel) }")
logger.info(f" imu.acc: { _shape(imu.acc) }")
logger.info(f" imu.pos: { _shape(imu.pos) }")
logger.info(f" imu.lin_vel: { _shape(imu.lin_vel) }")

View File

@@ -0,0 +1,29 @@
import numpy as np
from dataclasses import dataclass
@dataclass
class JointState:
pos: np.ndarray
vel: np.ndarray
force: np.ndarray
@dataclass
class BaseState:
pos: np.ndarray
quat: np.ndarray
vel: np.ndarray
ang_vel: np.ndarray
@dataclass
class IMUState:
quat: np.ndarray
ang_vel: np.ndarray
acc: np.ndarray
pos: np.ndarray
lin_vel: np.ndarray
@dataclass
class RobotProprioception:
joint: JointState
base: BaseState
imu: IMUState

View File

@@ -69,6 +69,7 @@ class Logger:
"""
self.logger = logging.getLogger(experiment_name + "_logger")
self.logger.setLevel(log_level)
self.logger.propagate = False
console_formatter = ColorFormatter( # console output format
fmt="%(asctime)s - %(color_level)s - %(filename)s:%(lineno)d - %(message)s",
@@ -94,19 +95,19 @@ class Logger:
self.logger.addHandler(fh)
def debug(self, msg, *args, **kwargs):
self.logger.debug(msg, *args, **kwargs)
self.logger.debug(msg, *args, **kwargs, stacklevel=2)
def info(self, msg, *args, **kwargs):
self.logger.info(msg, *args, **kwargs)
self.logger.info(msg, *args, **kwargs, stacklevel=2)
def warning(self, msg, *args, **kwargs):
self.logger.warning(msg, *args, **kwargs)
self.logger.warning(msg, *args, **kwargs, stacklevel=2)
def error(self, msg, *args, **kwargs):
self.logger.error(msg, *args, **kwargs)
self.logger.error(msg, *args, **kwargs, stacklevel=2)
def critical(self, msg, *args, **kwargs):
self.logger.critical(msg, *args, **kwargs)
self.logger.critical(msg, *args, **kwargs, stacklevel=2)
logger = Logger()