diff --git a/examples/simple/eval.py b/examples/simple/eval.py index 8efe434..a78b9f6 100644 --- a/examples/simple/eval.py +++ b/examples/simple/eval.py @@ -60,7 +60,7 @@ def main(): model = get_latest_model(log_path) # Setup environment - env = Go2SimpleEnv(num_envs=1, headless=False) + env = Go2SimpleEnv(num_envs=1, headless=False, env_mode="eval") env = RslRlWrapper(env) env.build() diff --git a/examples/simple/eval_deploy_generic.py b/examples/simple/eval_deploy_generic.py new file mode 100644 index 0000000..7c041c7 --- /dev/null +++ b/examples/simple/eval_deploy_generic.py @@ -0,0 +1,111 @@ +# !!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!! +# PLACE_HOLDER +# !!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!! + + +import os +import glob +import torch +import pickle +import argparse +from importlib import metadata +import genesis as gs + +from environment import Go2SimpleEnv + +try: + try: + if metadata.version("rsl-rl"): + raise ImportError + except metadata.PackageNotFoundError: + if metadata.version("rsl-rl-lib").startswith("1."): + raise ImportError +except (metadata.PackageNotFoundError, ImportError) as e: + raise ImportError("Please install install 'rsl-rl-lib>=2.2.4'.") from e +from rsl_rl.runners import OnPolicyRunner + +EXPERIMENT_NAME = "go2-simple" + +parser = argparse.ArgumentParser(add_help=True) +parser.add_argument("-d", "--device", type=str, default="gpu") +parser.add_argument("-e", "--exp_name", type=str, default=EXPERIMENT_NAME) +args = parser.parse_args() + + +def get_latest_model(log_dir: str) -> str: + """ + Get the last model from the log directory + """ + model_checkpoints = glob.glob(os.path.join(log_dir, "model_*.pt")) + if len(model_checkpoints) == 0: + print( + f"Warning: No model files found at '{log_dir}' (you might need to train more)." + ) + exit(1) + # Sort by the file with the highest number + sorted_models = sorted( + model_checkpoints, + key=lambda x: int(os.path.basename(x).split("_")[1].split(".")[0]), + ) + return sorted_models[-1] + + +def setup_observations(env: Go2SimpleEnv): + # Assign a function to each observation that will return real sensor data + obs = env.observation_manager.cfg + obs["angle_velocity"].fn = lambda env: torch.zeros(3) + obs["linear_velocity"].fn = lambda env: torch.zeros(3) + obs["projected_gravity"].fn = lambda env: torch.zeros(3) + obs["dof_position"].fn = lambda env: torch.zeros(12) + obs["dof_velocity"].fn = lambda env: torch.zeros(12) + # No need to update the actions observation, as that will be handled by the environment automatically + + +def main(): + # Processor backend (GPU or CPU) + backend = gs.gpu + if args.device == "cpu": + backend = gs.cpu + torch.set_default_device("cpu") + gs.init(logging_level="warning", backend=backend) + + # Load training configuration + log_path = f"./logs/{args.exp_name}" + [cfg] = pickle.load(open(f"{log_path}/cfgs.pkl", "rb")) + model = get_latest_model(log_path) + + # Setup environment + env = Go2SimpleEnv(num_envs=1, headless=False, mode="real") + env.build() + + # Update observations to use real sensors + setup_observations(env) + + # Load the trained policy + print("🎬 Loading last model...") + runner = OnPolicyRunner(env, cfg, log_path, device=gs.device) + runner.load(model) + policy = runner.get_inference_policy(device=gs.device) + + try: + obs, _ = env.reset() + with torch.no_grad(): + while True: + actions = policy(obs) + obs, _rews, _dones, _infos = env.step(actions) + + # Get actions to send to the actuators + actions = env.action_manager.get_actions() + # ...send the actuator values to the actuators + + except KeyboardInterrupt: + pass + except gs.GenesisException as e: + if e.message != "Viewer closed.": + raise e + except Exception as e: + raise e + + +if __name__ == "__main__": + main() diff --git a/examples/simple/eval_deploy_ros2.py b/examples/simple/eval_deploy_ros2.py new file mode 100644 index 0000000..4b7ac67 --- /dev/null +++ b/examples/simple/eval_deploy_ros2.py @@ -0,0 +1,112 @@ +import os +import glob +import torch +import pickle +import argparse +from importlib import metadata + +import rclpy +from ros2_interface import RosInterface + +from environment import Go2SimpleEnv + +try: + try: + if metadata.version("rsl-rl"): + raise ImportError + except metadata.PackageNotFoundError: + if metadata.version("rsl-rl-lib").startswith("1."): + raise ImportError +except (metadata.PackageNotFoundError, ImportError) as e: + raise ImportError("Please install install 'rsl-rl-lib>=2.2.4'.") from e +from rsl_rl.runners import OnPolicyRunner + +EXPERIMENT_NAME = "go2-simple" + +parser = argparse.ArgumentParser(add_help=True) +parser.add_argument("-d", "--device", type=str, default="gpu") +parser.add_argument("-e", "--exp_name", type=str, default=EXPERIMENT_NAME) +args = parser.parse_args() + + +def get_latest_model(log_dir: str) -> str: + """ + Get the last model from the log directory + """ + model_checkpoints = glob.glob(os.path.join(log_dir, "model_*.pt")) + if len(model_checkpoints) == 0: + print( + f"Warning: No model files found at '{log_dir}' (you might need to train more)." + ) + exit(1) + # Sort by the file with the highest number + sorted_models = sorted( + model_checkpoints, + key=lambda x: int(os.path.basename(x).split("_")[1].split(".")[0]), + ) + return sorted_models[-1] + + +def setup_observations(env: Go2SimpleEnv, ros_interface: RosInterface): + # Assign a function to each observation that will return real sensor data + obs = env.observation_manager.cfg + obs["angle_velocity"].fn = lambda env: ros_interface.get_angular_velocity() + obs["linear_velocity"].fn = lambda env: ros_interface.get_linear_velocity() + obs["projected_gravity"].fn = lambda env: torch.zeros(3) + obs["dof_position"].fn = lambda env: ros_interface.get_dofs_position() + obs["dof_velocity"].fn = lambda env: ros_interface.get_dofs_velocity() + # No need to update the actions observation, as that will be handled by the environment automatically + + +def main(): + # Processor backend (GPU or CPU) + rclpy.init() + if args.device == "cpu": + device = torch.device("cpu") + torch.set_default_device("cpu") + elif args.device == "gpu": + device = torch.device("cpu") + torch.set_default_device("cuda:0") + + # Load training configuration + log_path = f"./logs/{args.exp_name}" + [cfg] = pickle.load(open(f"{log_path}/cfgs.pkl", "rb")) + model = get_latest_model(log_path) + + # Setup environment + env = Go2SimpleEnv(num_envs=1, headless=False, mode="real") + env.build() + pos_joints = [] + vel_joints = [] + force_joints = [] + ros_interface = RosInterface(pos_joints, vel_joints, force_joints) + if rclpy.ok(): + rclpy.spin_once(ros_interface, timeout_sec=0.1) + + # Update observations to use real sensors + setup_observations(env, ros_interface=ros_interface) + + # Load the trained policy + print("🎬 Loading last model...") + runner = OnPolicyRunner(env, cfg, log_path, device=device) + runner.load(model) + policy = runner.get_inference_policy(device=device) + + try: + obs, _ = env.reset() + with torch.no_grad(): + while True and rclpy.ok(): + rclpy.spin_once(ros_interface, timeout_sec=0.1) + actions = policy(obs) + obs, _rews, _dones, _infos = env.step(actions) + # Get actions to send to the ros_interface + ros_interface._pos_actions = env.action_manager.get_actions() + + except KeyboardInterrupt: + pass + except Exception as e: + raise e + + +if __name__ == "__main__": + main() diff --git a/examples/simple/ros2_interface.py b/examples/simple/ros2_interface.py new file mode 100644 index 0000000..86a28ef --- /dev/null +++ b/examples/simple/ros2_interface.py @@ -0,0 +1,182 @@ +import rclpy +from rclpy.node import Node +from rclpy.clock import Clock +from sensor_msgs.msg import JointState +from geometry_msgs.msg import Twist, Pose +import torch + + +class RosInterface(Node): + def __init__(self, pos_joints, vel_joints, force_joints): + super().__init__("ros_interface") + self._ros_node = self + self._ros_clock = Clock() + + self._pos_joints = pos_joints + self._vel_joints = vel_joints + self._force_joints = force_joints + self._all_joints = pos_joints + vel_joints + force_joints + + self._num_pos_joints = len(pos_joints) + self._num_vel_joints = len(vel_joints) + self._num_force_joints = len(force_joints) + self._num_joints = ( + self._num_pos_joints + self._num_vel_joints + self._num_force_joints + ) + + self._pos_state = None + self._vel_state = None + self._force_state = None + + self._robot_lin_vel = None + self._robot_ang_vel = None + self._robot_pos = None + self._robot_quat = None + + self.setup_ros() + + def setup_ros(self): + """ + Setup ROS publishers and subscribers. + Should be called after env.build() so that joint names are available. + """ + # Initialize action buffers + self._pos_actions = torch.zeros(self._num_pos_joints) + self._vel_actions = torch.zeros(self._num_vel_joints) + self._force_actions = torch.zeros(self._num_force_joints) + + # Setup subscribers and publishers + self._setup_joint_action_publisher() + self._setup_robot_twist_subscriber() + self._setup_robot_pose_subscriber() + self._setup_joint_state_subscriber() + + def _current_timestep(self): + """ + Get the current sim time + """ + return self._ros_clock.now().to_msg() + + def _setup_joint_action_publisher(self): + print("Joint actions Publisher started") + + def joint_action_callback(): + joint_state_msg = JointState() + joint_state_msg.header.stamp = self._current_timestep() + joint_state_msg.name = self._all_joints + + pos_actions = [] + vel_actions = [] + force_actions = [] + for pos_idx in range(self._num_pos_joints): + pos_actions.append(self._pos_actions[pos_idx].item()) + vel_actions.append(None) + force_actions.append(None) + for vel_idx in range(self._num_vel_joints): + pos_actions.append(None) + vel_actions.append(self._vel_actions[vel_idx].item()) + force_actions.append(None) + for force_idx in range(self._num_force_joints): + pos_actions.append(None) + vel_actions.append(None) + force_actions.append(self._force_actions[force_idx].item()) + + joint_state_msg.position = pos_actions + joint_state_msg.velocity = vel_actions + joint_state_msg.effort = force_actions + self.joint_state_publisher.publish(joint_state_msg) + + self.joint_state_publisher = self._ros_node.create_publisher( + JointState, f"/joint_commands", 50 + ) + self.timer = self._ros_node.create_timer(0.01, joint_action_callback) + + def _setup_robot_twist_subscriber(self): + print("Robot twist Subscriber started") + + def robot_twist_callback(msg): + self._robot_lin_vel = msg.linear + self._robot_ang_vel = msg.angular + + self.robot_twist_subscriber = self._ros_node.create_subscription( + Twist, f"/robot_twist", robot_twist_callback, 100 + ) + + def _setup_robot_pose_subscriber(self): + print("Robot pose Subscriber started") + + def robot_pose_callback(msg): + self._robot_pos = torch.tensor( + [msg.position.x, msg.position.y, msg.position.z] + ) + self._robot_quat = torch.tensor( + [ + msg.orientation.w, + msg.orientation.x, + msg.orientation.y, + msg.orientation.z, + ] + ) + + self.robot_pose_subscriber = self._ros_node.create_subscription( + Pose, f"/robot_pose", robot_pose_callback, 100 + ) + + def _setup_joint_state_subscriber(self): + print("Joint state Subscriber started") + + def joint_state_callback(msg): + pos_vals = [] + vel_vals = [] + eff_vals = [] + # We need to fill values in the order of self.dofs_idx / self.joint_names + for joint_index in range(len(list(msg.name))): + if list(msg.position)[joint_index] is not None: + pos_vals.append(list(msg.position)[joint_index]) + else: + pos_vals.append(0.0) + if list(msg.velocity)[joint_index] is not None: + vel_vals.append(list(msg.velocity)[joint_index]) + else: + vel_vals.append(0.0) + if list(msg.effort)[joint_index] is not None: + eff_vals.append(list(msg.effort)[joint_index]) + else: + eff_vals.append(0.0) + + self._pos_state = torch.tensor(pos_vals) + self._vel_state = torch.tensor(vel_vals) + self._force_state = torch.tensor(eff_vals) + + self._ros_node.create_subscription( + JointState, f"joint_states", joint_state_callback, 100 + ) + + def get_angular_velocity(self): + if self._robot_ang_vel: + return torch.tensor( + [self._robot_ang_vel.x, self._robot_ang_vel.y, self._robot_ang_vel.z] + ).unsqueeze(0) + return torch.zeros((1, 3)) + + def get_linear_velocity(self): + if self._robot_lin_vel: + return torch.tensor( + [self._robot_lin_vel.x, self._robot_lin_vel.y, self._robot_lin_vel.z] + ).unsqueeze(0) + return torch.zeros((1, 3)) + + def get_dofs_position(self): + if self._pos_state is not None: + return self._pos_state.unsqueeze(0) + return torch.zeros((1, self._num_joints)) + + def get_dofs_velocity(self): + if self._vel_state is not None: + return self._vel_state.unsqueeze(0) + return torch.zeros((1, self._num_joints)) + + def get_dofs_force(self): + if self._force_state is not None: + return self._force_state.unsqueeze(0) + return torch.zeros((1, self._num_joints))