From 048f8af440ff10f8dd650d78291cd4068fb23ca0 Mon Sep 17 00:00:00 2001 From: Vaibhav Saxena Date: Wed, 21 Jan 2026 14:35:47 -0500 Subject: [PATCH] add absolute action to generated trajectories --- mimicgen/datagen/data_generator.py | 4 ++++ mimicgen/datagen/waypoint.py | 30 ++++++++++++++++++++++++++++ mimicgen/scripts/generate_dataset.py | 2 ++ mimicgen/utils/file_utils.py | 6 ++++++ 4 files changed, 42 insertions(+) diff --git a/mimicgen/datagen/data_generator.py b/mimicgen/datagen/data_generator.py index bc07081..24b721f 100644 --- a/mimicgen/datagen/data_generator.py +++ b/mimicgen/datagen/data_generator.py @@ -254,6 +254,7 @@ def generate( generated_obs = [] generated_datagen_infos = [] generated_actions = [] + generated_actions_abs = [] generated_success = False generated_src_demo_inds = [] # store selected src demo ind for each subtask in each trajectory generated_src_demo_labels = [] # like @generated_src_demo_inds, but padded to align with size of @generated_actions @@ -381,6 +382,7 @@ def generate( generated_obs += exec_results["observations"] generated_datagen_infos += exec_results["datagen_infos"] generated_actions.append(exec_results["actions"]) + generated_actions_abs.append(exec_results["actions_abs"]) generated_success = generated_success or exec_results["success"] generated_src_demo_inds.append(selected_src_demo_ind) generated_src_demo_labels.append(selected_src_demo_ind * np.ones((exec_results["actions"].shape[0], 1), dtype=int)) @@ -394,6 +396,7 @@ def generate( # merge numpy arrays if len(generated_actions) > 0: generated_actions = np.concatenate(generated_actions, axis=0) + generated_actions_abs = np.concatenate(generated_actions_abs, axis=0) generated_src_demo_labels = np.concatenate(generated_src_demo_labels, axis=0) results = dict( @@ -402,6 +405,7 @@ def generate( observations=generated_obs, datagen_infos=generated_datagen_infos, actions=generated_actions, + actions_abs=generated_actions_abs, success=generated_success, src_demo_inds=generated_src_demo_inds, src_demo_labels=generated_src_demo_labels, diff --git a/mimicgen/datagen/waypoint.py b/mimicgen/datagen/waypoint.py index fb1c534..5ced690 100644 --- a/mimicgen/datagen/waypoint.py +++ b/mimicgen/datagen/waypoint.py @@ -8,9 +8,11 @@ import json import numpy as np from copy import deepcopy +from scipy.spatial.transform import Rotation import mimicgen import mimicgen.utils.pose_utils as PoseUtils +import robosuite class Waypoint(object): @@ -346,6 +348,7 @@ def execute( states = [] actions = [] + actions_abs = [] observations = [] datagen_infos = [] success = { k: False for k in env.is_success() } # success metrics @@ -394,10 +397,36 @@ def execute( # step environment env.step(play_action) + # compute absolute actions by reading from robot controller + d_a = len(env.env.robots[0].action_limits[0]) + + # reshape to handle multi-robot case: (7,) -> (1, 7) or (14,) -> (2, 7) + stacked_actions = play_action.reshape(-1, d_a) + + # extract action remainder (gripper and any additional actions after index 6) + action_remainder = stacked_actions[:, 6:] + + abs_action_components = [] + for idx, robot in enumerate(env.env.robots): + if robosuite.__version__ < "1.5": + controller = robot.controller + goal_pos = controller.goal_pos + goal_ori = Rotation.from_matrix(controller.goal_ori).as_rotvec() + else: + controller = robot.part_controllers['right'] + goal_pos = controller.goal_pos + goal_ori = Rotation.from_matrix(controller.goal_ori).as_rotvec() + + # concatenate pos, ori, and action remainder for this robot + abs_action_components.append(np.concatenate([goal_pos, goal_ori, action_remainder[idx]])) + + play_action_abs = np.concatenate(abs_action_components) + # collect data states.append(state) play_action_record = play_action actions.append(play_action_record) + actions_abs.append(play_action_abs) observations.append(obs) datagen_infos.append(datagen_info) @@ -410,6 +439,7 @@ def execute( observations=observations, datagen_infos=datagen_infos, actions=np.array(actions), + actions_abs=np.array(actions_abs), success=bool(success["task"]), ) return results diff --git a/mimicgen/scripts/generate_dataset.py b/mimicgen/scripts/generate_dataset.py index 1dc0b2b..834f371 100644 --- a/mimicgen/scripts/generate_dataset.py +++ b/mimicgen/scripts/generate_dataset.py @@ -347,6 +347,7 @@ def generate_dataset( observations=(generated_traj["observations"] if mg_config.obs.collect_obs else None), datagen_info=generated_traj["datagen_infos"], actions=generated_traj["actions"], + actions_abs=generated_traj["actions_abs"], src_demo_inds=generated_traj["src_demo_inds"], src_demo_labels=generated_traj["src_demo_labels"], ) @@ -367,6 +368,7 @@ def generate_dataset( observations=(generated_traj["observations"] if mg_config.obs.collect_obs else None), datagen_info=generated_traj["datagen_infos"], actions=generated_traj["actions"], + actions_abs=generated_traj["actions_abs"], src_demo_inds=generated_traj["src_demo_inds"], src_demo_labels=generated_traj["src_demo_labels"], ) diff --git a/mimicgen/utils/file_utils.py b/mimicgen/utils/file_utils.py index e64c61b..ada0e10 100644 --- a/mimicgen/utils/file_utils.py +++ b/mimicgen/utils/file_utils.py @@ -229,6 +229,7 @@ def write_demo_to_hdf5( observations, datagen_info, actions, + actions_abs=None, src_demo_inds=None, src_demo_labels=None, ): @@ -244,6 +245,7 @@ def write_demo_to_hdf5( observations (list): list of observation dictionaries datagen_info (list): list of DatagenInfo instances actions (np.array): actions per timestep + actions_abs (np.array or None): absolute actions per timestep src_demo_inds (list or None): if provided, list of selected source demonstration indices for each subtask src_demo_labels (np.array or None): same as @src_demo_inds, but repeated to have a label for each timestep of the trajectory """ @@ -262,6 +264,10 @@ def write_demo_to_hdf5( # write actions ep_data_grp.create_dataset("actions", data=np.array(actions)) + + # write absolute actions if provided + if actions_abs is not None: + ep_data_grp.create_dataset("actions_abs", data=np.array(actions_abs)) # write simulator states if isinstance(states[0], dict):