v0.1.2
This commit is contained in:
@@ -3,6 +3,7 @@
|
|||||||
## 20251127
|
## 20251127
|
||||||
### v0.1.2
|
### v0.1.2
|
||||||
1. 修改logger
|
1. 修改logger
|
||||||
|
2. 完成mujoco的sensor数据, sim_data数据类处理
|
||||||
|
|
||||||
## 20251124-25
|
## 20251124-25
|
||||||
### v0.1.1
|
### v0.1.1
|
||||||
|
|||||||
@@ -15,5 +15,6 @@ from robogauge.utils.logger import logger
|
|||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
args = parse_args()
|
args = parse_args()
|
||||||
logger.create(args.experiment_name)
|
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: BasePipeline = task_register.make_pipeline(args.task_name, args=args)
|
||||||
pipeline.run()
|
pipeline.run()
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ import mujoco
|
|||||||
import mujoco.viewer
|
import mujoco.viewer
|
||||||
from dm_control import mjcf
|
from dm_control import mjcf
|
||||||
|
|
||||||
|
import re
|
||||||
import time
|
import time
|
||||||
import imageio
|
import imageio
|
||||||
import numpy as np
|
import numpy as np
|
||||||
@@ -18,6 +19,7 @@ import numpy as np
|
|||||||
from robogauge.utils.logger import logger
|
from robogauge.utils.logger import logger
|
||||||
from robogauge.utils.helpers import pares_path
|
from robogauge.utils.helpers import pares_path
|
||||||
from robogauge.tasks.simulator.mujoco_config import MujocoConfig
|
from robogauge.tasks.simulator.mujoco_config import MujocoConfig
|
||||||
|
from robogauge.tasks.simulator.sim_data import RobotProprioception, JointState, BaseState, IMUState
|
||||||
|
|
||||||
class MujocoSimulator:
|
class MujocoSimulator:
|
||||||
def __init__(self, sim_cfg: MujocoConfig):
|
def __init__(self, sim_cfg: MujocoConfig):
|
||||||
@@ -31,6 +33,7 @@ class MujocoSimulator:
|
|||||||
self.vid_writer = None
|
self.vid_writer = None
|
||||||
self.vid_count = 0
|
self.vid_count = 0
|
||||||
self._pause = True
|
self._pause = True
|
||||||
|
self.n_step = 0
|
||||||
|
|
||||||
def load(
|
def load(
|
||||||
self,
|
self,
|
||||||
@@ -76,10 +79,15 @@ class MujocoSimulator:
|
|||||||
logger.warning("Cannot save video in headless mode, disabling video saving.")
|
logger.warning("Cannot save video in headless mode, disabling video saving.")
|
||||||
self.cfg.render.save_video = False
|
self.cfg.render.save_video = False
|
||||||
if not self.headless:
|
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()
|
self.last_render_time = time.time()
|
||||||
if self.cfg.render.save_video:
|
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 = logger.log_dir / "videos"
|
||||||
vid_dir.mkdir(parents=True, exist_ok=True)
|
vid_dir.mkdir(parents=True, exist_ok=True)
|
||||||
@@ -91,6 +99,8 @@ class MujocoSimulator:
|
|||||||
logger.info(f"Saving simulation video to: {vid_path}")
|
logger.info(f"Saving simulation video to: {vid_path}")
|
||||||
self.vid_count += 1
|
self.vid_count += 1
|
||||||
self._pause = False
|
self._pause = False
|
||||||
|
self.n_step = 0
|
||||||
|
self.preload_sensors()
|
||||||
|
|
||||||
def key_callback(self, keycode):
|
def key_callback(self, keycode):
|
||||||
if keycode == 32:
|
if keycode == 32:
|
||||||
@@ -104,7 +114,9 @@ class MujocoSimulator:
|
|||||||
self.mj_physics.step()
|
self.mj_physics.step()
|
||||||
if self.viewer is not None:
|
if self.viewer is not None:
|
||||||
if self.viewer.is_running():
|
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:
|
if time_untile_next_render > 0:
|
||||||
time.sleep(time_untile_next_render)
|
time.sleep(time_untile_next_render)
|
||||||
self.viewer.sync()
|
self.viewer.sync()
|
||||||
@@ -117,8 +129,33 @@ class MujocoSimulator:
|
|||||||
logger.warning("Viewer closed by user, stop video recording.")
|
logger.warning("Viewer closed by user, stop video recording.")
|
||||||
self.close_viewer()
|
self.close_viewer()
|
||||||
|
|
||||||
info = {}
|
n_sensor = self.mj_model.nsensor
|
||||||
return info
|
|
||||||
|
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):
|
def reset(self):
|
||||||
""" Reset the simulator to initial state. """
|
""" Reset the simulator to initial state. """
|
||||||
@@ -141,3 +178,115 @@ class MujocoSimulator:
|
|||||||
self.vid_writer = None
|
self.vid_writer = None
|
||||||
self.renderer = None
|
self.renderer = None
|
||||||
logger.info("Closing video writer.")
|
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) }")
|
||||||
|
|||||||
29
robogauge/tasks/simulator/sim_data.py
Normal file
29
robogauge/tasks/simulator/sim_data.py
Normal 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
|
||||||
@@ -69,6 +69,7 @@ class Logger:
|
|||||||
"""
|
"""
|
||||||
self.logger = logging.getLogger(experiment_name + "_logger")
|
self.logger = logging.getLogger(experiment_name + "_logger")
|
||||||
self.logger.setLevel(log_level)
|
self.logger.setLevel(log_level)
|
||||||
|
self.logger.propagate = False
|
||||||
|
|
||||||
console_formatter = ColorFormatter( # console output format
|
console_formatter = ColorFormatter( # console output format
|
||||||
fmt="%(asctime)s - %(color_level)s - %(filename)s:%(lineno)d - %(message)s",
|
fmt="%(asctime)s - %(color_level)s - %(filename)s:%(lineno)d - %(message)s",
|
||||||
@@ -94,19 +95,19 @@ class Logger:
|
|||||||
self.logger.addHandler(fh)
|
self.logger.addHandler(fh)
|
||||||
|
|
||||||
def debug(self, msg, *args, **kwargs):
|
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):
|
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):
|
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):
|
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):
|
def critical(self, msg, *args, **kwargs):
|
||||||
self.logger.critical(msg, *args, **kwargs)
|
self.logger.critical(msg, *args, **kwargs, stacklevel=2)
|
||||||
|
|
||||||
logger = Logger()
|
logger = Logger()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user