From c81a9bd946f2fbe862f69abef893d86c776268f7 Mon Sep 17 00:00:00 2001 From: slecleach Date: Mon, 15 Sep 2025 17:33:46 -0400 Subject: [PATCH 1/4] adding bimanual setup --- judo/controller/__init__.py | 2 + judo/controller/controller.py | 4 +- judo/controller/overrides.py | 18 +- judo/models/xml/fr3_bimanual_pick.xml | 21 ++ judo/models/xml/fr3_components/fr3_bis.xml | 116 +++++++ .../xml/fr3_components/params_and_default.xml | 2 +- .../params_and_default_bimanual.xml | 105 ++++++ judo/models/xml/fr3_pick.xml | 3 +- judo/optimizers/__init__.py | 2 + judo/optimizers/overrides.py | 40 +++ judo/tasks/__init__.py | 2 + judo/tasks/fr3_bimanual_pick.py | 320 ++++++++++++++++++ 12 files changed, 629 insertions(+), 6 deletions(-) create mode 100644 judo/models/xml/fr3_bimanual_pick.xml create mode 100644 judo/models/xml/fr3_components/fr3_bis.xml create mode 100644 judo/models/xml/fr3_components/params_and_default_bimanual.xml create mode 100644 judo/tasks/fr3_bimanual_pick.py diff --git a/judo/controller/__init__.py b/judo/controller/__init__.py index 82c841d3..acd83ce8 100644 --- a/judo/controller/__init__.py +++ b/judo/controller/__init__.py @@ -6,6 +6,7 @@ set_default_caltech_leap_cube_overrides, set_default_cartpole_overrides, set_default_cylinder_push_overrides, + set_default_fr3_bimanual_pick_overrides, set_default_fr3_pick_overrides, set_default_leap_cube_down_overrides, set_default_leap_cube_overrides, @@ -21,6 +22,7 @@ set_default_caltech_leap_cube_overrides() set_default_cartpole_overrides() set_default_cylinder_push_overrides() +set_default_fr3_bimanual_pick_overrides() set_default_fr3_pick_overrides() set_default_leap_cube_overrides() set_default_leap_cube_down_overrides() diff --git a/judo/controller/controller.py b/judo/controller/controller.py index f261baf7..a35d3cc8 100644 --- a/judo/controller/controller.py +++ b/judo/controller/controller.py @@ -28,9 +28,9 @@ class ControllerConfig(OverridableConfig): """Base controller config.""" - horizon: float = 1.0 + horizon: float = 0.3 spline_order: Literal["zero", "linear", "cubic"] = "linear" - control_freq: float = 20.0 + control_freq: float = 5.0 max_opt_iters: int = 1 max_num_traces: int = 5 action_normalizer: Literal["none", "min_max", "running"] = "none" diff --git a/judo/controller/overrides.py b/judo/controller/overrides.py index aa2a14bf..acf21034 100644 --- a/judo/controller/overrides.py +++ b/judo/controller/overrides.py @@ -73,9 +73,23 @@ def set_default_fr3_pick_overrides() -> None: "fr3_pick", ControllerConfig, { - "horizon": 1.0, + "horizon": 0.3, + "spline_order": "linear", + "max_num_traces": 3, + "control_freq": 5.0, + }, + ) + + +def set_default_fr3_bimanual_pick_overrides() -> None: + """Sets the default task-specific controller config overrides for the fr3 pick task.""" + set_config_overrides( + "fr3_bimanual_pick", + ControllerConfig, + { + "horizon": 0.3, "spline_order": "linear", "max_num_traces": 3, - "control_freq": 20.0, + "control_freq": 5.0, }, ) diff --git a/judo/models/xml/fr3_bimanual_pick.xml b/judo/models/xml/fr3_bimanual_pick.xml new file mode 100644 index 00000000..76332894 --- /dev/null +++ b/judo/models/xml/fr3_bimanual_pick.xml @@ -0,0 +1,21 @@ + + + + + + + + + + + + + + + + + + + + + diff --git a/judo/models/xml/fr3_components/fr3_bis.xml b/judo/models/xml/fr3_components/fr3_bis.xml new file mode 100644 index 00000000..4e7e1ee3 --- /dev/null +++ b/judo/models/xml/fr3_components/fr3_bis.xml @@ -0,0 +1,116 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/judo/models/xml/fr3_components/params_and_default.xml b/judo/models/xml/fr3_components/params_and_default.xml index 494af3a9..8d9c4551 100644 --- a/judo/models/xml/fr3_components/params_and_default.xml +++ b/judo/models/xml/fr3_components/params_and_default.xml @@ -1,6 +1,6 @@ - diff --git a/judo/optimizers/overrides.py b/judo/optimizers/overrides.py index 4e0870d8..76ada102 100644 --- a/judo/optimizers/overrides.py +++ b/judo/optimizers/overrides.py @@ -230,8 +230,8 @@ def set_default_fr3_bimanual_pick_overrides() -> None: "fr3_bimanual_pick", PredictiveSamplingConfig, { - "num_nodes": 8, - "num_rollouts": 64, + "num_nodes": 5, + "num_rollouts": 16, "use_noise_ramp": True, "noise_ramp": 4.0, "sigma": 0.2, @@ -242,7 +242,7 @@ def set_default_fr3_bimanual_pick_overrides() -> None: CrossEntropyMethodConfig, { "num_nodes": 4, - "num_rollouts": 64, + "num_rollouts": 16, "num_elites": 3, "use_noise_ramp": True, "noise_ramp": 4.0, @@ -255,7 +255,7 @@ def set_default_fr3_bimanual_pick_overrides() -> None: MPPIConfig, { "num_nodes": 4, - "num_rollouts": 64, + "num_rollouts": 16, "use_noise_ramp": True, "noise_ramp": 4.0, "sigma": 0.01, diff --git a/judo/tasks/fr3_bimanual_pick.py b/judo/tasks/fr3_bimanual_pick.py index af1bdf72..2130bb58 100644 --- a/judo/tasks/fr3_bimanual_pick.py +++ b/judo/tasks/fr3_bimanual_pick.py @@ -12,10 +12,21 @@ from judo.tasks.base import Task, TaskConfig from judo.utils.fields import np_1d_field +# TODOS +# add trace for second arm DONE +# add collision avoidance reward DONE +# speed up the physics DONE +# remove the gripper from both arms DONE +# add a watter jug +# create a pick, lift and drop onto shelf task +# create a state machine for grasp object, use object to pull another object WONT DO + + +BOX_SIZE = 0.02 XML_PATH = str(MODEL_PATH / "xml/fr3_bimanual_pick.xml") QPOS_HOME = np.array( [ - 0.7, 0, 0.02, 1, 0, 0, 0, # object + 0.7, 0, BOX_SIZE, 1, 0, 0, 0, # object 0, -0.7854, 0.0, -2.3562, 0.0, 1.5708, 0.7854, # arm 0.04, 0.04, # gripper, equality constrained 0, -0.7854, 0.0, -2.3562, 0.0, 1.5708, 0.7854, # arm_bis @@ -75,6 +86,7 @@ class GlobalConfig: w_coll: float = 0.1 w_qvel: float = 0.005 w_open: float = 2.0 + w_collision_avoidance: float = 5.0 @slider("goal_radius", 0.005, 0.1, 0.005) @@ -127,16 +139,58 @@ def __init__(self, model_path: str = XML_PATH, sim_model_path: str | None = None arm_pos_adr = self.get_joint_position_start_index("fr3_joint1") self.arm_pos_slice = slice(arm_pos_adr, arm_pos_adr + 9) # 7 + 2 dofs for the gripper + arm_bis_pos_adr = self.get_joint_position_start_index("_fr3_joint1") + self.arm_bis_pos_slice = slice(arm_bis_pos_adr, arm_bis_pos_adr + 9) # 7 + 2 dofs for the gripper + # sensors self.left_finger_obj_adr = self.get_sensor_start_index("left_finger_obj") self.right_finger_obj_adr = self.get_sensor_start_index("right_finger_obj") self.left_finger_table_adr = self.get_sensor_start_index("left_finger_table") self.right_finger_table_adr = self.get_sensor_start_index("right_finger_table") self.grasp_site_adr = self.get_sensor_start_index("trace_grasp_site") - self.obj_table_adr = self.get_sensor_start_index("obj_table") self.ee_z_adr = self.get_sensor_start_index("ee_z") self.ee_z_slice = slice(self.ee_z_adr, self.ee_z_adr + 3) + self.left_finger_bis_obj_adr = self.get_sensor_start_index("_left_finger_obj") + self.right_finger_bis_obj_adr = self.get_sensor_start_index("_right_finger_obj") + self.left_finger_bis_table_adr = self.get_sensor_start_index("_left_finger_table") + self.right_finger_bis_table_adr = self.get_sensor_start_index("_right_finger_table") + self.grasp_site_bis_adr = self.get_sensor_start_index("_trace_grasp_site") + self.ee_z_bis_adr = self.get_sensor_start_index("_ee_z") + self.ee_z_bis_slice = slice(self.ee_z_bis_adr, self.ee_z_bis_adr + 3) + + self.obj_table_adr = self.get_sensor_start_index("obj_table") + + # collision_avoidance sensors + self.arm_maximal_pos_slice = [] + self.arm_bis_maximal_pos_slice = [] + for sensor_name in [ + "fr3_link1_pos", + "fr3_link2_pos", + "fr3_link3_pos", + "fr3_link4_pos", + "fr3_link5_pos", + "fr3_link6_pos", + "fr3_link7_pos", + ]: + idx = self.get_sensor_start_index(sensor_name) + np.arange(3) + self.arm_maximal_pos_slice.append(idx) + + for sensor_name in [ + "_fr3_link1_pos", + "_fr3_link2_pos", + "_fr3_link3_pos", + "_fr3_link4_pos", + "_fr3_link5_pos", + "_fr3_link6_pos", + "_fr3_link7_pos", + ]: + idx = self.get_sensor_start_index(sensor_name) + np.arange(3) + self.arm_bis_maximal_pos_slice.append(idx) + + self.arm_maximal_pos_slice = np.concatenate(self.arm_maximal_pos_slice) + self.arm_bis_maximal_pos_slice = np.concatenate(self.arm_bis_maximal_pos_slice) + # metadata that stores the current phase of the task self._data = mujoco.MjData(self.model) # used for computing hypothetical sensor data self.phase = Phase.LIFT # default phase @@ -161,7 +215,17 @@ def in_goal_xy(self, curr_state: np.ndarray, config: FR3BimanualPickConfig) -> n def check_sensor_dists( self, sensors: np.ndarray, - pair: Literal["left_finger_obj", "right_finger_obj", "left_finger_table", "right_finger_table", "obj_table"], + pair: Literal[ + "left_finger_obj", + "right_finger_obj", + "left_finger_table", + "right_finger_table", + "left_finger_obj_bis", + "right_finger_obj_bis", + "left_finger_table_bis", + "right_finger_table_bis", + "obj_table", + ], ) -> np.ndarray: """Computes the distance between a specified pair of bodies. @@ -180,11 +244,19 @@ def check_sensor_dists( i = self.left_finger_table_adr elif pair == "right_finger_table": i = self.right_finger_table_adr + elif pair == "left_finger_obj_bis": + i = self.left_finger_bis_obj_adr + elif pair == "right_finger_obj_bis": + i = self.right_finger_bis_obj_adr + elif pair == "left_finger_table_bis": + i = self.left_finger_bis_table_adr + elif pair == "right_finger_table_bis": + i = self.right_finger_bis_table_adr elif pair == "obj_table": i = self.obj_table_adr else: raise ValueError( - f"Invalid pair: {pair}. Must be one of 'left_finger_obj', 'right_finger_obj', or 'obj_table'." + f"Invalid pair: {pair}. Must be one of 'left_finger_obj', 'right_finger_obj', 'left_finger_table', 'right_finger_table', 'left_finger_obj_bis', 'right_finger_obj_bis', 'left_finger_table_bis', 'right_finger_table_bis', 'obj_table'." ) dist = sensors[:, :, i] return dist @@ -204,7 +276,7 @@ def pre_rollout(self, curr_state: np.ndarray, config: FR3BimanualPickConfig) -> # check whether the phase is MOVE # obj_in_air = curr_sensor[self.obj_table_adr] > 0 # object is not touching the table - obj_in_air = curr_state[self.obj_pos_adr + 2] > 0.02 + 1e-3 # object z position is above the table + obj_in_air = curr_state[self.obj_pos_adr + 2] > BOX_SIZE + 1e-3 # object z position is above the table if obj_in_air: phase = Phase.MOVE # if the object is in the air, we are in lift phase @@ -218,10 +290,11 @@ def pre_rollout(self, curr_state: np.ndarray, config: FR3BimanualPickConfig) -> # if in_goal_xy and obj_table_dist <= 0: # phase = Phase.HOMING obj_z_pos = curr_state[self.obj_pos_adr + 2] # z position of the object - if in_goal_xy and obj_z_pos <= 0.02 + 1e-3: # the cube is 4cm wide and we allow a tolerance + if in_goal_xy and obj_z_pos <= BOX_SIZE + 1e-3: # the cube is 4cm wide and we allow a tolerance phase = Phase.HOMING self.phase = phase + print(f"Phase: {self.phase}") def reward( self, @@ -251,6 +324,23 @@ def reward( grasp_site_pos = sensors[..., self.grasp_site_adr : self.grasp_site_adr + 3] # (num_rollouts, T, 3) ee_z_axis = sensors[..., self.ee_z_slice] # (num_rollouts, T, 3) + # collision avoidance + arm_maximal_pos = sensors[..., self.arm_maximal_pos_slice] # (num_rollouts, T, 7*3) + arm_bis_maximal_pos = sensors[..., self.arm_bis_maximal_pos_slice] # (num_rollouts, T, 7*3) + # reshape to (num_rollouts, T, 7, 3) + arm_maximal_pos = arm_maximal_pos.reshape(arm_maximal_pos.shape[0], arm_maximal_pos.shape[1], 7, 3) + arm_bis_maximal_pos = arm_bis_maximal_pos.reshape( + arm_bis_maximal_pos.shape[0], arm_bis_maximal_pos.shape[1], 7, 3 + ) + # tile and compute pairwise distances + arm_maximal_pos = np.expand_dims(arm_maximal_pos, axis=2) # (num_rollouts, T, 1, 7, 3) + arm_bis_maximal_pos = np.expand_dims(arm_bis_maximal_pos, axis=3) # (num_rollouts, T, 7, 1, 3) + # compute the distance between the two arms + arm_dist = np.linalg.norm(arm_maximal_pos - arm_bis_maximal_pos, axis=-1) # (num_rollouts, T, 7, 7) + arm_dist = np.mean(arm_dist, axis=(1, 2, 3)) # (num_rollouts,) + # compute the reward for collision avoidance + rew_collision_avoidance = -np.exp(-arm_dist) # (num_rollouts,) + # querying states obj_pos = states[..., self.obj_pos_slice] # (num_rollouts, T, 3) arm_pos = states[..., self.arm_pos_slice] # (num_rollouts, T, 9) @@ -302,6 +392,7 @@ def reward( w_coll = config.global_weights.w_coll w_qvel = config.global_weights.w_qvel w_open = config.global_weights.w_open + w_collision_avoidance = config.global_weights.w_collision_avoidance rew_upright = -np.linalg.norm(ee_z_axis - np.array([[[0.0, 0.0, -1.0]]]), axis=-1).sum(axis=-1) rew_coll = (1 - hand_touching).sum(axis=-1) # (num_rollouts,) @@ -309,7 +400,13 @@ def reward( rew_qvel = -(time_decay * qvel_norm).sum(axis=-1) rew_open = -((gripper_pos - 0.04) ** 2).sum(axis=-1) # encourage the gripper to be open - rewards += w_upright * rew_upright + w_coll * rew_coll + w_qvel * rew_qvel + w_open * rew_open + rewards += ( + w_upright * rew_upright + + w_coll * rew_coll + + w_qvel * rew_qvel + + w_open * rew_open + + w_collision_avoidance * rew_collision_avoidance + ) return rewards def reset(self) -> None: From fb6cdce4f3f5682cdaf63dcc6bfea485499df376 Mon Sep 17 00:00:00 2001 From: slecleach Date: Mon, 15 Sep 2025 23:14:08 -0400 Subject: [PATCH 3/4] trying homing --- judo/models/xml/fr3_bimanual_pick.xml | 10 ++--- judo/tasks/fr3_bimanual_pick.py | 58 +++++++++++++++++---------- 2 files changed, 41 insertions(+), 27 deletions(-) diff --git a/judo/models/xml/fr3_bimanual_pick.xml b/judo/models/xml/fr3_bimanual_pick.xml index 4f6e88e0..21004a53 100644 --- a/judo/models/xml/fr3_bimanual_pick.xml +++ b/judo/models/xml/fr3_bimanual_pick.xml @@ -8,12 +8,12 @@ + - - - - - + + + + diff --git a/judo/tasks/fr3_bimanual_pick.py b/judo/tasks/fr3_bimanual_pick.py index 2130bb58..66e8e55c 100644 --- a/judo/tasks/fr3_bimanual_pick.py +++ b/judo/tasks/fr3_bimanual_pick.py @@ -26,7 +26,7 @@ XML_PATH = str(MODEL_PATH / "xml/fr3_bimanual_pick.xml") QPOS_HOME = np.array( [ - 0.7, 0, BOX_SIZE, 1, 0, 0, 0, # object + 0.7, 0, 0.2, 1, 0, 0, 0, # object 0, -0.7854, 0.0, -2.3562, 0.0, 1.5708, 0.7854, # arm 0.04, 0.04, # gripper, equality constrained 0, -0.7854, 0.0, -2.3562, 0.0, 1.5708, 0.7854, # arm_bis @@ -141,6 +141,8 @@ def __init__(self, model_path: str = XML_PATH, sim_model_path: str | None = None arm_bis_pos_adr = self.get_joint_position_start_index("_fr3_joint1") self.arm_bis_pos_slice = slice(arm_bis_pos_adr, arm_bis_pos_adr + 9) # 7 + 2 dofs for the gripper + print("arm_pos_slice", self.arm_pos_slice) + print("arm_bis_pos_slice", self.arm_bis_pos_slice) # sensors self.left_finger_obj_adr = self.get_sensor_start_index("left_finger_obj") @@ -272,28 +274,29 @@ def pre_rollout(self, curr_state: np.ndarray, config: FR3BimanualPickConfig) -> # check the object z position # curr_sensor = self._data.sensordata # (total_sensor_dim,) - phase = Phase.LIFT # default phase - - # check whether the phase is MOVE - # obj_in_air = curr_sensor[self.obj_table_adr] > 0 # object is not touching the table - obj_in_air = curr_state[self.obj_pos_adr + 2] > BOX_SIZE + 1e-3 # object z position is above the table - if obj_in_air: - phase = Phase.MOVE # if the object is in the air, we are in lift phase - - # check whether the phase is PLACE - in_goal_xy = self.in_goal_xy(curr_state, config) - if in_goal_xy and obj_in_air: - phase = Phase.PLACE # if the object is in the goal xy, we are in place phase - - # check whether the phase is HOMING - # obj_table_dist = curr_sensor[self.obj_table_adr] - # if in_goal_xy and obj_table_dist <= 0: + # phase = Phase.LIFT # default phase + + # # check whether the phase is MOVE + # # obj_in_air = curr_sensor[self.obj_table_adr] > 0 # object is not touching the table + # obj_in_air = curr_state[self.obj_pos_adr + 2] > BOX_SIZE + 1e-3 # object z position is above the table + # if obj_in_air: + # phase = Phase.MOVE # if the object is in the air, we are in lift phase + + # # check whether the phase is PLACE + # in_goal_xy = self.in_goal_xy(curr_state, config) + # if in_goal_xy and obj_in_air: + # phase = Phase.PLACE # if the object is in the goal xy, we are in place phase + + # # check whether the phase is HOMING + # # obj_table_dist = curr_sensor[self.obj_table_adr] + # # if in_goal_xy and obj_table_dist <= 0: + # # phase = Phase.HOMING + # obj_z_pos = curr_state[self.obj_pos_adr + 2] # z position of the object + # if in_goal_xy and obj_z_pos <= BOX_SIZE + 1e-3: # the cube is 4cm wide and we allow a tolerance # phase = Phase.HOMING - obj_z_pos = curr_state[self.obj_pos_adr + 2] # z position of the object - if in_goal_xy and obj_z_pos <= BOX_SIZE + 1e-3: # the cube is 4cm wide and we allow a tolerance - phase = Phase.HOMING - self.phase = phase + # self.phase = phase + self.phase = Phase.HOMING print(f"Phase: {self.phase}") def reward( @@ -344,6 +347,7 @@ def reward( # querying states obj_pos = states[..., self.obj_pos_slice] # (num_rollouts, T, 3) arm_pos = states[..., self.arm_pos_slice] # (num_rollouts, T, 9) + arm_bis_pos = states[..., self.arm_bis_pos_slice] # (num_rollouts, T, 9) xy_pos = states[..., :2] # (num_rollouts, T, 2) z_obj = states[..., self.obj_pos_adr + 2] # (num_rollouts, T) qvel = states[..., self.model.nq : self.model.nq + self.model.nv] # (num_rollouts, T, nv) @@ -352,10 +356,17 @@ def reward( # distances and errors q_arm_goal = QPOS_HOME[self.arm_pos_slice] # (9,) + q_arm_bis_goal = QPOS_HOME[self.arm_bis_pos_slice] # (9,) + grasp_dist = ((grasp_site_pos - obj_pos) ** 2).sum(-1) # (num_rollouts, T) pick_height_err = (z_obj - config.pick_height) ** 2 # (num_rollouts, T) obj_goal_pos_dist = np.linalg.norm(xy_pos - config.goal_pos, axis=-1) # (num_rollouts, T) - home_dist = np.linalg.norm(arm_pos - q_arm_goal, axis=-1) # (num_rollouts, T) + + home_dist = 100 * np.linalg.norm(arm_pos - q_arm_goal, axis=-1) + 100 * np.linalg.norm( + arm_bis_pos - q_arm_bis_goal, axis=-1 + ) # (num_rollouts, T) + # home_dist = 100 * np.linalg.norm(arm_pos - q_arm_goal, axis=-1) + # home_dist = 100 * np.linalg.norm(arm_bis_pos - q_arm_bis_goal, axis=-1) # (num_rollouts, T) # contact checks left_finger_touching = left_finger_table_dist <= 0.0 # (num_rollouts, T) @@ -407,6 +418,9 @@ def reward( + w_open * rew_open + w_collision_avoidance * rew_collision_avoidance ) + print("rewards", rewards) + print("homing", home_dist.sum(axis=-1)) + print("rewards + homing", rewards + home_dist.sum(axis=-1)) return rewards def reset(self) -> None: From 3531de01c32005a27ce2e6c4ec04df612f632274 Mon Sep 17 00:00:00 2001 From: slecleach Date: Tue, 16 Sep 2025 12:57:34 -0400 Subject: [PATCH 4/4] first working version --- judo/controller/overrides.py | 4 +- judo/models/xml/fr3_bimanual_pick.xml | 18 +- judo/models/xml/fr3_components/fr3.xml | 4 +- judo/models/xml/fr3_components/fr3_bis.xml | 4 +- .../params_and_default_bimanual.xml | 54 ++--- judo/tasks/fr3_bimanual_pick.py | 194 ++++++++++-------- 6 files changed, 157 insertions(+), 121 deletions(-) diff --git a/judo/controller/overrides.py b/judo/controller/overrides.py index 500906fc..6380ae58 100644 --- a/judo/controller/overrides.py +++ b/judo/controller/overrides.py @@ -87,9 +87,9 @@ def set_default_fr3_bimanual_pick_overrides() -> None: "fr3_bimanual_pick", ControllerConfig, { - "horizon": 1.0, + "horizon": 1.3, "spline_order": "linear", "max_num_traces": 3, - "control_freq": 20.0, + "control_freq": 40.0, }, ) diff --git a/judo/models/xml/fr3_bimanual_pick.xml b/judo/models/xml/fr3_bimanual_pick.xml index 21004a53..4edb931c 100644 --- a/judo/models/xml/fr3_bimanual_pick.xml +++ b/judo/models/xml/fr3_bimanual_pick.xml @@ -4,16 +4,22 @@ - + + + + + + + - + - - - - + + + + diff --git a/judo/models/xml/fr3_components/fr3.xml b/judo/models/xml/fr3_components/fr3.xml index d8ae83de..8b0634d9 100644 --- a/judo/models/xml/fr3_components/fr3.xml +++ b/judo/models/xml/fr3_components/fr3.xml @@ -79,7 +79,7 @@ - + diff --git a/judo/models/xml/fr3_components/fr3_bis.xml b/judo/models/xml/fr3_components/fr3_bis.xml index 4e7e1ee3..af63a573 100644 --- a/judo/models/xml/fr3_components/fr3_bis.xml +++ b/judo/models/xml/fr3_components/fr3_bis.xml @@ -79,7 +79,7 @@ - + diff --git a/judo/models/xml/fr3_components/params_and_default_bimanual.xml b/judo/models/xml/fr3_components/params_and_default_bimanual.xml index 93fe76d1..a033216a 100644 --- a/judo/models/xml/fr3_components/params_and_default_bimanual.xml +++ b/judo/models/xml/fr3_components/params_and_default_bimanual.xml @@ -19,7 +19,7 @@ - + @@ -40,43 +40,43 @@ - + - - - - - - - - + + + + + + + + - - - - - - - - + + + + + + + + - - + - - + @@ -84,12 +84,12 @@ - - + - - + diff --git a/judo/tasks/fr3_bimanual_pick.py b/judo/tasks/fr3_bimanual_pick.py index 66e8e55c..8937f273 100644 --- a/judo/tasks/fr3_bimanual_pick.py +++ b/judo/tasks/fr3_bimanual_pick.py @@ -17,20 +17,33 @@ # add collision avoidance reward DONE # speed up the physics DONE # remove the gripper from both arms DONE -# add a watter jug +# add a watter jug DONE # create a pick, lift and drop onto shelf task # create a state machine for grasp object, use object to pull another object WONT DO - +# fix friction DONE +# fix shelf DONE +# add proper water jug +# add stay in place term for the object during the homing phase BOX_SIZE = 0.02 +LOW_OBJ_HEIGHT = 0.125 XML_PATH = str(MODEL_PATH / "xml/fr3_bimanual_pick.xml") QPOS_HOME = np.array( [ - 0.7, 0, 0.2, 1, 0, 0, 0, # object - 0, -0.7854, 0.0, -2.3562, 0.0, 1.5708, 0.7854, # arm - 0.04, 0.04, # gripper, equality constrained - 0, -0.7854, 0.0, -2.3562, 0.0, 1.5708, 0.7854, # arm_bis - 0.04, 0.04, # gripper_bis, equality constrained + 0.4, 0, 0.2, 1, 0, 0, 0, # object + -0.5, 0.2, 0.0, -2.3562, 0.0, 1.5708, 0.7854, # arm + # 0.04, 0.04, # gripper, equality constrained + 0.5, 0.2, 0.0, -2.3562, 0.0, 1.5708, 0.7854, # arm_bis + # 0.04, 0.04, # gripper_bis, equality constrained + ] +) # fmt: skip +QPOS_TERMINAL = np.array( + [ + 0.4, 0, 0.2, 1, 0, 0, 0, # object + 0, -np.pi/2, 0.0, -2.3562, 0.0, 1.5708, 0.7854, # arm + # 0.04, 0.04, # gripper, equality constrained + 0, -np.pi/2, 0.0, -2.3562, 0.0, 1.5708, 0.7854, # arm_bis + # 0.04, 0.04, # gripper_bis, equality constrained ] ) # fmt: skip @@ -60,8 +73,8 @@ class LiftConfig: class MoveConfig: """Reward configuration for the move phase of the FR3 pick task.""" - w_move_goal: float = 1.0 - w_move_close: float = 10.0 + w_move_goal: float = 5.0 + w_move_close: float = 5.0 @slider("w_place_table", 0.0, 10.0, 0.01) @@ -86,11 +99,12 @@ class GlobalConfig: w_coll: float = 0.1 w_qvel: float = 0.005 w_open: float = 2.0 - w_collision_avoidance: float = 5.0 + w_collision_avoidance: float = 125.0 -@slider("goal_radius", 0.005, 0.1, 0.005) +@slider("goal_radius", 0.0, 0.75, 0.01) @slider("pick_height", 0.0, 1.0, 0.01) +@slider("pick_height_tolerance", 0.0, 1.0, 0.01) @dataclass class FR3BimanualPickConfig(TaskConfig): """Reward configuration for FR3 pick task.""" @@ -102,17 +116,18 @@ class FR3BimanualPickConfig(TaskConfig): global_weights: GlobalConfig = field(default_factory=GlobalConfig) goal_pos: np.ndarray = np_1d_field( - np.array([0.6, 0.4]), - names=["x", "y"], - mins=[0.4, -1.0], - maxs=[1.0, 1.0], - steps=[0.01, 0.01], + np.array([1.1, 0.0, 0.55]), + names=["x", "y", "z"], + mins=[0.4, -1.0, 0.0], + maxs=[1.5, 1.0, 1.0], + steps=[0.01, 0.01, 0.01], vis_name="goal_position", - xyz_vis_indices=[0, 1, None], + xyz_vis_indices=[0, 1, 2], xyz_vis_defaults=[0.0, 0.0, 0.0], ) - goal_radius: float = 0.05 - pick_height: float = 0.3 + goal_radius: float = 0.25 + pick_height: float = 0.5 + pick_height_tolerance: float = 0.1 class FR3BimanualPick(Task[FR3BimanualPickConfig]): @@ -121,9 +136,6 @@ class FR3BimanualPick(Task[FR3BimanualPickConfig]): def __init__(self, model_path: str = XML_PATH, sim_model_path: str | None = None) -> None: """Initializes the LEAP cube rotation task.""" super().__init__(model_path, sim_model_path=sim_model_path) - self.reset_command = np.array( - [0, 0, 0, -1.57079, 0, 1.57079, -0.7853, 0.0, 0, 0, 0, -1.57079, 0, 1.57079, -0.7853, 0.0] - ) # object indices self.obj_pos_adr = self.get_joint_position_start_index("object_joint") @@ -137,26 +149,29 @@ def __init__(self, model_path: str = XML_PATH, sim_model_path: str | None = None # robot indices arm_pos_adr = self.get_joint_position_start_index("fr3_joint1") - self.arm_pos_slice = slice(arm_pos_adr, arm_pos_adr + 9) # 7 + 2 dofs for the gripper + # self.arm_pos_slice = slice(arm_pos_adr, arm_pos_adr + 9) # 7 + 2 dofs for the gripper + self.arm_pos_slice = slice(arm_pos_adr, arm_pos_adr + 7) # 7 + 2 dofs for the gripper arm_bis_pos_adr = self.get_joint_position_start_index("_fr3_joint1") - self.arm_bis_pos_slice = slice(arm_bis_pos_adr, arm_bis_pos_adr + 9) # 7 + 2 dofs for the gripper - print("arm_pos_slice", self.arm_pos_slice) - print("arm_bis_pos_slice", self.arm_bis_pos_slice) + # self.arm_bis_pos_slice = slice(arm_bis_pos_adr, arm_bis_pos_adr + 9) # 7 + 2 dofs for the gripper + self.arm_bis_pos_slice = slice(arm_bis_pos_adr, arm_bis_pos_adr + 7) # 7 + 2 dofs for the gripper + self.arms_pos_slice = slice(arm_pos_adr, arm_pos_adr + 14) + + self.reset_command = QPOS_HOME[self.arms_pos_slice] # sensors - self.left_finger_obj_adr = self.get_sensor_start_index("left_finger_obj") - self.right_finger_obj_adr = self.get_sensor_start_index("right_finger_obj") - self.left_finger_table_adr = self.get_sensor_start_index("left_finger_table") - self.right_finger_table_adr = self.get_sensor_start_index("right_finger_table") + # self.left_finger_obj_adr = self.get_sensor_start_index("left_finger_obj") + # self.right_finger_obj_adr = self.get_sensor_start_index("right_finger_obj") + # self.left_finger_table_adr = self.get_sensor_start_index("left_finger_table") + # self.right_finger_table_adr = self.get_sensor_start_index("right_finger_table") self.grasp_site_adr = self.get_sensor_start_index("trace_grasp_site") self.ee_z_adr = self.get_sensor_start_index("ee_z") self.ee_z_slice = slice(self.ee_z_adr, self.ee_z_adr + 3) - self.left_finger_bis_obj_adr = self.get_sensor_start_index("_left_finger_obj") - self.right_finger_bis_obj_adr = self.get_sensor_start_index("_right_finger_obj") - self.left_finger_bis_table_adr = self.get_sensor_start_index("_left_finger_table") - self.right_finger_bis_table_adr = self.get_sensor_start_index("_right_finger_table") + # self.left_finger_bis_obj_adr = self.get_sensor_start_index("_left_finger_obj") + # self.right_finger_bis_obj_adr = self.get_sensor_start_index("_right_finger_obj") + # self.left_finger_bis_table_adr = self.get_sensor_start_index("_left_finger_table") + # self.right_finger_bis_table_adr = self.get_sensor_start_index("_right_finger_table") self.grasp_site_bis_adr = self.get_sensor_start_index("_trace_grasp_site") self.ee_z_bis_adr = self.get_sensor_start_index("_ee_z") self.ee_z_bis_slice = slice(self.ee_z_bis_adr, self.ee_z_bis_adr + 3) @@ -274,29 +289,37 @@ def pre_rollout(self, curr_state: np.ndarray, config: FR3BimanualPickConfig) -> # check the object z position # curr_sensor = self._data.sensordata # (total_sensor_dim,) - # phase = Phase.LIFT # default phase + phase = Phase.LIFT # default phase - # # check whether the phase is MOVE - # # obj_in_air = curr_sensor[self.obj_table_adr] > 0 # object is not touching the table - # obj_in_air = curr_state[self.obj_pos_adr + 2] > BOX_SIZE + 1e-3 # object z position is above the table - # if obj_in_air: - # phase = Phase.MOVE # if the object is in the air, we are in lift phase + # check whether the phase is MOVE + # obj_in_air = curr_sensor[self.obj_table_adr] > 0 # object is not touching the table + obj_in_air = ( + curr_state[self.obj_pos_adr + 2] - config.pick_height >= -config.pick_height_tolerance + ) # object z position is close to the pick height + if obj_in_air: + phase = Phase.MOVE # if the object is in the air, we are in lift phase # # check whether the phase is PLACE # in_goal_xy = self.in_goal_xy(curr_state, config) # if in_goal_xy and obj_in_air: # phase = Phase.PLACE # if the object is in the goal xy, we are in place phase - # # check whether the phase is HOMING - # # obj_table_dist = curr_sensor[self.obj_table_adr] + # check whether the phase is HOMING + # obj_table_dist = curr_sensor[self.obj_table_adr] # # if in_goal_xy and obj_table_dist <= 0: # # phase = Phase.HOMING # obj_z_pos = curr_state[self.obj_pos_adr + 2] # z position of the object # if in_goal_xy and obj_z_pos <= BOX_SIZE + 1e-3: # the cube is 4cm wide and we allow a tolerance - # phase = Phase.HOMING - # self.phase = phase - self.phase = Phase.HOMING + obj_pos = curr_state[self.obj_pos_slice] + # print("obj_pos", obj_pos) + # print("config.goal_pos", config.goal_pos) + # print("config.goal_radius", config.goal_radius) + # print("norm", np.linalg.norm(obj_pos - config.goal_pos)) + if np.linalg.norm(obj_pos - config.goal_pos) <= config.goal_radius: + phase = Phase.HOMING + + self.phase = phase print(f"Phase: {self.phase}") def reward( @@ -321,11 +344,11 @@ def reward( * Qvel: The robot arm is not moving too fast. """ # querying sensors - left_finger_table_dist = self.check_sensor_dists(sensors, "left_finger_table") # noqa: F841 - right_finger_table_dist = self.check_sensor_dists(sensors, "right_finger_table") # noqa: F841 - obj_table_dist = self.check_sensor_dists(sensors, "obj_table") # noqa: F841 - grasp_site_pos = sensors[..., self.grasp_site_adr : self.grasp_site_adr + 3] # (num_rollouts, T, 3) - ee_z_axis = sensors[..., self.ee_z_slice] # (num_rollouts, T, 3) + # left_finger_table_dist = self.check_sensor_dists(sensors, "left_finger_table") # noqa: F841 + # right_finger_table_dist = self.check_sensor_dists(sensors, "right_finger_table") # noqa: F841 + # obj_table_dist = self.check_sensor_dists(sensors, "obj_table") # noqa: F841 + # grasp_site_pos = sensors[..., self.grasp_site_adr : self.grasp_site_adr + 3] # (num_rollouts, T, 3) + # ee_z_axis = sensors[..., self.ee_z_slice] # (num_rollouts, T, 3) # collision avoidance arm_maximal_pos = sensors[..., self.arm_maximal_pos_slice] # (num_rollouts, T, 7*3) @@ -336,10 +359,14 @@ def reward( arm_bis_maximal_pos.shape[0], arm_bis_maximal_pos.shape[1], 7, 3 ) # tile and compute pairwise distances - arm_maximal_pos = np.expand_dims(arm_maximal_pos, axis=2) # (num_rollouts, T, 1, 7, 3) - arm_bis_maximal_pos = np.expand_dims(arm_bis_maximal_pos, axis=3) # (num_rollouts, T, 7, 1, 3) + tiled_arm_maximal_pos = np.tile( + np.expand_dims(arm_maximal_pos, axis=2), (1, 1, 7, 1, 1) + ) # (num_rollouts, T, 1, 7, 3) + tiled_arm_bis_maximal_pos = np.tile( + np.expand_dims(arm_bis_maximal_pos, axis=3), (1, 1, 1, 7, 1) + ) # (num_rollouts, T, 7, 1, 3) # compute the distance between the two arms - arm_dist = np.linalg.norm(arm_maximal_pos - arm_bis_maximal_pos, axis=-1) # (num_rollouts, T, 7, 7) + arm_dist = np.linalg.norm(tiled_arm_maximal_pos - tiled_arm_bis_maximal_pos, axis=-1) # (num_rollouts, T, 7, 7) arm_dist = np.mean(arm_dist, axis=(1, 2, 3)) # (num_rollouts,) # compute the reward for collision avoidance rew_collision_avoidance = -np.exp(-arm_dist) # (num_rollouts,) @@ -348,79 +375,82 @@ def reward( obj_pos = states[..., self.obj_pos_slice] # (num_rollouts, T, 3) arm_pos = states[..., self.arm_pos_slice] # (num_rollouts, T, 9) arm_bis_pos = states[..., self.arm_bis_pos_slice] # (num_rollouts, T, 9) - xy_pos = states[..., :2] # (num_rollouts, T, 2) + # xy_pos = states[..., :2] # (num_rollouts, T, 2) z_obj = states[..., self.obj_pos_adr + 2] # (num_rollouts, T) qvel = states[..., self.model.nq : self.model.nq + self.model.nv] # (num_rollouts, T, nv) qvel_norm = np.linalg.norm(qvel, axis=-1) # (num_rollouts, T) - gripper_pos = arm_pos[..., -1] # (num_rollouts, T) + # gripper_pos = arm_pos[..., -1] # (num_rollouts, T) # distances and errors - q_arm_goal = QPOS_HOME[self.arm_pos_slice] # (9,) - q_arm_bis_goal = QPOS_HOME[self.arm_bis_pos_slice] # (9,) - - grasp_dist = ((grasp_site_pos - obj_pos) ** 2).sum(-1) # (num_rollouts, T) + q_arm_goal = QPOS_TERMINAL[self.arm_pos_slice] # (9,) + q_arm_bis_goal = QPOS_TERMINAL[self.arm_bis_pos_slice] # (9,) + + # grasp_dist = ((grasp_site_pos - obj_pos) ** 2).sum(-1) # (num_rollouts, T) + obj_lower_pos = obj_pos.copy() - np.array([0, 0, LOW_OBJ_HEIGHT]) + grasp_dist = np.linalg.norm(arm_maximal_pos[:, :, -1, :] - obj_lower_pos, axis=-1) # (num_rollouts, T) + grasp_dist_bis = np.linalg.norm(arm_bis_maximal_pos[:, :, -1, :] - obj_lower_pos, axis=-1) # (num_rollouts, T) + grasp_dist = grasp_dist + grasp_dist_bis pick_height_err = (z_obj - config.pick_height) ** 2 # (num_rollouts, T) - obj_goal_pos_dist = np.linalg.norm(xy_pos - config.goal_pos, axis=-1) # (num_rollouts, T) + obj_goal_pos_dist = np.linalg.norm(obj_pos - config.goal_pos, axis=-1) # (num_rollouts, T) - home_dist = 100 * np.linalg.norm(arm_pos - q_arm_goal, axis=-1) + 100 * np.linalg.norm( + home_dist = np.linalg.norm(arm_pos - q_arm_goal, axis=-1) + np.linalg.norm( arm_bis_pos - q_arm_bis_goal, axis=-1 ) # (num_rollouts, T) - # home_dist = 100 * np.linalg.norm(arm_pos - q_arm_goal, axis=-1) - # home_dist = 100 * np.linalg.norm(arm_bis_pos - q_arm_bis_goal, axis=-1) # (num_rollouts, T) # contact checks - left_finger_touching = left_finger_table_dist <= 0.0 # (num_rollouts, T) - right_finger_touching = right_finger_table_dist <= 0.0 # (num_rollouts, T) - hand_touching = left_finger_touching | right_finger_touching + # left_finger_touching = left_finger_table_dist <= 0.0 # (num_rollouts, T) + # right_finger_touching = right_finger_table_dist <= 0.0 # (num_rollouts, T) + # hand_touching = left_finger_touching | right_finger_touching # lift rewards if self.phase == Phase.LIFT: w_lift_close = config.lift_weights.w_lift_close w_lift_height = config.lift_weights.w_lift_height rewards = -(w_lift_close * grasp_dist + w_lift_height * pick_height_err).sum(axis=-1) + pass # move rewards elif self.phase == Phase.MOVE: w_move_goal = config.move_weights.w_move_goal w_move_close = config.move_weights.w_move_close rewards = -(w_move_goal * obj_goal_pos_dist + w_move_close * grasp_dist).sum(axis=-1) + pass # place rewards elif self.phase == Phase.PLACE: - w_place_table = config.place_weights.w_place_table - w_place_goal = config.place_weights.w_place_goal - rewards = -(+w_place_table * obj_table_dist + w_place_goal * obj_goal_pos_dist).sum(axis=-1) + # w_place_table = config.place_weights.w_place_table + # w_place_goal = config.place_weights.w_place_goal + # rewards = -(+w_place_table * obj_table_dist + w_place_goal * obj_goal_pos_dist).sum(axis=-1) + pass # homing rewards elif self.phase == Phase.HOMING: rewards = -home_dist.sum(axis=-1) + pass else: # should never happen raise ValueError(f"Invalid phase: {self.phase}. Must be one of {list(Phase)}.") # global rewards - w_upright = config.global_weights.w_upright - w_coll = config.global_weights.w_coll + # w_upright = config.global_weights.w_upright + # w_coll = config.global_weights.w_coll w_qvel = config.global_weights.w_qvel - w_open = config.global_weights.w_open + # w_open = config.global_weights.w_open w_collision_avoidance = config.global_weights.w_collision_avoidance - rew_upright = -np.linalg.norm(ee_z_axis - np.array([[[0.0, 0.0, -1.0]]]), axis=-1).sum(axis=-1) - rew_coll = (1 - hand_touching).sum(axis=-1) # (num_rollouts,) + # rew_upright = -np.linalg.norm(ee_z_axis - np.array([[[0.0, 0.0, -1.0]]]), axis=-1).sum(axis=-1) + # rew_coll = (1 - hand_touching).sum(axis=-1) # (num_rollouts,) time_decay = np.linspace(1.0, 0.0, states.shape[1]) # decay the velocity penalty over time rew_qvel = -(time_decay * qvel_norm).sum(axis=-1) - rew_open = -((gripper_pos - 0.04) ** 2).sum(axis=-1) # encourage the gripper to be open + # rew_open = -((gripper_pos - 0.04) ** 2).sum(axis=-1) # encourage the gripper to be open rewards += ( - w_upright * rew_upright - + w_coll * rew_coll - + w_qvel * rew_qvel - + w_open * rew_open + # w_upright * rew_upright + # + w_coll * rew_coll + +w_qvel * rew_qvel + # + w_open * rew_open + w_collision_avoidance * rew_collision_avoidance ) - print("rewards", rewards) - print("homing", home_dist.sum(axis=-1)) - print("rewards + homing", rewards + home_dist.sum(axis=-1)) return rewards def reset(self) -> None: