release
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
## Data processing scripts
|
||||
|
||||
These are some scripts used for processing the raw datasets from the benchmarks. We already pre-processed them and provide the final datasets. These scripts are for information only.
|
||||
@@ -0,0 +1,357 @@
|
||||
"""
|
||||
Filter avoid data based on modes.
|
||||
|
||||
Trajectories are normalized with filtered data, not the original data.
|
||||
"""
|
||||
|
||||
import os
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
import pickle
|
||||
import random
|
||||
import matplotlib.pyplot as plt
|
||||
from copy import deepcopy
|
||||
|
||||
from agent.dataset.d3il_dataset.avoiding_dataset import Avoiding_Dataset
|
||||
|
||||
|
||||
def make_dataset(
|
||||
load_path,
|
||||
save_dir,
|
||||
save_name_prefix,
|
||||
val_split,
|
||||
desired_modes,
|
||||
desired_mode_ratios,
|
||||
required_modes,
|
||||
avoid_modes,
|
||||
):
|
||||
print("Desired modes:", desired_modes)
|
||||
print("Required modes:", required_modes)
|
||||
print("Avoid modes:", avoid_modes)
|
||||
print("Desired mode ratios:", desired_mode_ratios)
|
||||
demo_dataset = Avoiding_Dataset(
|
||||
load_path,
|
||||
action_dim=2,
|
||||
obs_dim=4,
|
||||
max_len_data=200,
|
||||
)
|
||||
# from avoiding env
|
||||
level_distance = 0.18
|
||||
obstacle_offset = 0.075
|
||||
l1_ypos = -0.1
|
||||
l2_ypos = -0.1 + level_distance
|
||||
l3_ypos = -0.1 + 2 * level_distance
|
||||
# goal_ypos = -0.1 + 2.5 * level_distance
|
||||
l1_xpos = 0.5
|
||||
l2_top_xpos = 0.5 - obstacle_offset
|
||||
l2_bottom_xpos = 0.5 + obstacle_offset
|
||||
l3_top_xpos = 0.5 - 2 * obstacle_offset
|
||||
l3_mid_xpos = 0.5
|
||||
l3_bottom_xpos = 0.5 + 2 * obstacle_offset
|
||||
|
||||
def check_mode(x):
|
||||
r_x_pos = x[0]
|
||||
r_y_pos = x[1]
|
||||
mode_encoding = np.zeros((9))
|
||||
if r_y_pos - 0.01 <= l1_ypos <= r_y_pos + 0.01:
|
||||
if r_x_pos < l1_xpos:
|
||||
mode_encoding[0] = 1
|
||||
elif r_x_pos > l1_xpos:
|
||||
mode_encoding[1] = 1
|
||||
|
||||
if r_y_pos - 0.01 <= l2_ypos <= r_y_pos + 0.01:
|
||||
if r_x_pos < l2_top_xpos:
|
||||
mode_encoding[2] = 1
|
||||
elif l2_top_xpos < r_x_pos < l2_bottom_xpos:
|
||||
mode_encoding[3] = 1
|
||||
elif r_x_pos > l2_bottom_xpos:
|
||||
mode_encoding[4] = 1
|
||||
|
||||
# if r_y_pos - 0.015 <= self.l3_ypos and (not self.l3_passed):
|
||||
if r_y_pos >= l3_ypos:
|
||||
if r_x_pos < l3_top_xpos:
|
||||
mode_encoding[5] = 1
|
||||
if l3_top_xpos < r_x_pos < l3_mid_xpos:
|
||||
mode_encoding[6] = 1
|
||||
elif l3_mid_xpos < r_x_pos < l3_bottom_xpos:
|
||||
mode_encoding[7] = 1
|
||||
elif r_x_pos > l3_top_xpos:
|
||||
mode_encoding[8] = 1
|
||||
return mode_encoding
|
||||
|
||||
# extract length of each trajectory in the file
|
||||
full_traj_lengths = []
|
||||
full_actions = demo_dataset.actions
|
||||
full_obs = demo_dataset.observations
|
||||
masks = demo_dataset.masks
|
||||
action_dim = full_actions.shape[2]
|
||||
obs_dim = full_obs.shape[2]
|
||||
for ep in range(masks.shape[0]):
|
||||
full_traj_lengths.append(int(masks[ep].sum().item()))
|
||||
full_traj_lengths = np.array(full_traj_lengths)
|
||||
|
||||
# take the max and min of obs and action
|
||||
obs_min = np.zeros((obs_dim))
|
||||
obs_max = np.zeros((obs_dim))
|
||||
action_min = np.zeros((action_dim))
|
||||
action_max = np.zeros((action_dim))
|
||||
chosen_indices = []
|
||||
for i in tqdm(range(len(masks))):
|
||||
T = full_traj_lengths[i]
|
||||
obs_traj = full_obs[i, :T].numpy()
|
||||
action_traj = full_actions[i, :T].numpy()
|
||||
|
||||
# check if trajectory pass through desired hole
|
||||
flag_desired = False
|
||||
flag_required = [False for _ in required_modes] if required_modes else [True]
|
||||
flag_avoid = False
|
||||
for ob in obs_traj:
|
||||
modes = check_mode(ob)
|
||||
if any(modes[desired] for desired in desired_modes):
|
||||
desired_mode_idx = np.argmax(
|
||||
[modes[desired] for desired in desired_modes]
|
||||
)
|
||||
flag_desired = True
|
||||
if any(modes[avoid] for avoid in avoid_modes):
|
||||
flag_avoid = True
|
||||
break
|
||||
for j, required in enumerate(required_modes):
|
||||
if modes[required]:
|
||||
flag_required[j] = True
|
||||
if flag_avoid or not flag_desired or not all(flag_required):
|
||||
continue
|
||||
if desired_mode_ratios:
|
||||
if random.random() > desired_mode_ratios[desired_mode_idx]:
|
||||
continue
|
||||
chosen_indices.append(i)
|
||||
|
||||
obs_min = np.min(np.vstack((obs_min, np.min(obs_traj, axis=0))), axis=0)
|
||||
obs_max = np.max(np.vstack((obs_max, np.max(obs_traj, axis=0))), axis=0)
|
||||
action_min = np.min(
|
||||
np.vstack((action_min, np.min(action_traj, axis=0))), axis=0
|
||||
)
|
||||
action_max = np.max(
|
||||
np.vstack((action_max, np.max(action_traj, axis=0))), axis=0
|
||||
)
|
||||
if len(chosen_indices) == 0:
|
||||
raise ValueError("No data found for the desired/required modes")
|
||||
chosen_indices = np.array(chosen_indices)
|
||||
traj_lengths = full_traj_lengths[chosen_indices]
|
||||
actions = demo_dataset.actions[chosen_indices]
|
||||
obs = demo_dataset.observations[chosen_indices]
|
||||
max_traj_length = np.max(traj_lengths)
|
||||
|
||||
# split indices in train and val
|
||||
num_traj = len(traj_lengths)
|
||||
num_train = int(num_traj * (1 - val_split))
|
||||
train_indices = random.sample(range(num_traj), k=num_train)
|
||||
|
||||
logger.info("\n========== Basic Info ===========")
|
||||
logger.info("total transitions: {}".format(np.sum(traj_lengths)))
|
||||
logger.info("total trajectories: {}".format(len(traj_lengths)))
|
||||
logger.info(
|
||||
f"traj length mean/std: {np.mean(traj_lengths)}, {np.std(traj_lengths)}"
|
||||
)
|
||||
logger.info(f"traj length min/max: {np.min(traj_lengths)}, {np.max(traj_lengths)}")
|
||||
logger.info(f"obs min: {obs_min}")
|
||||
logger.info(f"obs max: {obs_max}")
|
||||
logger.info(f"action min: {action_min}")
|
||||
logger.info(f"action max: {action_max}")
|
||||
|
||||
# do over all indices
|
||||
out_train = {}
|
||||
keys = [
|
||||
"observations",
|
||||
"actions",
|
||||
"rewards",
|
||||
]
|
||||
total_timesteps = actions.shape[1]
|
||||
out_train["observations"] = np.empty((0, total_timesteps, obs_dim))
|
||||
out_train["actions"] = np.empty((0, total_timesteps, action_dim))
|
||||
out_train["rewards"] = np.empty((0, total_timesteps))
|
||||
out_train["traj_length"] = []
|
||||
out_val = deepcopy(out_train)
|
||||
for i in tqdm(range(len(traj_lengths))):
|
||||
if i in train_indices:
|
||||
out = out_train
|
||||
else:
|
||||
out = out_val
|
||||
T = traj_lengths[i]
|
||||
obs_traj = obs[i].numpy()
|
||||
action_traj = actions[i].numpy()
|
||||
|
||||
# scale to [-1, 1] for both ob and action
|
||||
obs_traj = 2 * (obs_traj - obs_min) / (obs_max - obs_min + 1e-6) - 1
|
||||
action_traj = (
|
||||
2 * (action_traj - action_min) / (action_max - action_min + 1e-6) - 1
|
||||
)
|
||||
|
||||
# get episode length
|
||||
traj_length = T
|
||||
out["traj_length"].append(traj_length)
|
||||
|
||||
# extract
|
||||
rewards = np.zeros(total_timesteps) # no reward from d3il dataset
|
||||
data_traj = {
|
||||
"observations": obs_traj,
|
||||
"actions": action_traj,
|
||||
"rewards": rewards,
|
||||
}
|
||||
for key in keys:
|
||||
traj = data_traj[key]
|
||||
out[key] = np.vstack((out[key], traj[None]))
|
||||
|
||||
# plot all trajectories and save in a figure
|
||||
def plot(out, name):
|
||||
def get_obj_xy_list():
|
||||
mid_pos = 0.5
|
||||
offset = 0.075
|
||||
first_level_y = -0.1
|
||||
level_distance = 0.18
|
||||
return [
|
||||
[mid_pos, first_level_y],
|
||||
[mid_pos - offset, first_level_y + level_distance],
|
||||
[mid_pos + offset, first_level_y + level_distance],
|
||||
[mid_pos - 2 * offset, first_level_y + 2 * level_distance],
|
||||
[mid_pos, first_level_y + 2 * level_distance],
|
||||
[mid_pos + 2 * offset, first_level_y + 2 * level_distance],
|
||||
]
|
||||
|
||||
pillar_xys = get_obj_xy_list()
|
||||
fig = plt.figure()
|
||||
all_trajs = out["observations"] # num x timestep x obs
|
||||
for traj, traj_length in zip(all_trajs, out["traj_length"]):
|
||||
# unnormalize
|
||||
traj = (traj + 1) / 2 # [-1, 1] -> [0, 1]
|
||||
traj = traj * (obs_max - obs_min) + obs_min
|
||||
plt.plot(
|
||||
traj[:traj_length, 2], traj[:traj_length, 3], color=(0.3, 0.3, 0.3)
|
||||
)
|
||||
plt.axhline(y=0.4, color=np.array([31, 119, 180]) / 255, linestyle="-")
|
||||
for xy in pillar_xys:
|
||||
circle = plt.Circle(xy, 0.01, color=(0.0, 0.0, 0.0), fill=True)
|
||||
plt.gca().add_patch(circle)
|
||||
plt.xlabel("X pos")
|
||||
plt.ylabel("Y pos")
|
||||
plt.xlim([0.2, 0.8])
|
||||
plt.ylim([-0.3, 0.5])
|
||||
ax = plt.gca()
|
||||
ax.set_aspect("equal", adjustable="box")
|
||||
ax.set_facecolor("white")
|
||||
plt.savefig(os.path.join(save_dir, name))
|
||||
plt.close(fig)
|
||||
|
||||
plot(out_train, name="train-trajs.png")
|
||||
plot(out_val, name="val-trajs.png")
|
||||
|
||||
# Save to np file
|
||||
save_train_path = os.path.join(save_dir, save_name_prefix + "train.pkl")
|
||||
save_val_path = os.path.join(save_dir, save_name_prefix + "val.pkl")
|
||||
with open(save_train_path, "wb") as f:
|
||||
pickle.dump(out_train, f)
|
||||
with open(save_val_path, "wb") as f:
|
||||
pickle.dump(out_val, f)
|
||||
normalization_save_path = os.path.join(
|
||||
save_dir, save_name_prefix + "normalization.npz"
|
||||
)
|
||||
np.savez(
|
||||
normalization_save_path,
|
||||
obs_min=obs_min,
|
||||
obs_max=obs_max,
|
||||
action_min=action_min,
|
||||
action_max=action_max,
|
||||
)
|
||||
|
||||
# debug
|
||||
logger.info("\n========== Final ===========")
|
||||
logger.info(
|
||||
f"Train - Number of episodes and transitions: {len(out_train['traj_length'])}, {np.sum(out_train['traj_length'])}"
|
||||
)
|
||||
logger.info(
|
||||
f"Val - Number of episodes and transitions: {len(out_val['traj_length'])}, {np.sum(out_val['traj_length'])}"
|
||||
)
|
||||
logger.info(
|
||||
f"Train - Mean/Std trajectory length: {np.mean(out_train['traj_length'])}, {np.std(out_train['traj_length'])}"
|
||||
)
|
||||
logger.info(
|
||||
f"Train - Max/Min trajectory length: {np.max(out_train['traj_length'])}, {np.min(out_train['traj_length'])}"
|
||||
)
|
||||
if val_split > 0:
|
||||
logger.info(
|
||||
f"Val - Mean/Std trajectory length: {np.mean(out_val['traj_length'])}, {np.std(out_val['traj_length'])}"
|
||||
)
|
||||
logger.info(
|
||||
f"Val - Max/Min trajectory length: {np.max(out_val['traj_length'])}, {np.min(out_val['traj_length'])}"
|
||||
)
|
||||
for obs_dim_ind in range(obs_dim):
|
||||
obs = out_train["observations"][:, :, obs_dim_ind]
|
||||
logger.info(
|
||||
f"Train - Obs dim {obs_dim_ind+1} mean {np.mean(obs)} std {np.std(obs)} min {np.min(obs)} max {np.max(obs)}"
|
||||
)
|
||||
for action_dim_ind in range(action_dim):
|
||||
action = out_train["actions"][:, :, action_dim_ind]
|
||||
logger.info(
|
||||
f"Train - Action dim {action_dim_ind+1} mean {np.mean(action)} std {np.std(action)} min {np.min(action)} max {np.max(action)}"
|
||||
)
|
||||
if val_split > 0:
|
||||
for obs_dim_ind in range(obs_dim):
|
||||
obs = out_val["observations"][:, :, obs_dim_ind]
|
||||
logger.info(
|
||||
f"Val - Obs dim {obs_dim_ind+1} mean {np.mean(obs)} std {np.std(obs)} min {np.min(obs)} max {np.max(obs)}"
|
||||
)
|
||||
for action_dim_ind in range(action_dim):
|
||||
action = out_val["actions"][:, :, action_dim_ind]
|
||||
logger.info(
|
||||
f"Val - Action dim {action_dim_ind+1} mean {np.mean(action)} std {np.std(action)} min {np.min(action)} max {np.max(action)}"
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--load_path", type=str, default=".")
|
||||
parser.add_argument("--save_dir", type=str, default=".")
|
||||
parser.add_argument("--save_name_prefix", type=str, default="")
|
||||
parser.add_argument("--val_split", type=float, default="0.2")
|
||||
parser.add_argument("--desired_modes", nargs="+", type=int)
|
||||
parser.add_argument("--desired_mode_ratios", nargs="+", type=float, default=[])
|
||||
parser.add_argument("--required_modes", nargs="+", type=int, default=[])
|
||||
parser.add_argument("--avoid_modes", nargs="+", type=int, default=[])
|
||||
args = parser.parse_args()
|
||||
if len(args.desired_mode_ratios) > 0:
|
||||
assert len(args.desired_modes) == len(
|
||||
args.desired_mode_ratios
|
||||
), "Desired modes and desired mode ratios should have the same length"
|
||||
|
||||
os.makedirs(args.save_dir, exist_ok=True)
|
||||
|
||||
import logging
|
||||
import datetime
|
||||
|
||||
os.makedirs(args.save_dir, exist_ok=True)
|
||||
log_path = os.path.join(
|
||||
args.save_dir,
|
||||
args.save_name_prefix
|
||||
+ f"{datetime.datetime.now().strftime('%Y_%m_%d_%H_%M_%S')}.log",
|
||||
)
|
||||
logger = logging.getLogger("get_D4RL_dataset")
|
||||
logger.setLevel(logging.INFO)
|
||||
file_handler = logging.FileHandler(log_path)
|
||||
file_handler.setLevel(logging.INFO) # Set the minimum level for this handler
|
||||
formatter = logging.Formatter(
|
||||
"%(asctime)s - %(name)s - %(levelname)s - %(message)s"
|
||||
)
|
||||
file_handler.setFormatter(formatter)
|
||||
logger.addHandler(file_handler)
|
||||
|
||||
make_dataset(
|
||||
args.load_path,
|
||||
args.save_dir,
|
||||
args.save_name_prefix,
|
||||
args.val_split,
|
||||
args.desired_modes,
|
||||
args.desired_mode_ratios,
|
||||
args.required_modes,
|
||||
args.avoid_modes,
|
||||
)
|
||||
@@ -0,0 +1,247 @@
|
||||
"""
|
||||
Download D4RL dataset and save it into our custom format so it can be loaded for diffusion training.
|
||||
|
||||
"""
|
||||
|
||||
import os
|
||||
import logging
|
||||
import gym
|
||||
import random
|
||||
from copy import deepcopy
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
import pickle
|
||||
|
||||
import d4rl.gym_mujoco # Import required to register environments
|
||||
|
||||
|
||||
def make_dataset(env_name, save_dir, save_name_prefix, val_split, logger):
|
||||
# Create the environment
|
||||
env = gym.make(env_name)
|
||||
|
||||
# d4rl abides by the OpenAI gym interface
|
||||
env.reset()
|
||||
env.step(env.action_space.sample())
|
||||
|
||||
# Each task is associated with a dataset
|
||||
# dataset contains observations, actions, rewards, terminals, and infos
|
||||
dataset = env.get_dataset()
|
||||
logger.info("\n========== Basic Info ===========")
|
||||
logger.info(f"Keys in the dataset: {dataset.keys()}")
|
||||
logger.info(f"Observation shape: {dataset['observations'].shape}")
|
||||
logger.info(f"Action shape: {dataset['actions'].shape}")
|
||||
terminal_indices = np.argwhere(dataset["terminals"])[:, 0]
|
||||
timeout_indices = np.argwhere(dataset["timeouts"])[:, 0]
|
||||
obs_dim = dataset["observations"].shape[1]
|
||||
action_dim = dataset["actions"].shape[1]
|
||||
done_indices = np.concatenate([terminal_indices, timeout_indices])
|
||||
done_indices = np.sort(done_indices)
|
||||
traj_lengths = []
|
||||
prev_index = 0
|
||||
for i in tqdm(range(len(done_indices))):
|
||||
# get episode length
|
||||
cur_index = done_indices[i]
|
||||
traj_lengths.append(cur_index - prev_index + 1)
|
||||
prev_index = cur_index + 1
|
||||
obs_min = np.min(dataset["observations"], axis=0)
|
||||
obs_max = np.max(dataset["observations"], axis=0)
|
||||
action_min = np.min(dataset["actions"], axis=0)
|
||||
action_max = np.max(dataset["actions"], axis=0)
|
||||
max_episode_steps = max(traj_lengths)
|
||||
logger.info("total transitions: {}".format(np.sum(traj_lengths)))
|
||||
logger.info("total trajectories: {}".format(len(traj_lengths)))
|
||||
logger.info(
|
||||
f"traj length mean/std: {np.mean(traj_lengths)}, {np.std(traj_lengths)}"
|
||||
)
|
||||
logger.info(f"traj length min/max: {np.min(traj_lengths)}, {np.max(traj_lengths)}")
|
||||
logger.info(f"obs min: {obs_min}")
|
||||
logger.info(f"obs max: {obs_max}")
|
||||
logger.info(f"action min: {action_min}")
|
||||
logger.info(f"action max: {action_max}")
|
||||
|
||||
# Subsample episodes by taking the first ones
|
||||
if args.max_episodes > 0:
|
||||
traj_lengths = traj_lengths[: args.max_episodes]
|
||||
done_indices = done_indices[: args.max_episodes]
|
||||
max_episode_steps = max(traj_lengths)
|
||||
|
||||
# split indices in train and val
|
||||
num_traj = len(traj_lengths)
|
||||
num_train = int(num_traj * (1 - val_split))
|
||||
train_indices = random.sample(range(num_traj), k=num_train)
|
||||
|
||||
# do over all indices
|
||||
out_train = {}
|
||||
keys = [
|
||||
"observations",
|
||||
"actions",
|
||||
"rewards",
|
||||
]
|
||||
out_train["observations"] = np.empty(
|
||||
(0, max_episode_steps, dataset["observations"].shape[-1])
|
||||
)
|
||||
out_train["actions"] = np.empty(
|
||||
(0, max_episode_steps, dataset["actions"].shape[-1])
|
||||
)
|
||||
out_train["rewards"] = np.empty((0, max_episode_steps))
|
||||
out_train["traj_length"] = []
|
||||
out_val = deepcopy(out_train)
|
||||
prev_index = 0
|
||||
train_episode_reward_all = []
|
||||
val_episode_reward_all = []
|
||||
for i in tqdm(range(len(done_indices))):
|
||||
if i in train_indices:
|
||||
out = out_train
|
||||
episode_reward_all = train_episode_reward_all
|
||||
else:
|
||||
out = out_val
|
||||
episode_reward_all = val_episode_reward_all
|
||||
|
||||
# get episode length
|
||||
cur_index = done_indices[i]
|
||||
traj_length = cur_index - prev_index + 1
|
||||
|
||||
# Skip if the episode has no reward
|
||||
if np.sum(dataset["rewards"][prev_index : cur_index + 1]) > 0:
|
||||
out["traj_length"].append(traj_length)
|
||||
|
||||
# apply padding to make all episodes have the same max steps
|
||||
for key in keys:
|
||||
traj = dataset[key][prev_index : cur_index + 1]
|
||||
|
||||
# also scale
|
||||
if key == "observations":
|
||||
traj = 2 * (traj - obs_min) / (obs_max - obs_min + 1e-6) - 1
|
||||
elif key == "actions":
|
||||
traj = (
|
||||
2 * (traj - action_min) / (action_max - action_min + 1e-6) - 1
|
||||
)
|
||||
|
||||
if traj.ndim == 1:
|
||||
traj = np.pad(
|
||||
traj,
|
||||
(0, max_episode_steps - len(traj)),
|
||||
mode="constant",
|
||||
constant_values=0,
|
||||
)
|
||||
else:
|
||||
traj = np.pad(
|
||||
traj,
|
||||
((0, max_episode_steps - traj.shape[0]), (0, 0)),
|
||||
mode="constant",
|
||||
constant_values=0,
|
||||
)
|
||||
out[key] = np.vstack((out[key], traj[None]))
|
||||
|
||||
# check reward
|
||||
episode_reward_all.append(np.sum(out["rewards"][-1]))
|
||||
else:
|
||||
print(f"skipping {i} / {len(done_indices)}")
|
||||
|
||||
# update prev index
|
||||
prev_index = cur_index + 1
|
||||
|
||||
# Save to np file
|
||||
save_train_path = os.path.join(save_dir, save_name_prefix + "train.pkl")
|
||||
save_val_path = os.path.join(save_dir, save_name_prefix + "val.pkl")
|
||||
with open(save_train_path, "wb") as f:
|
||||
pickle.dump(out_train, f)
|
||||
with open(save_val_path, "wb") as f:
|
||||
pickle.dump(out_val, f)
|
||||
normalization_save_path = os.path.join(
|
||||
save_dir, save_name_prefix + "normalization.npz"
|
||||
)
|
||||
np.savez(
|
||||
normalization_save_path,
|
||||
obs_min=obs_min,
|
||||
obs_max=obs_max,
|
||||
action_min=action_min,
|
||||
action_max=action_max,
|
||||
)
|
||||
|
||||
# debug
|
||||
logger.info("\n========== Final ===========")
|
||||
logger.info(
|
||||
f"Train - Number of episodes and transitions: {len(out_train['traj_length'])}, {np.sum(out_train['traj_length'])}"
|
||||
)
|
||||
logger.info(
|
||||
f"Val - Number of episodes and transitions: {len(out_val['traj_length'])}, {np.sum(out_val['traj_length'])}"
|
||||
)
|
||||
logger.info(
|
||||
f"Train - Mean/Std trajectory length: {np.mean(out_train['traj_length'])}, {np.std(out_train['traj_length'])}"
|
||||
)
|
||||
logger.info(
|
||||
f"Train - Max/Min trajectory length: {np.max(out_train['traj_length'])}, {np.min(out_train['traj_length'])}"
|
||||
)
|
||||
if val_split > 0:
|
||||
logger.info(
|
||||
f"Val - Mean/Std trajectory length: {np.mean(out_val['traj_length'])}, {np.std(out_val['traj_length'])}"
|
||||
)
|
||||
logger.info(
|
||||
f"Val - Max/Min trajectory length: {np.max(out_val['traj_length'])}, {np.min(out_val['traj_length'])}"
|
||||
)
|
||||
logger.info(
|
||||
f"Train - Mean/Std episode reward: {np.mean(train_episode_reward_all)}, {np.std(train_episode_reward_all)}"
|
||||
)
|
||||
if val_split > 0:
|
||||
logger.info(
|
||||
f"Val - Mean/Std episode reward: {np.mean(val_episode_reward_all)}, {np.std(val_episode_reward_all)}"
|
||||
)
|
||||
for obs_dim_ind in range(obs_dim):
|
||||
obs = out_train["observations"][:, :, obs_dim_ind]
|
||||
logger.info(
|
||||
f"Train - Obs dim {obs_dim_ind+1} mean {np.mean(obs)} std {np.std(obs)} min {np.min(obs)} max {np.max(obs)}"
|
||||
)
|
||||
for action_dim_ind in range(action_dim):
|
||||
action = out_train["actions"][:, :, action_dim_ind]
|
||||
logger.info(
|
||||
f"Train - Action dim {action_dim_ind+1} mean {np.mean(action)} std {np.std(action)} min {np.min(action)} max {np.max(action)}"
|
||||
)
|
||||
if val_split > 0:
|
||||
for obs_dim_ind in range(obs_dim):
|
||||
obs = out_val["observations"][:, :, obs_dim_ind]
|
||||
logger.info(
|
||||
f"Val - Obs dim {obs_dim_ind+1} mean {np.mean(obs)} std {np.std(obs)} min {np.min(obs)} max {np.max(obs)}"
|
||||
)
|
||||
for action_dim_ind in range(action_dim):
|
||||
action = out_val["actions"][:, :, action_dim_ind]
|
||||
logger.info(
|
||||
f"Val - Action dim {action_dim_ind+1} mean {np.mean(action)} std {np.std(action)} min {np.min(action)} max {np.max(action)}"
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--env_name", type=str, default="hopper-medium-v2")
|
||||
parser.add_argument("--save_dir", type=str, default=".")
|
||||
parser.add_argument("--save_name_prefix", type=str, default="")
|
||||
parser.add_argument("--val_split", type=float, default="0.2")
|
||||
parser.add_argument("--max_episodes", type=int, default="-1")
|
||||
args = parser.parse_args()
|
||||
|
||||
import datetime
|
||||
|
||||
# import logging.config
|
||||
if args.max_episodes > 0:
|
||||
args.save_name_prefix += f"max_episodes_{args.max_episodes}_"
|
||||
os.makedirs(args.save_dir, exist_ok=True)
|
||||
log_path = os.path.join(
|
||||
args.save_dir,
|
||||
args.save_name_prefix
|
||||
+ f"_{datetime.datetime.now().strftime('%Y_%m_%d_%H_%M_%S')}.log",
|
||||
)
|
||||
logger = logging.getLogger("get_D4RL_dataset")
|
||||
logger.setLevel(logging.INFO)
|
||||
file_handler = logging.FileHandler(log_path)
|
||||
file_handler.setLevel(logging.INFO) # Set the minimum level for this handler
|
||||
formatter = logging.Formatter(
|
||||
"%(asctime)s - %(name)s - %(levelname)s - %(message)s"
|
||||
)
|
||||
file_handler.setFormatter(formatter)
|
||||
logger.addHandler(file_handler)
|
||||
|
||||
make_dataset(
|
||||
args.env_name, args.save_dir, args.save_name_prefix, args.val_split, logger
|
||||
)
|
||||
@@ -0,0 +1,294 @@
|
||||
"""
|
||||
Process d3il dataset and save it into our custom format so it can be loaded for diffusion training.
|
||||
"""
|
||||
|
||||
import os
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
import pickle
|
||||
import random
|
||||
import matplotlib.pyplot as plt
|
||||
from copy import deepcopy
|
||||
|
||||
from agent.dataset.d3il_dataset.aligning_dataset import Aligning_Dataset
|
||||
from agent.dataset.d3il_dataset.avoiding_dataset import Avoiding_Dataset
|
||||
from agent.dataset.d3il_dataset.pushing_dataset import Pushing_Dataset
|
||||
from agent.dataset.d3il_dataset.sorting_dataset import Sorting_Dataset
|
||||
from agent.dataset.d3il_dataset.stacking_dataset import Stacking_Dataset
|
||||
|
||||
|
||||
def make_dataset(load_path, save_dir, save_name_prefix, env_type, val_split):
|
||||
if env_type == "align":
|
||||
demo_dataset = Aligning_Dataset(
|
||||
load_path,
|
||||
action_dim=3,
|
||||
obs_dim=20,
|
||||
max_len_data=512,
|
||||
)
|
||||
elif env_type == "avoid":
|
||||
demo_dataset = Avoiding_Dataset(
|
||||
load_path,
|
||||
action_dim=2,
|
||||
obs_dim=4,
|
||||
max_len_data=200,
|
||||
)
|
||||
elif env_type == "push":
|
||||
demo_dataset = Pushing_Dataset(
|
||||
load_path,
|
||||
action_dim=2,
|
||||
obs_dim=10,
|
||||
max_len_data=512,
|
||||
)
|
||||
elif env_type == "sort":
|
||||
# Can config number of boxes to be 2, 4, or 6.
|
||||
# TODO: add other numbers of boxes
|
||||
demo_dataset = Sorting_Dataset(
|
||||
load_path,
|
||||
action_dim=2,
|
||||
obs_dim=10,
|
||||
max_len_data=600,
|
||||
num_boxes=2,
|
||||
)
|
||||
elif env_type == "stack":
|
||||
demo_dataset = Stacking_Dataset(
|
||||
load_path,
|
||||
action_dim=8,
|
||||
obs_dim=20,
|
||||
max_len_data=1000,
|
||||
)
|
||||
else:
|
||||
raise ValueError("Invalid dataset type.")
|
||||
# extract length of each trajectory in the file
|
||||
traj_lengths = []
|
||||
actions = demo_dataset.actions
|
||||
obs = demo_dataset.observations
|
||||
masks = demo_dataset.masks
|
||||
action_dim = actions.shape[2]
|
||||
obs_dim = obs.shape[2]
|
||||
for ep in range(masks.shape[0]):
|
||||
traj_lengths.append(int(masks[ep].sum().item()))
|
||||
traj_lengths = np.array(traj_lengths)
|
||||
max_traj_length = np.max(traj_lengths)
|
||||
|
||||
# split indices in train and val
|
||||
num_traj = len(traj_lengths)
|
||||
num_train = int(num_traj * (1 - val_split))
|
||||
train_indices = random.sample(range(num_traj), k=num_train)
|
||||
|
||||
# take the max and min of obs and action
|
||||
obs_min = np.zeros((obs_dim))
|
||||
obs_max = np.zeros((obs_dim))
|
||||
action_min = np.zeros((action_dim))
|
||||
action_max = np.zeros((action_dim))
|
||||
for i in tqdm(range(len(traj_lengths))):
|
||||
T = traj_lengths[i]
|
||||
obs_traj = obs[i, :T].numpy()
|
||||
action_traj = actions[i, :T].numpy()
|
||||
obs_min = np.min(np.vstack((obs_min, np.min(obs_traj, axis=0))), axis=0)
|
||||
obs_max = np.max(np.vstack((obs_max, np.max(obs_traj, axis=0))), axis=0)
|
||||
action_min = np.min(
|
||||
np.vstack((action_min, np.min(action_traj, axis=0))), axis=0
|
||||
)
|
||||
action_max = np.max(
|
||||
np.vstack((action_max, np.max(action_traj, axis=0))), axis=0
|
||||
)
|
||||
logger.info("\n========== Basic Info ===========")
|
||||
logger.info("total transitions: {}".format(np.sum(traj_lengths)))
|
||||
logger.info("total trajectories: {}".format(len(traj_lengths)))
|
||||
logger.info(
|
||||
f"traj length mean/std: {np.mean(traj_lengths)}, {np.std(traj_lengths)}"
|
||||
)
|
||||
logger.info(f"traj length min/max: {np.min(traj_lengths)}, {np.max(traj_lengths)}")
|
||||
logger.info(f"obs min: {obs_min}")
|
||||
logger.info(f"obs max: {obs_max}")
|
||||
logger.info(f"action min: {action_min}")
|
||||
logger.info(f"action max: {action_max}")
|
||||
|
||||
# do over all indices
|
||||
out_train = {}
|
||||
keys = [
|
||||
"observations",
|
||||
"actions",
|
||||
"rewards",
|
||||
]
|
||||
total_timesteps = actions.shape[1]
|
||||
out_train["observations"] = np.empty((0, total_timesteps, obs_dim))
|
||||
out_train["actions"] = np.empty((0, total_timesteps, action_dim))
|
||||
out_train["rewards"] = np.empty((0, total_timesteps))
|
||||
out_train["traj_length"] = []
|
||||
out_val = deepcopy(out_train)
|
||||
for i in tqdm(range(len(traj_lengths))):
|
||||
if i in train_indices:
|
||||
out = out_train
|
||||
else:
|
||||
out = out_val
|
||||
|
||||
T = traj_lengths[i]
|
||||
obs_traj = obs[i].numpy()
|
||||
action_traj = actions[i].numpy()
|
||||
|
||||
# scale to [-1, 1] for both ob and action
|
||||
obs_traj = 2 * (obs_traj - obs_min) / (obs_max - obs_min + 1e-6) - 1
|
||||
action_traj = (
|
||||
2 * (action_traj - action_min) / (action_max - action_min + 1e-6) - 1
|
||||
)
|
||||
|
||||
# get episode length
|
||||
traj_length = T
|
||||
out["traj_length"].append(traj_length)
|
||||
|
||||
# extract
|
||||
rewards = np.zeros(total_timesteps) # no reward from d3il dataset
|
||||
data_traj = {
|
||||
"observations": obs_traj,
|
||||
"actions": action_traj,
|
||||
"rewards": rewards,
|
||||
}
|
||||
for key in keys:
|
||||
traj = data_traj[key]
|
||||
out[key] = np.vstack((out[key], traj[None]))
|
||||
|
||||
# plot all trajectories and save in a figure
|
||||
def plot(out, name):
|
||||
def get_obj_xy_list():
|
||||
mid_pos = 0.5
|
||||
offset = 0.075
|
||||
first_level_y = -0.1
|
||||
level_distance = 0.18
|
||||
return [
|
||||
[mid_pos, first_level_y],
|
||||
[mid_pos - offset, first_level_y + level_distance],
|
||||
[mid_pos + offset, first_level_y + level_distance],
|
||||
[mid_pos - 2 * offset, first_level_y + 2 * level_distance],
|
||||
[mid_pos, first_level_y + 2 * level_distance],
|
||||
[mid_pos + 2 * offset, first_level_y + 2 * level_distance],
|
||||
]
|
||||
|
||||
pillar_xys = get_obj_xy_list()
|
||||
fig = plt.figure()
|
||||
all_trajs = out["observations"] # num x timestep x obs
|
||||
for traj, traj_length in zip(all_trajs, out["traj_length"]):
|
||||
# unnormalize
|
||||
traj = (traj + 1) / 2 # [-1, 1] -> [0, 1]
|
||||
traj = traj * (obs_max - obs_min) + obs_min
|
||||
plt.plot(
|
||||
traj[:traj_length, 2], traj[:traj_length, 3], color=(0.3, 0.3, 0.3)
|
||||
)
|
||||
plt.axhline(y=0.4, color=np.array([31, 119, 180]) / 255, linestyle="-")
|
||||
for xy in pillar_xys:
|
||||
circle = plt.Circle(xy, 0.01, color=(0.0, 0.0, 0.0), fill=True)
|
||||
plt.gca().add_patch(circle)
|
||||
plt.xlabel("X pos")
|
||||
plt.ylabel("Y pos")
|
||||
plt.xlim([0.2, 0.8])
|
||||
plt.ylim([-0.3, 0.5])
|
||||
ax = plt.gca()
|
||||
ax.set_aspect("equal", adjustable="box")
|
||||
ax.set_facecolor("white")
|
||||
plt.savefig(os.path.join(save_dir, name))
|
||||
plt.close(fig)
|
||||
|
||||
plot(out_train, name="train-trajs.png")
|
||||
plot(out_val, name="val-trajs.png")
|
||||
|
||||
# Save to np file
|
||||
save_train_path = os.path.join(save_dir, save_name_prefix + "train.pkl")
|
||||
save_val_path = os.path.join(save_dir, save_name_prefix + "val.pkl")
|
||||
with open(save_train_path, "wb") as f:
|
||||
pickle.dump(out_train, f)
|
||||
with open(save_val_path, "wb") as f:
|
||||
pickle.dump(out_val, f)
|
||||
normalization_save_path = os.path.join(
|
||||
save_dir, save_name_prefix + "normalization.npz"
|
||||
)
|
||||
np.savez(
|
||||
normalization_save_path,
|
||||
obs_min=obs_min,
|
||||
obs_max=obs_max,
|
||||
action_min=action_min,
|
||||
action_max=action_max,
|
||||
)
|
||||
|
||||
# debug
|
||||
logger.info("\n========== Final ===========")
|
||||
logger.info(
|
||||
f"Train - Number of episodes and transitions: {len(out_train['traj_length'])}, {np.sum(out_train['traj_length'])}"
|
||||
)
|
||||
logger.info(
|
||||
f"Val - Number of episodes and transitions: {len(out_val['traj_length'])}, {np.sum(out_val['traj_length'])}"
|
||||
)
|
||||
logger.info(
|
||||
f"Train - Mean/Std trajectory length: {np.mean(out_train['traj_length'])}, {np.std(out_train['traj_length'])}"
|
||||
)
|
||||
logger.info(
|
||||
f"Train - Max/Min trajectory length: {np.max(out_train['traj_length'])}, {np.min(out_train['traj_length'])}"
|
||||
)
|
||||
if val_split > 0:
|
||||
logger.info(
|
||||
f"Val - Mean/Std trajectory length: {np.mean(out_val['traj_length'])}, {np.std(out_val['traj_length'])}"
|
||||
)
|
||||
logger.info(
|
||||
f"Val - Max/Min trajectory length: {np.max(out_val['traj_length'])}, {np.min(out_val['traj_length'])}"
|
||||
)
|
||||
for obs_dim_ind in range(obs_dim):
|
||||
obs = out_train["observations"][:, :, obs_dim_ind]
|
||||
logger.info(
|
||||
f"Train - Obs dim {obs_dim_ind+1} mean {np.mean(obs)} std {np.std(obs)} min {np.min(obs)} max {np.max(obs)}"
|
||||
)
|
||||
for action_dim_ind in range(action_dim):
|
||||
action = out_train["actions"][:, :, action_dim_ind]
|
||||
logger.info(
|
||||
f"Train - Action dim {action_dim_ind+1} mean {np.mean(action)} std {np.std(action)} min {np.min(action)} max {np.max(action)}"
|
||||
)
|
||||
if val_split > 0:
|
||||
for obs_dim_ind in range(obs_dim):
|
||||
obs = out_val["observations"][:, :, obs_dim_ind]
|
||||
logger.info(
|
||||
f"Val - Obs dim {obs_dim_ind+1} mean {np.mean(obs)} std {np.std(obs)} min {np.min(obs)} max {np.max(obs)}"
|
||||
)
|
||||
for action_dim_ind in range(action_dim):
|
||||
action = out_val["actions"][:, :, action_dim_ind]
|
||||
logger.info(
|
||||
f"Val - Action dim {action_dim_ind+1} mean {np.mean(action)} std {np.std(action)} min {np.min(action)} max {np.max(action)}"
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--load_path", type=str, default=".")
|
||||
parser.add_argument("--save_dir", type=str, default=".")
|
||||
parser.add_argument("--save_name_prefix", type=str, default="")
|
||||
parser.add_argument("--env_type", type=str, default="align")
|
||||
parser.add_argument("--val_split", type=float, default="0.2")
|
||||
args = parser.parse_args()
|
||||
|
||||
os.makedirs(args.save_dir, exist_ok=True)
|
||||
|
||||
import logging
|
||||
import datetime
|
||||
|
||||
os.makedirs(args.save_dir, exist_ok=True)
|
||||
log_path = os.path.join(
|
||||
args.save_dir,
|
||||
args.save_name_prefix
|
||||
+ f"_{datetime.datetime.now().strftime('%Y_%m_%d_%H_%M_%S')}.log",
|
||||
)
|
||||
logger = logging.getLogger("get_D4RL_dataset")
|
||||
logger.setLevel(logging.INFO)
|
||||
file_handler = logging.FileHandler(log_path)
|
||||
file_handler.setLevel(logging.INFO) # Set the minimum level for this handler
|
||||
formatter = logging.Formatter(
|
||||
"%(asctime)s - %(name)s - %(levelname)s - %(message)s"
|
||||
)
|
||||
file_handler.setFormatter(formatter)
|
||||
logger.addHandler(file_handler)
|
||||
|
||||
make_dataset(
|
||||
args.load_path,
|
||||
args.save_dir,
|
||||
args.save_name_prefix,
|
||||
args.env_type,
|
||||
args.val_split,
|
||||
)
|
||||
@@ -0,0 +1,421 @@
|
||||
"""
|
||||
Process robomimic dataset and save it into our custom format so it can be loaded for diffusion training.
|
||||
|
||||
Using some code from robomimic/robomimic/scripts/get_dataset_info.py
|
||||
|
||||
can-mh:
|
||||
total transitions: 62756
|
||||
total trajectories: 300
|
||||
traj length mean: 209.18666666666667
|
||||
traj length std: 114.42181532479817
|
||||
traj length min: 98
|
||||
traj length max: 1050
|
||||
action min: -1.0
|
||||
action max: 1.0
|
||||
|
||||
{
|
||||
"env_name": "PickPlaceCan",
|
||||
"env_version": "1.4.1",
|
||||
"type": 1,
|
||||
"env_kwargs": {
|
||||
"has_renderer": false,
|
||||
"has_offscreen_renderer": false,
|
||||
"ignore_done": true,
|
||||
"use_object_obs": true,
|
||||
"use_camera_obs": false,
|
||||
"control_freq": 20,
|
||||
"controller_configs": {
|
||||
"type": "OSC_POSE",
|
||||
"input_max": 1,
|
||||
"input_min": -1,
|
||||
"output_max": [
|
||||
0.05,
|
||||
0.05,
|
||||
0.05,
|
||||
0.5,
|
||||
0.5,
|
||||
0.5
|
||||
],
|
||||
"output_min": [
|
||||
-0.05,
|
||||
-0.05,
|
||||
-0.05,
|
||||
-0.5,
|
||||
-0.5,
|
||||
-0.5
|
||||
],
|
||||
"kp": 150,
|
||||
"damping": 1,
|
||||
"impedance_mode": "fixed",
|
||||
"kp_limits": [
|
||||
0,
|
||||
300
|
||||
],
|
||||
"damping_limits": [
|
||||
0,
|
||||
10
|
||||
],
|
||||
"position_limits": null,
|
||||
"orientation_limits": null,
|
||||
"uncouple_pos_ori": true,
|
||||
"control_delta": true,
|
||||
"interpolation": null,
|
||||
"ramp_ratio": 0.2
|
||||
},
|
||||
"robots": [
|
||||
"Panda"
|
||||
],
|
||||
"camera_depths": false,
|
||||
"camera_heights": 84,
|
||||
"camera_widths": 84,
|
||||
"reward_shaping": false
|
||||
}
|
||||
}
|
||||
|
||||
robomimic dataset normalizes action to [-1, 1], observation roughly? to [-1, 1]. Seems sometimes the upper value is a bit larger than 1 (but within 1.1).
|
||||
|
||||
"""
|
||||
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
import pickle
|
||||
|
||||
try:
|
||||
import h5py # not included in pyproject.toml
|
||||
except:
|
||||
print("Installing h5py")
|
||||
os.system("pip install h5py")
|
||||
import os
|
||||
import random
|
||||
from copy import deepcopy
|
||||
import logging
|
||||
|
||||
|
||||
def make_dataset(
|
||||
load_path,
|
||||
save_dir,
|
||||
save_name_prefix,
|
||||
val_split,
|
||||
normalize,
|
||||
):
|
||||
# Load hdf5 file from load_path
|
||||
with h5py.File(load_path, "r") as f:
|
||||
# put demonstration list in increasing episode order
|
||||
demos = sorted(list(f["data"].keys()))
|
||||
inds = np.argsort([int(elem[5:]) for elem in demos])
|
||||
demos = [demos[i] for i in inds]
|
||||
|
||||
if args.max_episodes > 0:
|
||||
demos = demos[: args.max_episodes]
|
||||
|
||||
# From generate_paper_configs.py: default observation is eef pose, gripper finger position, and object information, all of which are low-dim.
|
||||
low_dim_obs_names = [
|
||||
"robot0_eef_pos",
|
||||
"robot0_eef_quat",
|
||||
"robot0_gripper_qpos",
|
||||
]
|
||||
if "transport" in load_path:
|
||||
low_dim_obs_names += [
|
||||
"robot1_eef_pos",
|
||||
"robot1_eef_quat",
|
||||
"robot1_gripper_qpos",
|
||||
]
|
||||
if args.cameras is None: # state-only
|
||||
low_dim_obs_names.append("object")
|
||||
obs_dim = 0
|
||||
for low_dim_obs_name in low_dim_obs_names:
|
||||
dim = f["data/demo_0/obs/{}".format(low_dim_obs_name)].shape[1]
|
||||
obs_dim += dim
|
||||
logging.info(f"Using {low_dim_obs_name} with dim {dim} for observation")
|
||||
action_dim = f["data/demo_0/actions"].shape[1]
|
||||
logging.info(f"Total low-dim observation dim: {obs_dim}")
|
||||
logging.info(f"Action dim: {action_dim}")
|
||||
|
||||
# get basic stats
|
||||
traj_lengths = []
|
||||
obs_min = np.zeros((obs_dim))
|
||||
obs_max = np.zeros((obs_dim))
|
||||
action_min = np.zeros((action_dim))
|
||||
action_max = np.zeros((action_dim))
|
||||
for ep in demos:
|
||||
traj_lengths.append(f[f"data/{ep}/actions"].shape[0])
|
||||
obs = np.hstack(
|
||||
[
|
||||
f[f"data/{ep}/obs/{low_dim_obs_name}"][()]
|
||||
for low_dim_obs_name in low_dim_obs_names
|
||||
]
|
||||
)
|
||||
actions = f[f"data/{ep}/actions"]
|
||||
obs_min = np.minimum(obs_min, np.min(obs, axis=0))
|
||||
obs_max = np.maximum(obs_max, np.max(obs, axis=0))
|
||||
action_min = np.minimum(action_min, np.min(actions, axis=0))
|
||||
action_max = np.maximum(action_max, np.max(actions, axis=0))
|
||||
traj_lengths = np.array(traj_lengths)
|
||||
max_traj_length = np.max(traj_lengths)
|
||||
|
||||
# report statistics on the data
|
||||
logging.info("===== Basic stats =====")
|
||||
logging.info("total transitions: {}".format(np.sum(traj_lengths)))
|
||||
logging.info("total trajectories: {}".format(traj_lengths.shape[0]))
|
||||
logging.info(
|
||||
f"traj length mean/std: {np.mean(traj_lengths)}, {np.std(traj_lengths)}"
|
||||
)
|
||||
logging.info(
|
||||
f"traj length min/max: {np.min(traj_lengths)}, {np.max(traj_lengths)}"
|
||||
)
|
||||
logging.info(f"obs min: {obs_min}")
|
||||
logging.info(f"obs max: {obs_max}")
|
||||
logging.info(f"action min: {action_min}")
|
||||
logging.info(f"action max: {action_max}")
|
||||
|
||||
# deal with images
|
||||
if args.cameras is not None:
|
||||
img_shapes = []
|
||||
img_names = [] # not necessary but keep old implementation
|
||||
for camera in args.cameras:
|
||||
if f"{camera}_image" in f["data/demo_0/obs"]:
|
||||
img_shape = f["data/demo_0/obs/{}_image".format(camera)].shape[1:]
|
||||
img_shapes.append(img_shape)
|
||||
img_names.append(f"{camera}_image")
|
||||
# ensure all images have the same height and width
|
||||
assert all(
|
||||
[
|
||||
img_shape[0] == img_shapes[0][0]
|
||||
and img_shape[1] == img_shapes[0][1]
|
||||
for img_shape in img_shapes
|
||||
]
|
||||
)
|
||||
combined_img_shape = (
|
||||
img_shapes[0][0],
|
||||
img_shapes[0][1],
|
||||
sum([img_shape[2] for img_shape in img_shapes]),
|
||||
)
|
||||
logging.info(f"Image shapes: {img_shapes}")
|
||||
|
||||
# split indices in train and val
|
||||
num_traj = len(traj_lengths)
|
||||
num_train = int(num_traj * (1 - val_split))
|
||||
train_indices = random.sample(range(num_traj), k=num_train)
|
||||
|
||||
# do over all indices
|
||||
out_train = {}
|
||||
keys = [
|
||||
"observations",
|
||||
"actions",
|
||||
"rewards",
|
||||
]
|
||||
if args.cameras is not None:
|
||||
keys.append("images")
|
||||
out_train["observations"] = np.empty((0, max_traj_length, obs_dim))
|
||||
out_train["actions"] = np.empty((0, max_traj_length, action_dim))
|
||||
out_train["rewards"] = np.empty((0, max_traj_length))
|
||||
out_train["traj_length"] = []
|
||||
if args.cameras is not None:
|
||||
out_train["images"] = np.empty(
|
||||
(
|
||||
0,
|
||||
max_traj_length,
|
||||
*combined_img_shape,
|
||||
),
|
||||
dtype=np.uint8,
|
||||
)
|
||||
out_val = deepcopy(out_train)
|
||||
train_episode_reward_all = []
|
||||
val_episode_reward_all = []
|
||||
for i in tqdm(range(len(demos))):
|
||||
ep = demos[i]
|
||||
if i in train_indices:
|
||||
out = out_train
|
||||
else:
|
||||
out = out_val
|
||||
|
||||
# get episode length
|
||||
traj_length = f[f"data/{ep}"].attrs["num_samples"]
|
||||
out["traj_length"].append(traj_length)
|
||||
# print("Episode:", i, "Trajectory length:", traj_length)
|
||||
|
||||
# extract
|
||||
raw_actions = f[f"data/{ep}/actions"][()]
|
||||
rewards = f[f"data/{ep}/rewards"][()]
|
||||
raw_obs = np.hstack(
|
||||
[
|
||||
f[f"data/{ep}/obs/{low_dim_obs_name}"][()]
|
||||
for low_dim_obs_name in low_dim_obs_names
|
||||
]
|
||||
) # not normalized
|
||||
|
||||
# scale to [-1, 1] for both ob and action
|
||||
if normalize:
|
||||
obs = 2 * (raw_obs - obs_min) / (obs_max - obs_min + 1e-6) - 1
|
||||
actions = (
|
||||
2 * (raw_actions - action_min) / (action_max - action_min + 1e-6)
|
||||
- 1
|
||||
)
|
||||
else:
|
||||
obs = raw_obs
|
||||
actions = raw_actions
|
||||
|
||||
data_traj = {
|
||||
"observations": obs,
|
||||
"actions": actions,
|
||||
"rewards": rewards,
|
||||
}
|
||||
if args.cameras is not None: # no normalization
|
||||
data_traj["images"] = np.concatenate(
|
||||
(
|
||||
[
|
||||
f["data/{}/obs/{}".format(ep, img_name)][()]
|
||||
for img_name in img_names
|
||||
]
|
||||
),
|
||||
axis=-1,
|
||||
)
|
||||
|
||||
# apply padding to make all episodes have the same max steps
|
||||
# later when we load this dataset, we will use the traj_length to slice the data
|
||||
for key in keys:
|
||||
traj = data_traj[key]
|
||||
if traj.ndim == 1:
|
||||
pad_width = (0, max_traj_length - len(traj))
|
||||
elif traj.ndim == 2:
|
||||
pad_width = ((0, max_traj_length - traj.shape[0]), (0, 0))
|
||||
elif traj.ndim == 4:
|
||||
pad_width = (
|
||||
(0, max_traj_length - traj.shape[0]),
|
||||
(0, 0),
|
||||
(0, 0),
|
||||
(0, 0),
|
||||
)
|
||||
else:
|
||||
raise ValueError("Unsupported dimension")
|
||||
traj = np.pad(
|
||||
traj,
|
||||
pad_width,
|
||||
mode="constant",
|
||||
constant_values=0,
|
||||
)
|
||||
out[key] = np.vstack((out[key], traj[None]))
|
||||
|
||||
# check reward
|
||||
if i in train_indices:
|
||||
train_episode_reward_all.append(np.sum(data_traj["rewards"]))
|
||||
else:
|
||||
val_episode_reward_all.append(np.sum(data_traj["rewards"]))
|
||||
|
||||
# Save to np file
|
||||
save_train_path = os.path.join(save_dir, save_name_prefix + "train.pkl")
|
||||
save_val_path = os.path.join(save_dir, save_name_prefix + "val.pkl")
|
||||
with open(save_train_path, "wb") as f:
|
||||
pickle.dump(out_train, f)
|
||||
with open(save_val_path, "wb") as f:
|
||||
pickle.dump(out_val, f)
|
||||
if normalize:
|
||||
normalization_save_path = os.path.join(
|
||||
save_dir, save_name_prefix + "normalization.npz"
|
||||
)
|
||||
np.savez(
|
||||
normalization_save_path,
|
||||
obs_min=obs_min,
|
||||
obs_max=obs_max,
|
||||
action_min=action_min,
|
||||
action_max=action_max,
|
||||
)
|
||||
|
||||
# debug
|
||||
logging.info("\n========== Final ===========")
|
||||
logging.info(
|
||||
f"Train - Number of episodes and transitions: {len(out_train['traj_length'])}, {np.sum(out_train['traj_length'])}"
|
||||
)
|
||||
logging.info(
|
||||
f"Val - Number of episodes and transitions: {len(out_val['traj_length'])}, {np.sum(out_val['traj_length'])}"
|
||||
)
|
||||
logging.info(
|
||||
f"Train - Mean/Std trajectory length: {np.mean(out_train['traj_length'])}, {np.std(out_train['traj_length'])}"
|
||||
)
|
||||
logging.info(
|
||||
f"Train - Max/Min trajectory length: {np.max(out_train['traj_length'])}, {np.min(out_train['traj_length'])}"
|
||||
)
|
||||
logging.info(
|
||||
f"Train - Mean/Std episode reward: {np.mean(train_episode_reward_all)}, {np.std(train_episode_reward_all)}"
|
||||
)
|
||||
if val_split > 0:
|
||||
logging.info(
|
||||
f"Val - Mean/Std trajectory length: {np.mean(out_val['traj_length'])}, {np.std(out_val['traj_length'])}"
|
||||
)
|
||||
logging.info(
|
||||
f"Val - Max/Min trajectory length: {np.max(out_val['traj_length'])}, {np.min(out_val['traj_length'])}"
|
||||
)
|
||||
logging.info(
|
||||
f"Val - Mean/Std episode reward: {np.mean(val_episode_reward_all)}, {np.std(val_episode_reward_all)}"
|
||||
)
|
||||
for obs_dim_ind in range(obs_dim):
|
||||
obs = out_train["observations"][:, :, obs_dim_ind]
|
||||
logging.info(
|
||||
f"Train - Obs dim {obs_dim_ind+1} mean {np.mean(obs)} std {np.std(obs)} min {np.min(obs)} max {np.max(obs)}"
|
||||
)
|
||||
for action_dim_ind in range(action_dim):
|
||||
action = out_train["actions"][:, :, action_dim_ind]
|
||||
logging.info(
|
||||
f"Train - Action dim {action_dim_ind+1} mean {np.mean(action)} std {np.std(action)} min {np.min(action)} max {np.max(action)}"
|
||||
)
|
||||
if val_split > 0:
|
||||
for obs_dim_ind in range(obs_dim):
|
||||
obs = out_val["observations"][:, :, obs_dim_ind]
|
||||
logging.info(
|
||||
f"Val - Obs dim {obs_dim_ind+1} mean {np.mean(obs)} std {np.std(obs)} min {np.min(obs)} max {np.max(obs)}"
|
||||
)
|
||||
for action_dim_ind in range(action_dim):
|
||||
action = out_val["actions"][:, :, action_dim_ind]
|
||||
logging.info(
|
||||
f"Val - Action dim {action_dim_ind+1} mean {np.mean(action)} std {np.std(action)} min {np.min(action)} max {np.max(action)}"
|
||||
)
|
||||
# logging.info("Train - Observation shape:", out_train["observations"].shape)
|
||||
# logging.info("Train - Action shape:", out_train["actions"].shape)
|
||||
# logging.info("Train - Reward shape:", out_train["rewards"].shape)
|
||||
# logging.info("Val - Observation shape:", out_val["observations"].shape)
|
||||
# logging.info("Val - Action shape:", out_val["actions"].shape)
|
||||
# logging.info("Val - Reward shape:", out_val["rewards"].shape)
|
||||
# if use_img:
|
||||
# logging.info("Image shapes:", img_shapes)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--load_path", type=str, default=".")
|
||||
parser.add_argument("--save_dir", type=str, default=".")
|
||||
parser.add_argument("--save_name_prefix", type=str, default="")
|
||||
parser.add_argument("--val_split", type=float, default="0.2")
|
||||
parser.add_argument("--max_episodes", type=int, default="-1")
|
||||
parser.add_argument("--normalize", action="store_true")
|
||||
parser.add_argument("--cameras", nargs="*", default=None)
|
||||
args = parser.parse_args()
|
||||
|
||||
import datetime
|
||||
|
||||
if args.max_episodes > 0:
|
||||
args.save_name_prefix += f"max_episodes_{args.max_episodes}_"
|
||||
|
||||
os.makedirs(args.save_dir, exist_ok=True)
|
||||
log_path = os.path.join(
|
||||
args.save_dir,
|
||||
args.save_name_prefix
|
||||
+ f"_{datetime.datetime.now().strftime('%Y_%m_%d_%H_%M_%S')}.log",
|
||||
)
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(message)s",
|
||||
handlers=[
|
||||
logging.FileHandler(log_path, mode="w"),
|
||||
logging.StreamHandler(),
|
||||
],
|
||||
)
|
||||
|
||||
make_dataset(
|
||||
args.load_path,
|
||||
args.save_dir,
|
||||
args.save_name_prefix,
|
||||
args.val_split,
|
||||
args.normalize,
|
||||
)
|
||||
@@ -0,0 +1,432 @@
|
||||
def get_dataset_download_url(cfg):
|
||||
env = cfg.env
|
||||
# Gym
|
||||
if env == "hopper-medium-v2":
|
||||
return "https://drive.google.com/drive/u/1/folders/18Ti-92XVq3sE24K096WAxjC_SCCngeHd"
|
||||
elif env == "walker2d-medium-v2":
|
||||
return "https://drive.google.com/drive/u/1/folders/1BJu8NklriunDHsDrLT6fEpcro3_2IPFf"
|
||||
elif env == "halfcheetah-medium-v2":
|
||||
return "https://drive.google.com/drive/u/1/folders/1Drel26tiuQ9oD3YNf1eyy0UVaf5SQj-U"
|
||||
# D3IL
|
||||
elif env == "avoid" and cfg.mode == "d56_r12": # M1
|
||||
return "https://drive.google.com/drive/u/1/folders/1ZAPvLQwv2y4Q98UDVKXFT4fvGF5yhD_o"
|
||||
elif env == "avoid" and cfg.mode == "d57_r12": # M2
|
||||
return "https://drive.google.com/drive/u/1/folders/1wyJi1Zbnd6JNy4WGszHBH40A0bbl-vkd"
|
||||
elif env == "avoid" and cfg.mode == "d58_r12": # M3
|
||||
return "https://drive.google.com/drive/u/1/folders/1mNXCIPnCO_FDBlEj95InA9eWJM2XcEEj"
|
||||
# Robomimic
|
||||
elif env == "lift" and "img" not in cfg.train_dataset_path: # state
|
||||
return "https://drive.google.com/drive/u/1/folders/1lbXgMKBTAiFdJqPZqWXpwjEyrVW16MBu"
|
||||
elif env == "lift" and "img" in cfg.train_dataset_path: # img
|
||||
return "https://drive.google.com/drive/u/1/folders/1H-UncdzHx6wd5NWVzrQyftfGls7KGz1O"
|
||||
elif env == "can" and "img" not in cfg.train_dataset_path:
|
||||
return "https://drive.google.com/drive/u/1/folders/1J1qSvsDEf40jnMZY9W0r6ww--E3MdmK3"
|
||||
elif env == "can" and "img" in cfg.train_dataset_path:
|
||||
return "https://drive.google.com/drive/u/1/folders/1VGp_5xXXb1-GJutdSc6AZSzXNk-6_vRz"
|
||||
elif env == "square" and "img" not in cfg.train_dataset_path:
|
||||
return "https://drive.google.com/drive/u/1/folders/1mVVNOJ6wt2EXoapF7PKkcqxbsB9gvK-B"
|
||||
elif env == "square" and "img" in cfg.train_dataset_path:
|
||||
return "https://drive.google.com/drive/u/1/folders/1-aGqVeKLIzJCEst8p0ZTjfjkrXFfJLxa"
|
||||
elif env == "transport" and "img" not in cfg.train_dataset_path:
|
||||
return "https://drive.google.com/drive/u/1/folders/1EVHmFx-YdX4MEE1EjwduVayvaH9vvqvK"
|
||||
elif env == "transport" and "img" in cfg.train_dataset_path:
|
||||
return "https://drive.google.com/drive/u/1/folders/1cOkAZQmmETYEPFrnnX0EuD6mv0kUfMO2"
|
||||
# Furniture-Bench
|
||||
elif env == "one_leg_low_dim":
|
||||
return "https://drive.google.com/drive/u/1/folders/1v4LG2D1fS8id5hqE7Jjt7MYNFUEBNyh4"
|
||||
elif env == "one_leg_med_dim":
|
||||
return "https://drive.google.com/drive/u/1/folders/1ohDuMSgCqGN1CSh1cI_8A0ia3DVHzj3w"
|
||||
elif env == "lamp_low_dim":
|
||||
return "https://drive.google.com/drive/u/1/folders/14MqDUmuNmTFcBtKw5gcvx7nuir0zmF7V"
|
||||
elif env == "lamp_med_dim":
|
||||
return "https://drive.google.com/drive/u/1/folders/1bhOoN0xet4ga0rOIvcRHUFoaCjPRNRLf"
|
||||
elif env == "round_table_low_dim":
|
||||
return "https://drive.google.com/drive/u/1/folders/15oF3qiqGzlT_98FDoTIVtGmLtHd5SIbi"
|
||||
elif env == "round_table_med_dim":
|
||||
return "https://drive.google.com/drive/u/1/folders/1U27xjdRrKlLC8E33o7jMFZ1HF5P_Soik"
|
||||
# unknown
|
||||
else:
|
||||
raise ValueError(f"Unknown environment {env}")
|
||||
|
||||
|
||||
def get_normalization_download_url(cfg):
|
||||
env = cfg.env_name
|
||||
# Gym
|
||||
if env == "hopper-medium-v2":
|
||||
return "https://drive.google.com/file/d/1HHZ2X6r5io6hjG-MHVFFoPMJV2fJ3lis/view?usp=drive_link"
|
||||
elif env == "walker2d-medium-v2":
|
||||
return "https://drive.google.com/file/d/1NSX7t3DFKaBj5HNpv91Oo5h6oXTk0zoo/view?usp=drive_link"
|
||||
elif env == "halfcheetah-medium-v2":
|
||||
return "https://drive.google.com/file/d/1LlwCMfy1b5e8jSx99CV3lWhcrQWrI2Jm/view?usp=drive_link"
|
||||
# D3IL
|
||||
elif env == "avoiding-m5" and cfg.mode == "d56_r12": # M1
|
||||
return "https://drive.google.com/file/d/1PubKaPabbiSdWYpGmouDhYfXp4QwNHFG/view?usp=drive_link"
|
||||
elif env == "avoiding-m5" and cfg.mode == "d57_r12": # M2
|
||||
return "https://drive.google.com/file/d/1Hoohw8buhsLzXoqivMA6IzKS5Izlj07_/view?usp=drive_link"
|
||||
elif env == "avoiding-m5" and cfg.mode == "d58_r12": # M3
|
||||
return "https://drive.google.com/file/d/1qt7apV52C9Tflsc-A55J6uDMHzaFa1wN/view?usp=drive_link"
|
||||
# Robomimic
|
||||
elif env == "lift" and "img" not in cfg.normalization_path: # state
|
||||
return "https://drive.google.com/file/d/1d3WjwRds-7I5bBFpZuY27OT9ycb8r_QM/view?usp=drive_link"
|
||||
elif env == "lift" and "img" in cfg.normalization_path: # img
|
||||
return "https://drive.google.com/file/d/15GnKDIK8VasvUHahcvEeK_uEs1J9i0ja/view?usp=drive_link"
|
||||
elif env == "can" and "img" not in cfg.normalization_path:
|
||||
return "https://drive.google.com/file/d/14FxHk9zQ-5ulAO26a6xrdvc-gRkL36FR/view?usp=drive_link"
|
||||
elif env == "can" and "img" in cfg.normalization_path:
|
||||
return "https://drive.google.com/file/d/1APAB6W10ECaNVVL72F0C2wC6oJhhbjWX/view?usp=drive_link"
|
||||
elif env == "square" and "img" not in cfg.normalization_path:
|
||||
return "https://drive.google.com/file/d/1FFMqWVv0145OJjbA_iglkWywmdK22Za-/view?usp=drive_link"
|
||||
elif env == "square" and "img" in cfg.normalization_path:
|
||||
return "https://drive.google.com/file/d/1jq5atfHdu-ZMQ8YaFjwcctbJESBYDgP8/view?usp=drive_link"
|
||||
elif env == "transport" and "img" not in cfg.normalization_path:
|
||||
return "https://drive.google.com/file/d/1EmC80gIgLoqQ8kRPH5r0mDVqX6NqHO3p/view?usp=drive_link"
|
||||
elif env == "transport" and "img" in cfg.normalization_path:
|
||||
return "https://drive.google.com/file/d/1LBgvIacNzbXCZXWKYanddqiotevG3hmA/view?usp=drive_link"
|
||||
# Furniture-Bench
|
||||
elif env == "one_leg_low_dim":
|
||||
return "https://drive.google.com/file/d/1fbYxau8Z0tifeuu_06UKdRRz8zozxMUE/view?usp=drive_link"
|
||||
elif env == "one_leg_med_dim":
|
||||
return "https://drive.google.com/file/d/1bT-MIG99-uwSXL3VwCz6kK9Kg9t6a7BZ/view?usp=drive_link"
|
||||
elif env == "lamp_low_dim":
|
||||
return "https://drive.google.com/file/d/1KRyCcCT63-cikz6Q_McXM0iVYx9M2E_o/view?usp=drive_link"
|
||||
elif env == "lamp_med_dim":
|
||||
return "https://drive.google.com/file/d/1WBHKbA6BtSfF3qfRIkyuWga5T9X3fEvg/view?usp=drive_link"
|
||||
elif env == "round_table_low_dim":
|
||||
return "https://drive.google.com/file/d/1CwfNLI5KXEkl_jVgtAY7-xfq7jED-xc0/view?usp=drive_link"
|
||||
elif env == "round_table_med_dim":
|
||||
return "https://drive.google.com/file/d/1vNxzok6f6HROGxQR2eoThsN62hVjpayn/view?usp=drive_link"
|
||||
# unknown
|
||||
else:
|
||||
raise ValueError(f"Unknown environment {env}")
|
||||
|
||||
|
||||
def get_checkpoint_download_url(cfg):
|
||||
path = cfg.base_policy_path
|
||||
######################################
|
||||
#### Gym
|
||||
######################################
|
||||
if (
|
||||
"hopper-medium-v2_pre_diffusion_mlp_ta4_td20/2024-06-12_23-10-05/checkpoint/state_3000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1uV3beg2YuBRh11t7jnRsyd_M9bsFGOkA/view?usp=drive_link"
|
||||
elif (
|
||||
"walker2d-medium-v2_pre_diffusion_mlp_ta4_td20/2024-06-12_23-06-12/checkpoint/state_3000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1StopetttozWba4l9u0VT_JW1rYe6kJx2/view?usp=drive_link"
|
||||
elif (
|
||||
"halfcheetah-medium-v2_pre_diffusion_mlp_ta4_td20/2024-06-12_23-04-42/checkpoint/state_3000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1o9ryyeZQAsaB4ffUTCJkIaGCi0frL3G4/view?usp=drive_link"
|
||||
######################################
|
||||
#### D3IL
|
||||
######################################
|
||||
elif (
|
||||
"avoid_d56_r12_pre_diffusion_mlp_ta4_td20/2024-07-06_22-50-07/checkpoint/state_10000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1JdEOG0KsCA9EX9zq09DE0xkTB4xy6DNp/view?usp=drive_link"
|
||||
elif (
|
||||
"avoid_d56_r12_pre_gaussian_mlp_ta4/2024-07-07_01-35-48/checkpoint/state_10000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/138wEi_rVV5HpcwgH6_3BlXQ1dhgZN05L/view?usp=drive_link"
|
||||
elif (
|
||||
"avoid_d56_r12_pre_gmm_mlp_ta4/2024-07-10_14-30-00/checkpoint/state_10000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1krvEINP6UBfnJG9bEge53h3MNJkTQXe_/view?usp=drive_link"
|
||||
elif (
|
||||
"avoid_d57_r12_pre_diffusion_mlp_ta4_td20/2024-07-07_13-12-09/checkpoint/state_15000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1wAmRHzRZ4O5z_ZWDhVg5JFZbugo4cHKA/view?usp=drive_link"
|
||||
elif (
|
||||
"avoid_d57_r12_pre_gaussian_mlp_ta4/2024-07-07_02-15-50/checkpoint/state_10000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1hf_647bJ0EMRhArsfxStSkhgkGaWGhmD/view?usp=drive_link"
|
||||
elif (
|
||||
"avoid_d57_r12_pre_gmm_mlp_ta4/2024-07-10_15-44-32/checkpoint/state_10000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1CE4AcNJp2UITIHpuLUwxp7cDH8Y8jsYC/view?usp=drive_link"
|
||||
elif (
|
||||
"avoid_d58_r12_pre_diffusion_mlp_ta4_td20/2024-07-07_13-54-54/checkpoint/state_15000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1w5X1lJZd0wI6E2XRdj839TqJ_9EziGRx/view?usp=drive_link"
|
||||
elif (
|
||||
"avoid_d58_r12_pre_gaussian_mlp_ta4/2024-07-07_13-11-49/checkpoint/state_10000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1YIVEN0Ykica9dj_DxxPrB4QJLgv-Ne0N/view?usp=drive_link"
|
||||
elif (
|
||||
"avoid_d58_r12_pre_gmm_mlp_ta4/2024-07-10_17-01-50/checkpoint/state_10000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/174tadrqjfxJdOgsMjNbg53goZtIkpW3f/view?usp=drive_link"
|
||||
######################################
|
||||
#### Robomimic-Lift
|
||||
######################################
|
||||
elif (
|
||||
"lift_pre_diffusion_unet_ta4_td20/2024-06-29_02-49-45/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1T-NGgBmT-UmcVWADygXj873IyWLewvsU/view?usp=drive_link"
|
||||
elif (
|
||||
"lift_pre_diffusion_mlp_ta4_td20/2024-06-28_14-47-58/checkpoint/state_5000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1Ngr-DNxoB9XNCZ2O-NF5p60NzmYlzmWG/view?usp=drive_link"
|
||||
elif (
|
||||
"lift_pre_diffusion_mlp_img_ta4_td100/2024-07-30_22-24-35/checkpoint/state_2500.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/19hqNicwKKKDrlS5UMr-FLRu51EAW8Z51/view?usp=drive_link"
|
||||
elif (
|
||||
"lift_pre_gaussian_mlp_ta4/2024-06-28_14-48-24/checkpoint/state_5000.pt" in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/157x5_XJy3ZyaPz_7Vr_opXAZ_mPvlmhp/view?usp=drive_link"
|
||||
elif (
|
||||
"lift_pre_gaussian_mlp_img_ta4/2024-07-28_23-00-48/checkpoint/state_500.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1Uae7K3Hv9XzaAGljG2fjsxYyOEgL03zv/view?usp=drive_link"
|
||||
elif (
|
||||
"lift_pre_gaussian_transformer_ta4/2024-06-28_14-49-23/checkpoint/state_5000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1Z_C8ureDDPXqpDUaRMVKBt3VgfZlIQdE/view?usp=drive_link"
|
||||
elif "lift_pre_gmm_mlp_ta4/2024-06-28_15-30-32/checkpoint/state_5000.pt" in path:
|
||||
return "https://drive.google.com/file/d/1wFvBoIOaCJVibqEIYSjJxAbxjGkH7RMk/view?usp=drive_link"
|
||||
elif (
|
||||
"lift_pre_gmm_transformer_ta4/2024-06-28_14-51-23/checkpoint/state_5000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1w_3WOXS51debWc1ShgaO2VQx49dj9ky2/view?usp=drive_link"
|
||||
######################################
|
||||
#### Robomimic-Can
|
||||
######################################
|
||||
elif (
|
||||
"can_pre_diffusion_unet_ta4_td20/2024-06-29_02-49-45/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1s346KCe2aar_tXX7u8rzjRF3kpwVpH5c/view?usp=drive_link"
|
||||
elif (
|
||||
"can_pre_diffusion_mlp_ta4_td20/2024-06-28_13-29-54/checkpoint/state_5000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1L1ZLD1u1Y1YJmRLGzScXbQ02wGS-_cWo/view?usp=drive_link"
|
||||
elif (
|
||||
"can_pre_diffusion_mlp_img_ta4_td100/2024-07-30_22-23-55/checkpoint/state_5000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1siIKDGVHld3ZH8vDqgu9iE5H9KoaYq_Z/view?usp=drive_link"
|
||||
elif (
|
||||
"can_pre_gaussian_mlp_ta4/2024-06-28_13-31-00/checkpoint/state_5000.pt" in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1bA-A0p0KnHwrVO3MqjuYdV3dZuuZICMy/view?usp=drive_link"
|
||||
elif (
|
||||
"can_pre_gaussian_mlp_img_ta4/2024-07-28_21-54-40/checkpoint/state_1000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/16vUUyVO9DvnDyPtSGiZx4iQP4JxuSAoS/view?usp=drive_link"
|
||||
elif (
|
||||
"can_pre_gaussian_transformer_ta4/2024-06-28_13-42-20/checkpoint/state_5000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1cGf7s8aS5grZsGRb5PYZiPogcPRr_lrf/view?usp=drive_link"
|
||||
elif "can_pre_gmm_mlp_ta4/2024-06-28_13-32-19/checkpoint/state_5000.pt" in path:
|
||||
return "https://drive.google.com/file/d/1KVx8-KICiIHstcjsvhPlQc_pZZ6RxXpd/view?usp=drive_link"
|
||||
elif (
|
||||
"can_pre_gmm_transformer_ta4/2024-06-28_13-43-21/checkpoint/state_5000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1xSgwGG40zdoO2DDSM79l0rMHeNmaifnq/view?usp=drive_link"
|
||||
######################################
|
||||
#### Robomimic-Square
|
||||
######################################
|
||||
elif (
|
||||
"square_pre_diffusion_unet_ta4_td20/2024-06-29_02-48-45/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/11IEgQe0LFI23hn1Cwf6Z_YfJdDilVc0z/view?usp=drive_link"
|
||||
elif (
|
||||
"square_pre_diffusion_mlp_ta4_td20/2024-07-10_01-46-16/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1lP9mNe2AxMigfOywcaHOOR7FxQ-KR_Ee/view?usp=drive_link"
|
||||
elif (
|
||||
"square_pre_diffusion_mlp_img_ta4_td100/2024-07-30_22-27-34/checkpoint/state_4000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1miud5SX41xjPoW8yqClRhaFx0Q7EWl3V/view?usp=drive_link"
|
||||
elif (
|
||||
"square_pre_gaussian_mlp_ta4/2024-06-28_15-02-32/checkpoint/state_5000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1rETiyXLz7YgYoHKwLEZa7dFabu4gfJR8/view?usp=drive_link"
|
||||
elif (
|
||||
"square_pre_gaussian_mlp_img_ta4/2024-07-30_18-44-32/checkpoint/state_4000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1myB6FOAmt6c6x3ScGKXRULgx826tZVS4/view?usp=drive_link"
|
||||
elif (
|
||||
"square_pre_gaussian_transformer_ta4/2024-06-28_15-02-39/checkpoint/state_5000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1JJQ7KbRWWBB09PLwNRriAUEAl9vjFeCW/view?usp=drive_link"
|
||||
elif "square_pre_gmm_mlp_ta4/2024-06-28_15-03-08/checkpoint/state_5000.pt" in path:
|
||||
return "https://drive.google.com/file/d/10ujnnOk2Pn-yjE7iW9-RYbApo_aKiNRq/view?usp=drive_link"
|
||||
elif (
|
||||
"square_pre_gmm_transformer_ta4/2024-06-28_15-03-15/checkpoint/state_5000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1LczXhgeNtQfqySsfGNbbviPrlLwyh-E3/view?usp=drive_link"
|
||||
######################################
|
||||
#### Robomimic-Transport
|
||||
######################################
|
||||
elif (
|
||||
"transport_pre_diffusion_unet_ta16_td20/2024-07-04_02-20-53/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1MNGT8j9x1uudugGUcia-xwP_7f7xVY4K/view?usp=drive_link"
|
||||
elif (
|
||||
"transport_pre_diffusion_mlp_ta8_td20/2024-07-08_11-18-59/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1240FmcDPg_VyXReEtePjBN4MML-OT21C/view?usp=drive_link"
|
||||
elif (
|
||||
"transport_pre_diffusion_mlp_img_ta8_td100/2024-07-30_22-30-06/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1_UmjNv27_w49XC_EO1WexgreJmC0SEdF/view?usp=drive_link"
|
||||
elif (
|
||||
"transport_pre_gaussian_mlp_ta8/2024-07-10_01-50-52/checkpoint/state_5000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1NOHDNHu1sTabxBd4DZrSrTyX9l1I2GuO/view?usp=drive_link"
|
||||
elif (
|
||||
"transport_pre_gaussian_mlp_img_ta8/2024-07-30_21-39-34/checkpoint/state_5000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1XYOIwOEgOVoUdMmxWuBRAcA5BOS777Kc/view?usp=drive_link"
|
||||
elif (
|
||||
"transport_pre_gaussian_transformer_ta8/2024-06-28_15-18-16/checkpoint/state_5000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1hnt9lX5bg82iFFsAk5FY3TktPXsAhUFU/view?usp=drive_link"
|
||||
elif (
|
||||
"transport_pre_gmm_mlp_ta8/2024-07-10_01-51-21/checkpoint/state_5000.pt" in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1da9yLIu5ahq-ZgIIsG7wqehA5VopFTnt/view?usp=drive_link"
|
||||
elif (
|
||||
"transport_pre_gmm_transformer_ta8/2024-06-28_15-18-43/checkpoint/state_5000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1c0S7WX-U1Kn6n-wWZEQSR1alLPW2eZmi/view?usp=drive_link"
|
||||
######################################
|
||||
#### Furniture-One_leg
|
||||
######################################
|
||||
elif (
|
||||
"one_leg_low_dim_pre_diffusion_mlp_ta8_td100/2024-07-22_20-01-16/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1tP0i53EMwNyw_bS2lt9WtQDIAoTxn40g/view?usp=drive_link"
|
||||
elif (
|
||||
"one_leg_low_dim_pre_diffusion_unet_ta16_td100/2024-07-03_22-23-38/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/13nOr7EI79RqdRuoc-je0LoTX5Y-IaO9J/view?usp=drive_link"
|
||||
elif (
|
||||
"one_leg_low_dim_pre_gaussian_mlp_ta8/2024-06-26_23-43-02/checkpoint/state_3000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1DeHj-IEX3a2ZLXFlV0MYe1N_5qupod_w/view?usp=drive_link"
|
||||
elif (
|
||||
"one_leg_med_dim_pre_diffusion_mlp_ta8_td100/2024-07-23_01-28-11/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1aEzObl4EkOBKSs0wI2MGmkslzIjZy2_L/view?usp=drive_link"
|
||||
elif (
|
||||
"one_leg_med_dim_pre_diffusion_unet_ta16_td100/2024-07-04_02-16-16/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1pSwp_IUDSQ15OszChCkrrTQYdvG3pHc1/view?usp=drive_link"
|
||||
elif (
|
||||
"one_leg_med_dim_pre_gaussian_mlp_ta8/2024-06-28_16-27-02/checkpoint/state_3000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1lC4PKRPn4tRWR9VW5fSfCu3MKnRvYqD6/view?usp=drive_link"
|
||||
######################################
|
||||
#### Furniture-Lamp
|
||||
######################################
|
||||
elif (
|
||||
"lamp_low_dim_pre_diffusion_mlp_ta8_td100/2024-07-23_01-28-20/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1md5m4vpe5MmStu-fOk8snF-52BYGZpSc/view?usp=drive_link"
|
||||
elif (
|
||||
"lamp_low_dim_pre_diffusion_unet_ta16_td100/2024-07-04_02-16-48/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/103lUuvyKPvp97hzBYnUDAWgor6LsUf21/view?usp=drive_link"
|
||||
elif (
|
||||
"lamp_low_dim_pre_gaussian_mlp_ta8/2024-06-28_16-26-51/checkpoint/state_3000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1tK6NuZf4_xtTlZksIR9oQxiFEADFRVA9/view?usp=drive_link"
|
||||
elif (
|
||||
"lamp_med_dim_pre_diffusion_mlp_ta8_td100/2024-07-23_01-28-20/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1yVUXYCK_vFhQ7mawxnCGxuHS6RgWWDjj/view?usp=drive_link"
|
||||
elif (
|
||||
"lamp_med_dim_pre_diffusion_unet_ta16_td100/2024-07-04_02-17-21/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1_qd47U50on-T1Pqojlfy7S3Cl6bLm_D1/view?usp=drive_link"
|
||||
elif (
|
||||
"lamp_med_dim_pre_gaussian_mlp_ta8/2024-06-28_16-26-56/checkpoint/state_3000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1d16HmHgidtXForqoN5QJg_e152XCaJM5/view?usp=drive_link"
|
||||
######################################
|
||||
#### Furniture-Round_table
|
||||
######################################
|
||||
elif (
|
||||
"round_table_low_dim_pre_diffusion_mlp_ta8_td100/2024-07-23_01-28-26/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1iJyqJr84AtszGqPeSysN70A-SmP73mKu/view?usp=drive_link"
|
||||
elif (
|
||||
"round_table_low_dim_pre_diffusion_unet_ta16_td100/2024-07-04_02-19-48/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1F3RFgLcFemU-IDUrLCXMzmQXknC2ISgJ/view?usp=drive_link"
|
||||
elif (
|
||||
"round_table_low_dim_pre_gaussian_mlp_ta8/2024-06-28_16-26-51/checkpoint/state_3000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1eX2u0cvu_zveblrg2htFomAxSToJmHyL/view?usp=drive_link"
|
||||
elif (
|
||||
"round_table_med_dim_pre_diffusion_mlp_ta8_td100/2024-07-23_01-28-29/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1DDvHEqb7WqVgT3_W6gYCNDQrvGyreatP/view?usp=drive_link"
|
||||
elif (
|
||||
"round_table_med_dim_pre_diffusion_unet_ta16_td100/2024-07-04_02-20-21/checkpoint/state_8000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1hzungyvt3Uc-XbrCzVghzdr59zDf6FlD/view?usp=drive_link"
|
||||
elif (
|
||||
"round_table_med_dim_pre_gaussian_mlp_ta8/2024-06-28_16-26-38/checkpoint/state_3000.pt"
|
||||
in path
|
||||
):
|
||||
return "https://drive.google.com/file/d/1LvELZ7A-whxKk1Oq7S8M9uBiSrWPAGQT/view?usp=drive_link"
|
||||
# unknown --- this means the user trained own policy but specifies the wrong path
|
||||
else:
|
||||
return None
|
||||
Executable
+50
@@ -0,0 +1,50 @@
|
||||
#!/bin/bash
|
||||
|
||||
##################### Paths #####################
|
||||
|
||||
# Set default paths
|
||||
DEFAULT_DATA_DIR="${PWD}/data"
|
||||
DEFAULT_LOG_DIR="${PWD}/log"
|
||||
|
||||
# Prompt the user for input, allowing overrides
|
||||
read -p "Enter the desired data directory [default: ${DEFAULT_DATA_DIR}], leave empty to use default: " DATA_DIR
|
||||
DPPO_DATA_DIR=${DATA_DIR:-$DEFAULT_DATA_DIR} # Use user input or default if input is empty
|
||||
|
||||
read -p "Enter the desired logging directory [default: ${DEFAULT_LOG_DIR}], leave empty to use default: " LOG_DIR
|
||||
DPPO_LOG_DIR=${LOG_DIR:-$DEFAULT_LOG_DIR} # Use user input or default if input is empty
|
||||
|
||||
# Export to current session
|
||||
export DPPO_DATA_DIR="$DPPO_DATA_DIR"
|
||||
export DPPO_LOG_DIR="$DPPO_LOG_DIR"
|
||||
|
||||
# Confirm the paths with the user
|
||||
echo "Data directory set to: $DPPO_DATA_DIR"
|
||||
echo "Log directory set to: $DPPO_LOG_DIR"
|
||||
|
||||
# Append environment variables to .bashrc
|
||||
echo "export DPPO_DATA_DIR=\"$DPPO_DATA_DIR\"" >> ~/.bashrc
|
||||
echo "export DPPO_LOG_DIR=\"$DPPO_LOG_DIR\"" >> ~/.bashrc
|
||||
|
||||
echo "Environment variables DPPO_DATA_DIR and DPPO_LOG_DIR added to .bashrc and applied to the current session."
|
||||
|
||||
##################### WandB #####################
|
||||
|
||||
# Prompt the user for input, allowing overrides
|
||||
read -p "Enter your WandB entity (username or team name), leave empty to skip: " ENTITY
|
||||
|
||||
# Check if ENTITY is not empty
|
||||
if [ -n "$ENTITY" ]; then
|
||||
# If ENTITY is not empty, set the environment variable
|
||||
export DPPO_WANDB_ENTITY="$ENTITY"
|
||||
|
||||
# Confirm the entity with the user
|
||||
echo "WandB entity set to: $DPPO_WANDB_ENTITY"
|
||||
|
||||
# Append environment variable to .bashrc
|
||||
echo "export DPPO_WANDB_ENTITY=\"$ENTITY\"" >> ~/.bashrc
|
||||
|
||||
echo "Environment variable DPPO_WANDB_ENTITY added to .bashrc and applied to the current session."
|
||||
else
|
||||
# If ENTITY is empty, skip setting the environment variable
|
||||
echo "No WandB entity provided. Please set wandb=null when running scripts to disable wandb logging and avoid error."
|
||||
fi
|
||||
@@ -0,0 +1,34 @@
|
||||
"""
|
||||
Visualize Avoid environment from D3IL in MuJoCo GUI
|
||||
|
||||
"""
|
||||
|
||||
import gym
|
||||
import gym_avoiding
|
||||
import imageio
|
||||
|
||||
# from gym_avoiding_env.gym_avoiding.envs.avoiding import ObstacleAvoidanceEnv
|
||||
# from envs.gym_avoiding_env.gym_avoiding.envs.avoiding import ObstacleAvoidanceEnv
|
||||
from gym.envs import make as make_
|
||||
import numpy as np
|
||||
|
||||
# env = ObstacleAvoidanceEnv(render=False)
|
||||
env = make_("avoiding-v0", render=True)
|
||||
# env.start() # no need to start() any more, already run in init() now
|
||||
env.reset()
|
||||
print(env.action_space)
|
||||
|
||||
# video_writer = imageio.get_writer("test_d3il.mp4", fps=30)
|
||||
while 1:
|
||||
obs, reward, done, info = env.step(np.array([0.02, 0.1]))
|
||||
print("Reward:", reward)
|
||||
# video_img = env.render(
|
||||
# mode="rgb_array",
|
||||
# # height=640,
|
||||
# # width=480,
|
||||
# # camera_name=self.render_camera_name,
|
||||
# )
|
||||
# video_writer.append_data(video_img)
|
||||
if input("Press space to stop, or any other key to continue") == " ":
|
||||
break
|
||||
# video_writer.close()
|
||||
@@ -0,0 +1,36 @@
|
||||
"""
|
||||
Test Robomimic rendering, no GUI
|
||||
|
||||
"""
|
||||
|
||||
import os
|
||||
import time
|
||||
from gym import spaces
|
||||
import robosuite as suite
|
||||
|
||||
os.environ["MUJOCO_GL"] = "egl"
|
||||
if __name__ == "__main__":
|
||||
env = suite.make(
|
||||
env_name="TwoArmTransport",
|
||||
robots=["Panda", "Panda"],
|
||||
has_renderer=False,
|
||||
has_offscreen_renderer=True,
|
||||
use_camera_obs=True,
|
||||
camera_heights=96,
|
||||
camera_widths=96,
|
||||
camera_names="shouldercamera0",
|
||||
render_gpu_device_id=0,
|
||||
horizon=20,
|
||||
)
|
||||
obs, done = env.reset(), False
|
||||
print("Finished resetting!")
|
||||
low, high = env.action_spec
|
||||
action_space = spaces.Box(low=low, high=high)
|
||||
steps, time_stamp = 0, time.time()
|
||||
while True:
|
||||
while not done:
|
||||
obs, reward, done, info = env.step(action_space.sample())
|
||||
steps += 1
|
||||
obs, done = env.reset(), False
|
||||
print(f"FPS: {steps / (time.time() - time_stamp)}")
|
||||
steps, time_stamp = 0, time.time()
|
||||
@@ -0,0 +1,91 @@
|
||||
"""
|
||||
Launcher for all experiments. Download pre-training data, normalization statistics, and pre-trained checkpoints if needed.
|
||||
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
import pretty_errors
|
||||
import logging
|
||||
|
||||
import math
|
||||
import hydra
|
||||
from omegaconf import OmegaConf
|
||||
import gdown
|
||||
from download_url import (
|
||||
get_dataset_download_url,
|
||||
get_normalization_download_url,
|
||||
get_checkpoint_download_url,
|
||||
)
|
||||
|
||||
# allows arbitrary python code execution in configs using the ${eval:''} resolver
|
||||
OmegaConf.register_new_resolver("eval", eval, replace=True)
|
||||
OmegaConf.register_new_resolver("round_up", math.ceil)
|
||||
OmegaConf.register_new_resolver("round_down", math.floor)
|
||||
|
||||
# suppress d4rl import error
|
||||
os.environ["D4RL_SUPPRESS_IMPORT_ERROR"] = "1"
|
||||
|
||||
# add logger
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
# use line-buffering for both stdout and stderr
|
||||
sys.stdout = open(sys.stdout.fileno(), mode="w", buffering=1)
|
||||
sys.stderr = open(sys.stderr.fileno(), mode="w", buffering=1)
|
||||
|
||||
|
||||
@hydra.main(
|
||||
version_base=None,
|
||||
config_path=os.path.join(
|
||||
os.getcwd(), "cfg"
|
||||
), # possibly overwritten by --config-path
|
||||
)
|
||||
def main(cfg: OmegaConf):
|
||||
# resolve immediately so all the ${now:} resolvers will use the same time.
|
||||
OmegaConf.resolve(cfg)
|
||||
|
||||
# For pre-training: download dataset if needed
|
||||
if "train_dataset_path" in cfg and not os.path.exists(cfg.train_dataset_path):
|
||||
download_url = get_dataset_download_url(cfg)
|
||||
download_target = os.path.dirname(cfg.train_dataset_path)
|
||||
log.info(f"Downloading dataset from {download_url} to {download_target}")
|
||||
gdown.download_folder(url=download_url, output=download_target)
|
||||
|
||||
# For for-tuning: download normalization if needed
|
||||
if "normalization_path" in cfg and not os.path.exists(cfg.normalization_path):
|
||||
download_url = get_normalization_download_url(cfg)
|
||||
download_target = cfg.normalization_path
|
||||
dir_name = os.path.dirname(download_target)
|
||||
if not os.path.exists(dir_name):
|
||||
os.makedirs(dir_name)
|
||||
log.info(
|
||||
f"Downloading normalization statistics from {download_url} to {download_target}"
|
||||
)
|
||||
gdown.download(url=download_url, output=download_target, fuzzy=True)
|
||||
|
||||
# For for-tuning: download checkpoint if needed
|
||||
if "base_policy_path" in cfg and not os.path.exists(cfg.base_policy_path):
|
||||
download_url = get_checkpoint_download_url(cfg)
|
||||
if download_url is None:
|
||||
raise ValueError(
|
||||
f"Unknown checkpoint path. Did you specify the correct path to the policy you trained?"
|
||||
)
|
||||
download_target = cfg.base_policy_path
|
||||
dir_name = os.path.dirname(download_target)
|
||||
if not os.path.exists(dir_name):
|
||||
os.makedirs(dir_name)
|
||||
log.info(f"Downloading checkpoint from {download_url} to {download_target}")
|
||||
gdown.download(url=download_url, output=download_target, fuzzy=True)
|
||||
|
||||
# Deal with isaacgym needs to be imported before torch
|
||||
if "env" in cfg and "env_type" in cfg.env and cfg.env.env_type == "furniture":
|
||||
import furniture_bench
|
||||
|
||||
# run agent
|
||||
cls = hydra.utils.get_class(cfg._target_)
|
||||
agent = cls(cfg)
|
||||
agent.run()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user