Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions mimicgen/datagen/data_generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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))
Expand All @@ -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(
Expand All @@ -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,
Expand Down
30 changes: 30 additions & 0 deletions mimicgen/datagen/waypoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -346,6 +348,7 @@ def execute(

states = []
actions = []
actions_abs = []
observations = []
datagen_infos = []
success = { k: False for k in env.is_success() } # success metrics
Expand Down Expand Up @@ -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)

Expand All @@ -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
2 changes: 2 additions & 0 deletions mimicgen/scripts/generate_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"],
)
Expand All @@ -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"],
)
Expand Down
6 changes: 6 additions & 0 deletions mimicgen/utils/file_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -229,6 +229,7 @@ def write_demo_to_hdf5(
observations,
datagen_info,
actions,
actions_abs=None,
src_demo_inds=None,
src_demo_labels=None,
):
Expand All @@ -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
"""
Expand All @@ -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):
Expand Down