release
This commit is contained in:
@@ -0,0 +1,18 @@
|
||||
import os
|
||||
|
||||
FRAMEWORK_DIR = os.path.dirname(__file__)
|
||||
|
||||
|
||||
def sim_framework_path(*args) -> str:
|
||||
"""
|
||||
Abstraction from os.path.join()
|
||||
Builds absolute paths from relative path strings with SIM_FRAMEWORK/ as root.
|
||||
If args already contains an absolute path, it is used as root for the subsequent joins
|
||||
Args:
|
||||
*args:
|
||||
|
||||
Returns:
|
||||
absolute path
|
||||
|
||||
"""
|
||||
return os.path.abspath(os.path.join(FRAMEWORK_DIR, *args))
|
||||
@@ -0,0 +1,342 @@
|
||||
import random
|
||||
from typing import Optional, Callable, Any
|
||||
import logging
|
||||
|
||||
import os
|
||||
import glob
|
||||
|
||||
try:
|
||||
import cv2 # not included in pyproject.toml
|
||||
except:
|
||||
print("Installing cv2")
|
||||
os.system("pip install opencv-python")
|
||||
import torch
|
||||
import pickle
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
|
||||
from agent.dataset.d3il_dataset.base_dataset import TrajectoryDataset
|
||||
from agent.dataset.d3il_dataset import sim_framework_path
|
||||
|
||||
|
||||
class Aligning_Dataset(TrajectoryDataset):
|
||||
def __init__(
|
||||
self,
|
||||
data_directory: os.PathLike,
|
||||
device="cpu",
|
||||
obs_dim: int = 20,
|
||||
action_dim: int = 2,
|
||||
max_len_data: int = 256,
|
||||
window_size: int = 1,
|
||||
):
|
||||
|
||||
super().__init__(
|
||||
data_directory=data_directory,
|
||||
device=device,
|
||||
obs_dim=obs_dim,
|
||||
action_dim=action_dim,
|
||||
max_len_data=max_len_data,
|
||||
window_size=window_size,
|
||||
)
|
||||
|
||||
logging.info("Loading Robot Push Dataset")
|
||||
|
||||
inputs = []
|
||||
actions = []
|
||||
masks = []
|
||||
|
||||
# data_dir = sim_framework_path(data_directory)
|
||||
# state_files = glob.glob(data_dir + "/env*")
|
||||
|
||||
rp_data_dir = sim_framework_path("data/aligning/all_data/state")
|
||||
state_files = np.load(sim_framework_path(data_directory), allow_pickle=True)
|
||||
|
||||
for file in state_files:
|
||||
|
||||
with open(os.path.join(rp_data_dir, file), "rb") as f:
|
||||
env_state = pickle.load(f)
|
||||
|
||||
# lengths.append(len(env_state['robot']['des_c_pos']))
|
||||
|
||||
zero_obs = np.zeros((1, self.max_len_data, self.obs_dim), dtype=np.float32)
|
||||
zero_action = np.zeros(
|
||||
(1, self.max_len_data, self.action_dim), dtype=np.float32
|
||||
)
|
||||
zero_mask = np.zeros((1, self.max_len_data), dtype=np.float32)
|
||||
|
||||
# robot and box positions
|
||||
robot_des_pos = env_state["robot"]["des_c_pos"]
|
||||
|
||||
robot_c_pos = env_state["robot"]["c_pos"]
|
||||
|
||||
push_box_pos = env_state["push-box"]["pos"]
|
||||
push_box_quat = env_state["push-box"]["quat"]
|
||||
|
||||
target_box_pos = env_state["target-box"]["pos"]
|
||||
target_box_quat = env_state["target-box"]["quat"]
|
||||
|
||||
# target_box_pos = np.zeros(push_box_pos.shape)
|
||||
# target_box_quat = np.zeros(push_box_quat.shape)
|
||||
# target_box_pos[:] = push_box_pos[-1:]
|
||||
# target_box_quat[:] = push_box_quat[-1:]
|
||||
|
||||
input_state = np.concatenate(
|
||||
(
|
||||
robot_des_pos,
|
||||
robot_c_pos,
|
||||
push_box_pos,
|
||||
push_box_quat,
|
||||
target_box_pos,
|
||||
target_box_quat,
|
||||
),
|
||||
axis=-1,
|
||||
)
|
||||
|
||||
vel_state = robot_des_pos[1:] - robot_des_pos[:-1]
|
||||
|
||||
valid_len = len(input_state) - 1
|
||||
|
||||
zero_obs[0, :valid_len, :] = input_state[:-1]
|
||||
zero_action[0, :valid_len, :] = vel_state
|
||||
zero_mask[0, :valid_len] = 1
|
||||
|
||||
inputs.append(zero_obs)
|
||||
actions.append(zero_action)
|
||||
masks.append(zero_mask)
|
||||
|
||||
# shape: B, T, n
|
||||
self.observations = torch.from_numpy(np.concatenate(inputs)).to(device).float()
|
||||
self.actions = torch.from_numpy(np.concatenate(actions)).to(device).float()
|
||||
self.masks = torch.from_numpy(np.concatenate(masks)).to(device).float()
|
||||
|
||||
self.num_data = len(self.observations)
|
||||
|
||||
self.slices = self.get_slices()
|
||||
|
||||
def get_slices(self):
|
||||
slices = []
|
||||
|
||||
min_seq_length = np.inf
|
||||
for i in range(self.num_data):
|
||||
T = self.get_seq_length(i)
|
||||
min_seq_length = min(T, min_seq_length)
|
||||
|
||||
if T - self.window_size < 0:
|
||||
print(
|
||||
f"Ignored short sequence #{i}: len={T}, window={self.window_size}"
|
||||
)
|
||||
else:
|
||||
slices += [
|
||||
(i, start, start + self.window_size)
|
||||
for start in range(T - self.window_size + 1)
|
||||
] # slice indices follow convention [start, end)
|
||||
|
||||
return slices
|
||||
|
||||
def get_seq_length(self, idx):
|
||||
return int(self.masks[idx].sum().item())
|
||||
|
||||
def get_all_actions(self):
|
||||
result = []
|
||||
# mask out invalid actions
|
||||
for i in range(len(self.masks)):
|
||||
T = int(self.masks[i].sum().item())
|
||||
result.append(self.actions[i, :T, :])
|
||||
return torch.cat(result, dim=0)
|
||||
|
||||
def get_all_observations(self):
|
||||
result = []
|
||||
# mask out invalid observations
|
||||
for i in range(len(self.masks)):
|
||||
T = int(self.masks[i].sum().item())
|
||||
result.append(self.observations[i, :T, :])
|
||||
return torch.cat(result, dim=0)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.slices)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
|
||||
i, start, end = self.slices[idx]
|
||||
|
||||
obs = self.observations[i, start:end]
|
||||
act = self.actions[i, start:end]
|
||||
mask = self.masks[i, start:end]
|
||||
|
||||
return obs, act, mask
|
||||
|
||||
|
||||
class Aligning_Img_Dataset(TrajectoryDataset):
|
||||
def __init__(
|
||||
self,
|
||||
data_directory: os.PathLike,
|
||||
device="cpu",
|
||||
obs_dim: int = 20,
|
||||
action_dim: int = 2,
|
||||
max_len_data: int = 256,
|
||||
window_size: int = 1,
|
||||
):
|
||||
|
||||
super().__init__(
|
||||
data_directory=data_directory,
|
||||
device=device,
|
||||
obs_dim=obs_dim,
|
||||
action_dim=action_dim,
|
||||
max_len_data=max_len_data,
|
||||
window_size=window_size,
|
||||
)
|
||||
|
||||
logging.info("Loading Robot Push Dataset")
|
||||
|
||||
inputs = []
|
||||
actions = []
|
||||
masks = []
|
||||
|
||||
data_dir = sim_framework_path("environments/dataset/data/aligning/all_data")
|
||||
|
||||
state_files = np.load(sim_framework_path(data_directory), allow_pickle=True)
|
||||
|
||||
bp_cam_imgs = []
|
||||
inhand_cam_imgs = []
|
||||
|
||||
for file in tqdm(state_files[:3]):
|
||||
|
||||
with open(os.path.join(data_dir, "state", file), "rb") as f:
|
||||
env_state = pickle.load(f)
|
||||
|
||||
# lengths.append(len(env_state['robot']['des_c_pos']))
|
||||
|
||||
zero_obs = np.zeros((1, self.max_len_data, self.obs_dim), dtype=np.float32)
|
||||
zero_action = np.zeros(
|
||||
(1, self.max_len_data, self.action_dim), dtype=np.float32
|
||||
)
|
||||
zero_mask = np.zeros((1, self.max_len_data), dtype=np.float32)
|
||||
|
||||
# robot and box positions
|
||||
robot_des_pos = env_state["robot"]["des_c_pos"]
|
||||
robot_c_pos = env_state["robot"]["c_pos"]
|
||||
|
||||
file_name = os.path.basename(file).split(".")[0]
|
||||
|
||||
###############################################################
|
||||
bp_images = []
|
||||
bp_imgs = glob.glob(data_dir + "/images/bp-cam/" + file_name + "/*")
|
||||
bp_imgs.sort(key=lambda x: int(os.path.basename(x).split(".")[0]))
|
||||
|
||||
for img in bp_imgs:
|
||||
image = cv2.imread(img).astype(np.float32)
|
||||
image = image.transpose((2, 0, 1)) / 255.0
|
||||
|
||||
image = torch.from_numpy(image).to(self.device).float().unsqueeze(0)
|
||||
|
||||
bp_images.append(image)
|
||||
|
||||
bp_images = torch.concatenate(bp_images, dim=0)
|
||||
################################################################
|
||||
inhand_imgs = glob.glob(data_dir + "/images/inhand-cam/" + file_name + "/*")
|
||||
inhand_imgs.sort(key=lambda x: int(os.path.basename(x).split(".")[0]))
|
||||
inhand_images = []
|
||||
for img in inhand_imgs:
|
||||
image = cv2.imread(img).astype(np.float32)
|
||||
image = image.transpose((2, 0, 1)) / 255.0
|
||||
|
||||
image = torch.from_numpy(image).to(self.device).float().unsqueeze(0)
|
||||
|
||||
inhand_images.append(image)
|
||||
inhand_images = torch.concatenate(inhand_images, dim=0)
|
||||
##################################################################
|
||||
|
||||
# push_box_pos = env_state['push-box']['pos']
|
||||
# push_box_quat = env_state['push-box']['quat']
|
||||
#
|
||||
# target_box_pos = env_state['target-box']['pos']
|
||||
# target_box_quat = env_state['target-box']['quat']
|
||||
|
||||
# target_box_pos = np.zeros(push_box_pos.shape)
|
||||
# target_box_quat = np.zeros(push_box_quat.shape)
|
||||
# target_box_pos[:] = push_box_pos[-1:]
|
||||
# target_box_quat[:] = push_box_quat[-1:]
|
||||
|
||||
# input_state = np.concatenate((robot_des_pos), axis=-1)
|
||||
|
||||
vel_state = robot_des_pos[1:] - robot_des_pos[:-1]
|
||||
|
||||
valid_len = len(vel_state)
|
||||
|
||||
zero_obs[0, :valid_len, :] = robot_des_pos[:-1]
|
||||
zero_action[0, :valid_len, :] = vel_state
|
||||
zero_mask[0, :valid_len] = 1
|
||||
|
||||
bp_cam_imgs.append(bp_images)
|
||||
inhand_cam_imgs.append(inhand_images)
|
||||
|
||||
inputs.append(zero_obs)
|
||||
actions.append(zero_action)
|
||||
masks.append(zero_mask)
|
||||
|
||||
self.bp_cam_imgs = bp_cam_imgs
|
||||
self.inhand_cam_imgs = inhand_cam_imgs
|
||||
|
||||
# shape: B, T, n
|
||||
self.observations = torch.from_numpy(np.concatenate(inputs)).to(device).float()
|
||||
self.actions = torch.from_numpy(np.concatenate(actions)).to(device).float()
|
||||
self.masks = torch.from_numpy(np.concatenate(masks)).to(device).float()
|
||||
|
||||
self.num_data = len(self.observations)
|
||||
|
||||
self.slices = self.get_slices()
|
||||
|
||||
def get_slices(self):
|
||||
slices = []
|
||||
|
||||
min_seq_length = np.inf
|
||||
for i in range(self.num_data):
|
||||
T = self.get_seq_length(i)
|
||||
min_seq_length = min(T, min_seq_length)
|
||||
|
||||
if T - self.window_size < 0:
|
||||
print(
|
||||
f"Ignored short sequence #{i}: len={T}, window={self.window_size}"
|
||||
)
|
||||
else:
|
||||
slices += [
|
||||
(i, start, start + self.window_size)
|
||||
for start in range(T - self.window_size + 1)
|
||||
] # slice indices follow convention [start, end)
|
||||
|
||||
return slices
|
||||
|
||||
def get_seq_length(self, idx):
|
||||
return int(self.masks[idx].sum().item())
|
||||
|
||||
def get_all_actions(self):
|
||||
result = []
|
||||
# mask out invalid actions
|
||||
for i in range(len(self.masks)):
|
||||
T = int(self.masks[i].sum().item())
|
||||
result.append(self.actions[i, :T, :])
|
||||
return torch.cat(result, dim=0)
|
||||
|
||||
def get_all_observations(self):
|
||||
result = []
|
||||
# mask out invalid observations
|
||||
for i in range(len(self.masks)):
|
||||
T = int(self.masks[i].sum().item())
|
||||
result.append(self.observations[i, :T, :])
|
||||
return torch.cat(result, dim=0)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.slices)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
|
||||
i, start, end = self.slices[idx]
|
||||
|
||||
obs = self.observations[i, start:end]
|
||||
act = self.actions[i, start:end]
|
||||
mask = self.masks[i, start:end]
|
||||
|
||||
bp_imgs = self.bp_cam_imgs[i][start:end]
|
||||
inhand_imgs = self.inhand_cam_imgs[i][start:end]
|
||||
|
||||
return bp_imgs, inhand_imgs, obs, act, mask
|
||||
@@ -0,0 +1,126 @@
|
||||
import logging
|
||||
|
||||
import os
|
||||
import torch
|
||||
import pickle
|
||||
import numpy as np
|
||||
|
||||
from agent.dataset.d3il_dataset.base_dataset import TrajectoryDataset
|
||||
|
||||
|
||||
class Avoiding_Dataset(TrajectoryDataset):
|
||||
def __init__(
|
||||
self,
|
||||
data_directory: os.PathLike,
|
||||
device="cpu",
|
||||
obs_dim: int = 20,
|
||||
action_dim: int = 2,
|
||||
max_len_data: int = 256,
|
||||
window_size: int = 1,
|
||||
):
|
||||
|
||||
super().__init__(
|
||||
data_directory=data_directory,
|
||||
device=device,
|
||||
obs_dim=obs_dim,
|
||||
action_dim=action_dim,
|
||||
max_len_data=max_len_data,
|
||||
window_size=window_size,
|
||||
)
|
||||
|
||||
logging.info("Loading Sorting Dataset")
|
||||
|
||||
inputs = []
|
||||
actions = []
|
||||
masks = []
|
||||
|
||||
data_dir = data_directory
|
||||
state_files = os.listdir(data_dir)
|
||||
|
||||
for file in state_files:
|
||||
with open(os.path.join(data_dir, file), "rb") as f:
|
||||
env_state = pickle.load(f)
|
||||
|
||||
zero_obs = np.zeros((1, self.max_len_data, self.obs_dim), dtype=np.float32)
|
||||
zero_action = np.zeros(
|
||||
(1, self.max_len_data, self.action_dim), dtype=np.float32
|
||||
)
|
||||
zero_mask = np.zeros((1, self.max_len_data), dtype=np.float32)
|
||||
|
||||
# robot and box posistion
|
||||
robot_des_pos = env_state["robot"]["des_c_pos"][:, :2]
|
||||
robot_c_pos = env_state["robot"]["c_pos"][:, :2]
|
||||
|
||||
input_state = np.concatenate((robot_des_pos, robot_c_pos), axis=-1)
|
||||
|
||||
vel_state = robot_des_pos[1:] - robot_des_pos[:-1]
|
||||
valid_len = len(vel_state)
|
||||
|
||||
zero_obs[0, :valid_len, :] = input_state[:-1]
|
||||
zero_action[0, :valid_len, :] = vel_state
|
||||
zero_mask[0, :valid_len] = 1
|
||||
|
||||
inputs.append(zero_obs)
|
||||
actions.append(zero_action)
|
||||
masks.append(zero_mask)
|
||||
|
||||
# shape: B, T, n
|
||||
self.observations = torch.from_numpy(np.concatenate(inputs)).to(device).float()
|
||||
self.actions = torch.from_numpy(np.concatenate(actions)).to(device).float()
|
||||
self.masks = torch.from_numpy(np.concatenate(masks)).to(device).float()
|
||||
|
||||
self.num_data = len(self.observations)
|
||||
|
||||
self.slices = self.get_slices()
|
||||
|
||||
def get_slices(self):
|
||||
slices = []
|
||||
|
||||
min_seq_length = np.inf
|
||||
for i in range(self.num_data):
|
||||
T = self.get_seq_length(i)
|
||||
min_seq_length = min(T, min_seq_length)
|
||||
|
||||
if T - self.window_size < 0:
|
||||
print(
|
||||
f"Ignored short sequence #{i}: len={T}, window={self.window_size}"
|
||||
)
|
||||
else:
|
||||
slices += [
|
||||
(i, start, start + self.window_size)
|
||||
for start in range(T - self.window_size + 1)
|
||||
] # slice indices follow convention [start, end)
|
||||
|
||||
return slices
|
||||
|
||||
def get_seq_length(self, idx):
|
||||
return int(self.masks[idx].sum().item())
|
||||
|
||||
def get_all_actions(self):
|
||||
result = []
|
||||
# mask out invalid actions
|
||||
for i in range(len(self.masks)):
|
||||
T = int(self.masks[i].sum().item())
|
||||
result.append(self.actions[i, :T, :])
|
||||
return torch.cat(result, dim=0)
|
||||
|
||||
def get_all_observations(self):
|
||||
result = []
|
||||
# mask out invalid observations
|
||||
for i in range(len(self.masks)):
|
||||
T = int(self.masks[i].sum().item())
|
||||
result.append(self.observations[i, :T, :])
|
||||
return torch.cat(result, dim=0)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.slices)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
|
||||
i, start, end = self.slices[idx]
|
||||
|
||||
obs = self.observations[i, start:end]
|
||||
act = self.actions[i, start:end]
|
||||
mask = self.masks[i, start:end]
|
||||
|
||||
return obs, act, mask
|
||||
@@ -0,0 +1,54 @@
|
||||
import abc
|
||||
import os
|
||||
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
|
||||
class TrajectoryDataset(Dataset, abc.ABC):
|
||||
"""
|
||||
A dataset containing trajectories.
|
||||
TrajectoryDataset[i] returns: (observations, actions, mask)
|
||||
observations: Tensor[T, ...], T frames of observations
|
||||
actions: Tensor[T, ...], T frames of actions
|
||||
mask: Tensor[T]: 0: invalid; 1: valid
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
data_directory: os.PathLike,
|
||||
device="cpu",
|
||||
obs_dim: int = 20,
|
||||
action_dim: int = 2,
|
||||
max_len_data: int = 256,
|
||||
window_size: int = 1,
|
||||
):
|
||||
|
||||
self.data_directory = data_directory
|
||||
self.device = device
|
||||
|
||||
self.max_len_data = max_len_data
|
||||
self.action_dim = action_dim
|
||||
self.obs_dim = obs_dim
|
||||
|
||||
self.window_size = window_size
|
||||
|
||||
@abc.abstractmethod
|
||||
def get_seq_length(self, idx):
|
||||
"""
|
||||
Returns the length of the idx-th trajectory.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abc.abstractmethod
|
||||
def get_all_actions(self):
|
||||
"""
|
||||
Returns all actions from all trajectories, concatenated on dim 0 (time).
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abc.abstractmethod
|
||||
def get_all_observations(self):
|
||||
"""
|
||||
Returns all actions from all trajectories, concatenated on dim 0 (time).
|
||||
"""
|
||||
raise NotImplementedError
|
||||
@@ -0,0 +1,350 @@
|
||||
import itertools
|
||||
|
||||
import numpy as np
|
||||
|
||||
"""
|
||||
From OpenAIGym Please see there under mujoco/Robots
|
||||
"""
|
||||
|
||||
# For testing whether a number is close to zero
|
||||
_FLOAT_EPS = np.finfo(np.float64).eps
|
||||
_EPS4 = _FLOAT_EPS * 4.0
|
||||
|
||||
|
||||
def get_quaternion_error(curr_quat, des_quat):
|
||||
"""
|
||||
Calculates the difference between the current quaternion and the desired quaternion.
|
||||
See Siciliano textbook page 140 Eq 3.91
|
||||
|
||||
:param curr_quat: current quaternion
|
||||
:param des_quat: desired quaternion
|
||||
:return: difference between current quaternion and desired quaternion
|
||||
"""
|
||||
quatError = np.zeros((3,))
|
||||
|
||||
quatError[0] = (
|
||||
curr_quat[0] * des_quat[1]
|
||||
- des_quat[0] * curr_quat[1]
|
||||
- curr_quat[3] * des_quat[2]
|
||||
+ curr_quat[2] * des_quat[3]
|
||||
)
|
||||
|
||||
quatError[1] = (
|
||||
curr_quat[0] * des_quat[2]
|
||||
- des_quat[0] * curr_quat[2]
|
||||
+ curr_quat[3] * des_quat[1]
|
||||
- curr_quat[1] * des_quat[3]
|
||||
)
|
||||
|
||||
quatError[2] = (
|
||||
curr_quat[0] * des_quat[3]
|
||||
- des_quat[0] * curr_quat[3]
|
||||
- curr_quat[2] * des_quat[1]
|
||||
+ curr_quat[1] * des_quat[2]
|
||||
)
|
||||
|
||||
return quatError
|
||||
|
||||
|
||||
def euler2mat(euler):
|
||||
"""Convert Euler Angles to Rotation Matrix. See rotation.py for notes"""
|
||||
euler = np.asarray(euler, dtype=np.float64)
|
||||
assert euler.shape[-1] == 3, "Invalid shaped euler {}".format(euler)
|
||||
|
||||
ai, aj, ak = -euler[..., 2], -euler[..., 1], -euler[..., 0]
|
||||
si, sj, sk = np.sin(ai), np.sin(aj), np.sin(ak)
|
||||
ci, cj, ck = np.cos(ai), np.cos(aj), np.cos(ak)
|
||||
cc, cs = ci * ck, ci * sk
|
||||
sc, ss = si * ck, si * sk
|
||||
|
||||
mat = np.empty(euler.shape[:-1] + (3, 3), dtype=np.float64)
|
||||
mat[..., 2, 2] = cj * ck
|
||||
mat[..., 2, 1] = sj * sc - cs
|
||||
mat[..., 2, 0] = sj * cc + ss
|
||||
mat[..., 1, 2] = cj * sk
|
||||
mat[..., 1, 1] = sj * ss + cc
|
||||
mat[..., 1, 0] = sj * cs - sc
|
||||
mat[..., 0, 2] = -sj
|
||||
mat[..., 0, 1] = cj * si
|
||||
mat[..., 0, 0] = cj * ci
|
||||
return mat
|
||||
|
||||
|
||||
def euler2quat(euler):
|
||||
"""Convert Euler Angles to Quaternions. See rotation.py for notes"""
|
||||
euler = np.asarray(euler, dtype=np.float64)
|
||||
assert euler.shape[-1] == 3, "Invalid shape euler {}".format(euler)
|
||||
|
||||
ai, aj, ak = euler[..., 2] / 2, -euler[..., 1] / 2, euler[..., 0] / 2
|
||||
si, sj, sk = np.sin(ai), np.sin(aj), np.sin(ak)
|
||||
ci, cj, ck = np.cos(ai), np.cos(aj), np.cos(ak)
|
||||
cc, cs = ci * ck, ci * sk
|
||||
sc, ss = si * ck, si * sk
|
||||
|
||||
quat = np.empty(euler.shape[:-1] + (4,), dtype=np.float64)
|
||||
quat[..., 0] = cj * cc + sj * ss
|
||||
quat[..., 3] = cj * sc - sj * cs
|
||||
quat[..., 2] = -(cj * ss + sj * cc)
|
||||
quat[..., 1] = cj * cs - sj * sc
|
||||
return quat
|
||||
|
||||
|
||||
def mat2euler(mat):
|
||||
"""Convert Rotation Matrix to Euler Angles. See rotation.py for notes"""
|
||||
mat = np.asarray(mat, dtype=np.float64)
|
||||
assert mat.shape[-2:] == (3, 3), "Invalid shape matrix {}".format(mat)
|
||||
|
||||
cy = np.sqrt(mat[..., 2, 2] * mat[..., 2, 2] + mat[..., 1, 2] * mat[..., 1, 2])
|
||||
condition = cy > _EPS4
|
||||
euler = np.empty(mat.shape[:-1], dtype=np.float64)
|
||||
euler[..., 2] = np.where(
|
||||
condition,
|
||||
-np.arctan2(mat[..., 0, 1], mat[..., 0, 0]),
|
||||
-np.arctan2(-mat[..., 1, 0], mat[..., 1, 1]),
|
||||
)
|
||||
euler[..., 1] = np.where(
|
||||
condition, -np.arctan2(-mat[..., 0, 2], cy), -np.arctan2(-mat[..., 0, 2], cy)
|
||||
)
|
||||
euler[..., 0] = np.where(
|
||||
condition, -np.arctan2(mat[..., 1, 2], mat[..., 2, 2]), 0.0
|
||||
)
|
||||
return euler
|
||||
|
||||
|
||||
def mat2quat(mat):
|
||||
"""Convert Rotation Matrix to Quaternion. See rotation.py for notes"""
|
||||
mat = np.asarray(mat, dtype=np.float64)
|
||||
assert mat.shape[-2:] == (3, 3), "Invalid shape matrix {}".format(mat)
|
||||
|
||||
Qxx, Qyx, Qzx = mat[..., 0, 0], mat[..., 0, 1], mat[..., 0, 2]
|
||||
Qxy, Qyy, Qzy = mat[..., 1, 0], mat[..., 1, 1], mat[..., 1, 2]
|
||||
Qxz, Qyz, Qzz = mat[..., 2, 0], mat[..., 2, 1], mat[..., 2, 2]
|
||||
# Fill only lower half of symmetric matrix
|
||||
K = np.zeros(mat.shape[:-2] + (4, 4), dtype=np.float64)
|
||||
K[..., 0, 0] = Qxx - Qyy - Qzz
|
||||
K[..., 1, 0] = Qyx + Qxy
|
||||
K[..., 1, 1] = Qyy - Qxx - Qzz
|
||||
K[..., 2, 0] = Qzx + Qxz
|
||||
K[..., 2, 1] = Qzy + Qyz
|
||||
K[..., 2, 2] = Qzz - Qxx - Qyy
|
||||
K[..., 3, 0] = Qyz - Qzy
|
||||
K[..., 3, 1] = Qzx - Qxz
|
||||
K[..., 3, 2] = Qxy - Qyx
|
||||
K[..., 3, 3] = Qxx + Qyy + Qzz
|
||||
K /= 3.0
|
||||
# TODO: vectorize this -- probably could be made faster
|
||||
q = np.empty(K.shape[:-2] + (4,))
|
||||
it = np.nditer(q[..., 0], flags=["multi_index"])
|
||||
while not it.finished:
|
||||
# Use Hermitian eigenvectors, values for speed
|
||||
vals, vecs = np.linalg.eigh(K[it.multi_index])
|
||||
# Select largest eigenvector, reorder to w,x,y,z quaternion
|
||||
q[it.multi_index] = vecs[[3, 0, 1, 2], np.argmax(vals)]
|
||||
# Prefer quaternion with positive w
|
||||
# (q * -1 corresponds to same rotation as q)
|
||||
if q[it.multi_index][0] < 0:
|
||||
q[it.multi_index] *= -1
|
||||
it.iternext()
|
||||
return q
|
||||
|
||||
|
||||
def quat2euler(quat):
|
||||
"""Convert Quaternion to Euler Angles. See rotation.py for notes"""
|
||||
return mat2euler(quat2mat(quat))
|
||||
|
||||
|
||||
def subtract_euler(e1, e2):
|
||||
assert e1.shape == e2.shape
|
||||
assert e1.shape[-1] == 3
|
||||
q1 = euler2quat(e1)
|
||||
q2 = euler2quat(e2)
|
||||
q_diff = quat_mul(q1, quat_conjugate(q2))
|
||||
return quat2euler(q_diff)
|
||||
|
||||
|
||||
def quat2mat(quat):
|
||||
"""Convert Quaternion to Euler Angles. See rotation.py for notes"""
|
||||
quat = np.asarray(quat, dtype=np.float64)
|
||||
assert quat.shape[-1] == 4, "Invalid shape quat {}".format(quat)
|
||||
|
||||
w, x, y, z = quat[..., 0], quat[..., 1], quat[..., 2], quat[..., 3]
|
||||
Nq = np.sum(quat * quat, axis=-1)
|
||||
s = 2.0 / Nq
|
||||
X, Y, Z = x * s, y * s, z * s
|
||||
wX, wY, wZ = w * X, w * Y, w * Z
|
||||
xX, xY, xZ = x * X, x * Y, x * Z
|
||||
yY, yZ, zZ = y * Y, y * Z, z * Z
|
||||
|
||||
mat = np.empty(quat.shape[:-1] + (3, 3), dtype=np.float64)
|
||||
mat[..., 0, 0] = 1.0 - (yY + zZ)
|
||||
mat[..., 0, 1] = xY - wZ
|
||||
mat[..., 0, 2] = xZ + wY
|
||||
mat[..., 1, 0] = xY + wZ
|
||||
mat[..., 1, 1] = 1.0 - (xX + zZ)
|
||||
mat[..., 1, 2] = yZ - wX
|
||||
mat[..., 2, 0] = xZ - wY
|
||||
mat[..., 2, 1] = yZ + wX
|
||||
mat[..., 2, 2] = 1.0 - (xX + yY)
|
||||
return np.where((Nq > _FLOAT_EPS)[..., np.newaxis, np.newaxis], mat, np.eye(3))
|
||||
|
||||
|
||||
def quat_conjugate(q):
|
||||
inv_q = -q
|
||||
inv_q[..., 0] *= -1
|
||||
return inv_q
|
||||
|
||||
|
||||
def quat_mul(q0, q1):
|
||||
assert q0.shape == q1.shape
|
||||
assert q0.shape[-1] == 4
|
||||
assert q1.shape[-1] == 4
|
||||
|
||||
w0 = q0[..., 0]
|
||||
x0 = q0[..., 1]
|
||||
y0 = q0[..., 2]
|
||||
z0 = q0[..., 3]
|
||||
|
||||
w1 = q1[..., 0]
|
||||
x1 = q1[..., 1]
|
||||
y1 = q1[..., 2]
|
||||
z1 = q1[..., 3]
|
||||
|
||||
w = w0 * w1 - x0 * x1 - y0 * y1 - z0 * z1
|
||||
x = w0 * x1 + x0 * w1 + y0 * z1 - z0 * y1
|
||||
y = w0 * y1 + y0 * w1 + z0 * x1 - x0 * z1
|
||||
z = w0 * z1 + z0 * w1 + x0 * y1 - y0 * x1
|
||||
q = np.array([w, x, y, z])
|
||||
if q.ndim == 2:
|
||||
q = q.swapaxes(0, 1)
|
||||
assert q.shape == q0.shape
|
||||
return q
|
||||
|
||||
|
||||
def quat_rot_vec(q, v0):
|
||||
q_v0 = np.array([0, v0[0], v0[1], v0[2]])
|
||||
q_v = quat_mul(q, quat_mul(q_v0, quat_conjugate(q)))
|
||||
v = q_v[1:]
|
||||
return v
|
||||
|
||||
|
||||
def quat_identity():
|
||||
return np.array([1, 0, 0, 0])
|
||||
|
||||
|
||||
def quat2axisangle(quat):
|
||||
theta = 0
|
||||
axis = np.array([0, 0, 1])
|
||||
sin_theta = np.linalg.norm(quat[1:])
|
||||
|
||||
if sin_theta > 0.0001:
|
||||
theta = 2 * np.arcsin(sin_theta)
|
||||
theta *= 1 if quat[0] >= 0 else -1
|
||||
axis = quat[1:] / sin_theta
|
||||
|
||||
return axis, theta
|
||||
|
||||
|
||||
def euler2point_euler(euler):
|
||||
_euler = euler.copy()
|
||||
if len(_euler.shape) < 2:
|
||||
_euler = np.expand_dims(_euler, 0)
|
||||
assert _euler.shape[1] == 3
|
||||
_euler_sin = np.sin(_euler)
|
||||
_euler_cos = np.cos(_euler)
|
||||
return np.concatenate([_euler_sin, _euler_cos], axis=-1)
|
||||
|
||||
|
||||
def point_euler2euler(euler):
|
||||
_euler = euler.copy()
|
||||
if len(_euler.shape) < 2:
|
||||
_euler = np.expand_dims(_euler, 0)
|
||||
assert _euler.shape[1] == 6
|
||||
angle = np.arctan(_euler[..., :3] / _euler[..., 3:])
|
||||
angle[_euler[..., 3:] < 0] += np.pi
|
||||
return angle
|
||||
|
||||
|
||||
def quat2point_quat(quat):
|
||||
# Should be in qw, qx, qy, qz
|
||||
_quat = quat.copy()
|
||||
if len(_quat.shape) < 2:
|
||||
_quat = np.expand_dims(_quat, 0)
|
||||
assert _quat.shape[1] == 4
|
||||
angle = np.arccos(_quat[:, [0]]) * 2
|
||||
xyz = _quat[:, 1:]
|
||||
xyz[np.squeeze(np.abs(np.sin(angle / 2))) >= 1e-5] = (xyz / np.sin(angle / 2))[
|
||||
np.squeeze(np.abs(np.sin(angle / 2))) >= 1e-5
|
||||
]
|
||||
return np.concatenate([np.sin(angle), np.cos(angle), xyz], axis=-1)
|
||||
|
||||
|
||||
def point_quat2quat(quat):
|
||||
_quat = quat.copy()
|
||||
if len(_quat.shape) < 2:
|
||||
_quat = np.expand_dims(_quat, 0)
|
||||
assert _quat.shape[1] == 5
|
||||
angle = np.arctan(_quat[:, [0]] / _quat[:, [1]])
|
||||
qw = np.cos(angle / 2)
|
||||
|
||||
qxyz = _quat[:, 2:]
|
||||
qxyz[np.squeeze(np.abs(np.sin(angle / 2))) >= 1e-5] = (qxyz * np.sin(angle / 2))[
|
||||
np.squeeze(np.abs(np.sin(angle / 2))) >= 1e-5
|
||||
]
|
||||
return np.concatenate([qw, qxyz], axis=-1)
|
||||
|
||||
|
||||
def normalize_angles(angles):
|
||||
"""Puts angles in [-pi, pi] range."""
|
||||
angles = angles.copy()
|
||||
if angles.size > 0:
|
||||
angles = (angles + np.pi) % (2 * np.pi) - np.pi
|
||||
assert -np.pi - 1e-6 <= angles.min() and angles.max() <= np.pi + 1e-6
|
||||
return angles
|
||||
|
||||
|
||||
def round_to_straight_angles(angles):
|
||||
"""Returns closest angle modulo 90 degrees"""
|
||||
angles = np.round(angles / (np.pi / 2)) * (np.pi / 2)
|
||||
return normalize_angles(angles)
|
||||
|
||||
|
||||
def get_parallel_rotations():
|
||||
mult90 = [0, np.pi / 2, -np.pi / 2, np.pi]
|
||||
parallel_rotations = []
|
||||
for euler in itertools.product(mult90, repeat=3):
|
||||
canonical = mat2euler(euler2mat(euler))
|
||||
canonical = np.round(canonical / (np.pi / 2))
|
||||
if canonical[0] == -2:
|
||||
canonical[0] = 2
|
||||
if canonical[2] == -2:
|
||||
canonical[2] = 2
|
||||
canonical *= np.pi / 2
|
||||
if all([(canonical != rot).any() for rot in parallel_rotations]):
|
||||
parallel_rotations += [canonical]
|
||||
assert len(parallel_rotations) == 24
|
||||
return parallel_rotations
|
||||
|
||||
|
||||
def posRotMat2TFMat(pos, rot_mat):
|
||||
"""Converts a position and a 3x3 rotation matrix to a 4x4 transformation matrix"""
|
||||
t_mat = np.eye(4)
|
||||
t_mat[:3, :3] = rot_mat
|
||||
t_mat[:3, 3] = np.array(pos)
|
||||
return t_mat
|
||||
|
||||
|
||||
def mat2posQuat(mat):
|
||||
"""Converts a 4x4 rotation matrix to a position and a quaternion"""
|
||||
pos = mat[:3, 3]
|
||||
quat = mat2quat(mat[:3, :3])
|
||||
return pos, quat
|
||||
|
||||
|
||||
def wxyz_to_xyzw(quat):
|
||||
"""Converts WXYZ Quaternions to XYZW Quaternions"""
|
||||
return np.roll(quat, -1)
|
||||
|
||||
|
||||
def xyzw_to_wxyz(quat):
|
||||
"""Converts XYZW Quaternions to WXYZ Quaternions"""
|
||||
return np.roll(quat, 1)
|
||||
@@ -0,0 +1,161 @@
|
||||
import logging
|
||||
|
||||
import os
|
||||
import torch
|
||||
import pickle
|
||||
import numpy as np
|
||||
|
||||
from agent.dataset.d3il_dataset.base_dataset import TrajectoryDataset
|
||||
from agent.dataset.d3il_dataset import sim_framework_path
|
||||
from .geo_transform import quat2euler
|
||||
|
||||
|
||||
class Pushing_Dataset(TrajectoryDataset):
|
||||
def __init__(
|
||||
self,
|
||||
data_directory: os.PathLike,
|
||||
device="cpu",
|
||||
obs_dim: int = 20,
|
||||
action_dim: int = 2,
|
||||
max_len_data: int = 256,
|
||||
window_size: int = 1,
|
||||
):
|
||||
|
||||
super().__init__(
|
||||
data_directory=data_directory,
|
||||
device=device,
|
||||
obs_dim=obs_dim,
|
||||
action_dim=action_dim,
|
||||
max_len_data=max_len_data,
|
||||
window_size=window_size,
|
||||
)
|
||||
|
||||
logging.info("Loading Block Push Dataset")
|
||||
|
||||
inputs = []
|
||||
actions = []
|
||||
masks = []
|
||||
|
||||
# for root, dirs, files in os.walk(self.data_directory):
|
||||
#
|
||||
# for mode_dir in dirs:
|
||||
|
||||
# state_files = glob.glob(os.path.join(root, mode_dir) + "/env*")
|
||||
# data_dir = os.path.join(sim_framework_path(data_directory), "local")
|
||||
# data_dir = sim_framework_path(data_directory)
|
||||
# state_files = glob.glob(data_dir + "/env*")
|
||||
|
||||
bp_data_dir = sim_framework_path("data/pushing/all_data")
|
||||
|
||||
state_files = np.load(sim_framework_path(data_directory), allow_pickle=True)
|
||||
|
||||
for file in state_files:
|
||||
with open(os.path.join(bp_data_dir, file), "rb") as f:
|
||||
env_state = pickle.load(f)
|
||||
|
||||
# lengths.append(len(env_state['robot']['des_c_pos']))
|
||||
|
||||
zero_obs = np.zeros((1, self.max_len_data, self.obs_dim), dtype=np.float32)
|
||||
zero_action = np.zeros(
|
||||
(1, self.max_len_data, self.action_dim), dtype=np.float32
|
||||
)
|
||||
zero_mask = np.zeros((1, self.max_len_data), dtype=np.float32)
|
||||
|
||||
# robot and box positions
|
||||
robot_des_pos = env_state["robot"]["des_c_pos"][:, :2]
|
||||
|
||||
robot_c_pos = env_state["robot"]["c_pos"][:, :2]
|
||||
|
||||
red_box_pos = env_state["red-box"]["pos"][:, :2]
|
||||
red_box_quat = np.tan(quat2euler(env_state["red-box"]["quat"])[:, -1:])
|
||||
|
||||
green_box_pos = env_state["green-box"]["pos"][:, :2]
|
||||
green_box_quat = np.tan(quat2euler(env_state["green-box"]["quat"])[:, -1:])
|
||||
|
||||
red_target_pos = env_state["red-target"]["pos"][:, :2]
|
||||
green_target_pos = env_state["green-target"]["pos"][:, :2]
|
||||
|
||||
input_state = np.concatenate(
|
||||
(
|
||||
robot_des_pos,
|
||||
robot_c_pos,
|
||||
red_box_pos,
|
||||
red_box_quat,
|
||||
green_box_pos,
|
||||
green_box_quat,
|
||||
),
|
||||
axis=-1,
|
||||
)
|
||||
|
||||
vel_state = robot_des_pos[1:] - robot_des_pos[:-1]
|
||||
|
||||
valid_len = len(input_state) - 1
|
||||
|
||||
zero_obs[0, :valid_len, :] = input_state[:-1]
|
||||
zero_action[0, :valid_len, :] = vel_state
|
||||
zero_mask[0, :valid_len] = 1
|
||||
|
||||
inputs.append(zero_obs)
|
||||
actions.append(zero_action)
|
||||
masks.append(zero_mask)
|
||||
|
||||
# shape: B, T, n
|
||||
self.observations = torch.from_numpy(np.concatenate(inputs)).to(device).float()
|
||||
self.actions = torch.from_numpy(np.concatenate(actions)).to(device).float()
|
||||
self.masks = torch.from_numpy(np.concatenate(masks)).to(device).float()
|
||||
|
||||
self.num_data = len(self.observations)
|
||||
|
||||
self.slices = self.get_slices()
|
||||
|
||||
def get_slices(self):
|
||||
slices = []
|
||||
|
||||
min_seq_length = np.inf
|
||||
for i in range(self.num_data):
|
||||
T = self.get_seq_length(i)
|
||||
min_seq_length = min(T, min_seq_length)
|
||||
|
||||
if T - self.window_size < 0:
|
||||
print(
|
||||
f"Ignored short sequence #{i}: len={T}, window={self.window_size}"
|
||||
)
|
||||
else:
|
||||
slices += [
|
||||
(i, start, start + self.window_size)
|
||||
for start in range(T - self.window_size + 1)
|
||||
] # slice indices follow convention [start, end)
|
||||
|
||||
return slices
|
||||
|
||||
def get_seq_length(self, idx):
|
||||
return int(self.masks[idx].sum().item())
|
||||
|
||||
def get_all_actions(self):
|
||||
result = []
|
||||
# mask out invalid actions
|
||||
for i in range(len(self.masks)):
|
||||
T = int(self.masks[i].sum().item())
|
||||
result.append(self.actions[i, :T, :])
|
||||
return torch.cat(result, dim=0)
|
||||
|
||||
def get_all_observations(self):
|
||||
result = []
|
||||
# mask out invalid observations
|
||||
for i in range(len(self.masks)):
|
||||
T = int(self.masks[i].sum().item())
|
||||
result.append(self.observations[i, :T, :])
|
||||
return torch.cat(result, dim=0)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.slices)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
|
||||
i, start, end = self.slices[idx]
|
||||
|
||||
obs = self.observations[i, start:end]
|
||||
act = self.actions[i, start:end]
|
||||
mask = self.masks[i, start:end]
|
||||
|
||||
return obs, act, mask
|
||||
@@ -0,0 +1,491 @@
|
||||
import logging
|
||||
|
||||
import os
|
||||
import glob
|
||||
import torch
|
||||
import pickle
|
||||
import numpy as np
|
||||
|
||||
try:
|
||||
import cv2 # not included in pyproject.toml
|
||||
except:
|
||||
print("Installing cv2")
|
||||
os.system("pip install opencv-python")
|
||||
from tqdm import tqdm
|
||||
|
||||
from agent.dataset.d3il_dataset.base_dataset import TrajectoryDataset
|
||||
from agent.dataset.d3il_dataset import sim_framework_path
|
||||
from .geo_transform import quat2euler
|
||||
|
||||
|
||||
class Sorting_Dataset(TrajectoryDataset):
|
||||
def __init__(
|
||||
self,
|
||||
data_directory: os.PathLike,
|
||||
device="cpu",
|
||||
obs_dim: int = 20,
|
||||
action_dim: int = 2,
|
||||
max_len_data: int = 256,
|
||||
window_size: int = 1,
|
||||
num_boxes: int = 2,
|
||||
):
|
||||
|
||||
super().__init__(
|
||||
data_directory=data_directory,
|
||||
device=device,
|
||||
obs_dim=obs_dim,
|
||||
action_dim=action_dim,
|
||||
max_len_data=max_len_data,
|
||||
window_size=window_size,
|
||||
)
|
||||
|
||||
logging.info("Loading Sorting Dataset")
|
||||
|
||||
inputs = []
|
||||
actions = []
|
||||
masks = []
|
||||
|
||||
# for root, dirs, files in os.walk(self.data_directory):
|
||||
#
|
||||
# for mode_dir in dirs:
|
||||
|
||||
# state_files = glob.glob(os.path.join(root, mode_dir) + "/env*")
|
||||
# data_dir = os.path.join(sim_framework_path(data_directory), "local")
|
||||
# data_dir = sim_framework_path(data_directory)
|
||||
# state_files = glob.glob(data_dir + "/env*")
|
||||
|
||||
# random.seed(0)
|
||||
|
||||
# data_dir = sim_framework_path(data_directory)
|
||||
# state_files = os.listdir(data_dir)
|
||||
#
|
||||
# random.shuffle(state_files)
|
||||
#
|
||||
# if data == "train":
|
||||
# env_state_files = state_files[50:]
|
||||
# elif data == "eval":
|
||||
# env_state_files = state_files[:50]
|
||||
# else:
|
||||
# assert False, "wrong data type"
|
||||
|
||||
if num_boxes == 2:
|
||||
data_dir = sim_framework_path("data/sorting/2_boxes/state")
|
||||
elif num_boxes == 4:
|
||||
data_dir = sim_framework_path("data/sorting/4_boxes/state")
|
||||
elif num_boxes == 6:
|
||||
data_dir = sim_framework_path("data/sorting/6_boxes/state")
|
||||
else:
|
||||
assert False, "check num boxes"
|
||||
|
||||
state_files = np.load(sim_framework_path(data_directory), allow_pickle=True)
|
||||
|
||||
for file in state_files:
|
||||
with open(os.path.join(data_dir, file), "rb") as f:
|
||||
env_state = pickle.load(f)
|
||||
|
||||
# lengths.append(len(env_state['robot']['des_c_pos']))
|
||||
|
||||
zero_obs = np.zeros((1, self.max_len_data, self.obs_dim), dtype=np.float32)
|
||||
zero_action = np.zeros(
|
||||
(1, self.max_len_data, self.action_dim), dtype=np.float32
|
||||
)
|
||||
zero_mask = np.zeros((1, self.max_len_data), dtype=np.float32)
|
||||
|
||||
# robot and box posistion
|
||||
robot_des_pos = env_state["robot"]["des_c_pos"][:, :2]
|
||||
robot_c_pos = env_state["robot"]["c_pos"][:, :2]
|
||||
|
||||
if num_boxes == 2:
|
||||
red_box1_pos = env_state["red-box1"]["pos"][:, :2]
|
||||
red_box1_quat = np.tan(
|
||||
quat2euler(env_state["red-box1"]["quat"])[:, -1:]
|
||||
)
|
||||
|
||||
blue_box1_pos = env_state["blue-box1"]["pos"][:, :2]
|
||||
blue_box1_quat = np.tan(
|
||||
quat2euler(env_state["blue-box1"]["quat"])[:, -1:]
|
||||
)
|
||||
|
||||
input_state = np.concatenate(
|
||||
(
|
||||
robot_des_pos,
|
||||
robot_c_pos,
|
||||
red_box1_pos,
|
||||
red_box1_quat,
|
||||
blue_box1_pos,
|
||||
blue_box1_quat,
|
||||
),
|
||||
axis=-1,
|
||||
)
|
||||
|
||||
elif num_boxes == 4:
|
||||
|
||||
red_box1_pos = env_state["red-box1"]["pos"][:, :2]
|
||||
red_box1_quat = np.tan(
|
||||
quat2euler(env_state["red-box1"]["quat"])[:, -1:]
|
||||
)
|
||||
|
||||
red_box2_pos = env_state["red-box2"]["pos"][:, :2]
|
||||
red_box2_quat = np.tan(
|
||||
quat2euler(env_state["red-box2"]["quat"])[:, -1:]
|
||||
)
|
||||
|
||||
blue_box1_pos = env_state["blue-box1"]["pos"][:, :2]
|
||||
blue_box1_quat = np.tan(
|
||||
quat2euler(env_state["blue-box1"]["quat"])[:, -1:]
|
||||
)
|
||||
|
||||
blue_box2_pos = env_state["blue-box2"]["pos"][:, :2]
|
||||
blue_box2_quat = np.tan(
|
||||
quat2euler(env_state["blue-box2"]["quat"])[:, -1:]
|
||||
)
|
||||
|
||||
input_state = np.concatenate(
|
||||
(
|
||||
robot_des_pos,
|
||||
robot_c_pos,
|
||||
red_box1_pos,
|
||||
red_box1_quat,
|
||||
red_box2_pos,
|
||||
red_box2_quat,
|
||||
blue_box1_pos,
|
||||
blue_box1_quat,
|
||||
blue_box2_pos,
|
||||
blue_box2_quat,
|
||||
),
|
||||
axis=-1,
|
||||
)
|
||||
|
||||
elif num_boxes == 6:
|
||||
|
||||
red_box1_pos = env_state["red-box1"]["pos"][:, :2]
|
||||
red_box1_quat = np.tan(
|
||||
quat2euler(env_state["red-box1"]["quat"])[:, -1:]
|
||||
)
|
||||
|
||||
red_box2_pos = env_state["red-box2"]["pos"][:, :2]
|
||||
red_box2_quat = np.tan(
|
||||
quat2euler(env_state["red-box2"]["quat"])[:, -1:]
|
||||
)
|
||||
|
||||
red_box3_pos = env_state["red-box3"]["pos"][:, :2]
|
||||
red_box3_quat = np.tan(
|
||||
quat2euler(env_state["red-box3"]["quat"])[:, -1:]
|
||||
)
|
||||
|
||||
blue_box1_pos = env_state["blue-box1"]["pos"][:, :2]
|
||||
blue_box1_quat = np.tan(
|
||||
quat2euler(env_state["blue-box1"]["quat"])[:, -1:]
|
||||
)
|
||||
|
||||
blue_box2_pos = env_state["blue-box2"]["pos"][:, :2]
|
||||
blue_box2_quat = np.tan(
|
||||
quat2euler(env_state["blue-box2"]["quat"])[:, -1:]
|
||||
)
|
||||
|
||||
blue_box3_pos = env_state["blue-box3"]["pos"][:, :2]
|
||||
blue_box3_quat = np.tan(
|
||||
quat2euler(env_state["blue-box3"]["quat"])[:, -1:]
|
||||
)
|
||||
|
||||
input_state = np.concatenate(
|
||||
(
|
||||
robot_des_pos,
|
||||
robot_c_pos,
|
||||
red_box1_pos,
|
||||
red_box1_quat,
|
||||
red_box2_pos,
|
||||
red_box2_quat,
|
||||
red_box3_pos,
|
||||
red_box3_quat,
|
||||
blue_box1_pos,
|
||||
blue_box1_quat,
|
||||
blue_box2_pos,
|
||||
blue_box2_quat,
|
||||
blue_box3_pos,
|
||||
blue_box3_quat,
|
||||
),
|
||||
axis=-1,
|
||||
)
|
||||
|
||||
else:
|
||||
assert False, "check num boxes"
|
||||
|
||||
vel_state = robot_des_pos[1:] - robot_des_pos[:-1]
|
||||
valid_len = len(vel_state)
|
||||
|
||||
zero_obs[0, :valid_len, :] = input_state[:-1]
|
||||
zero_action[0, :valid_len, :] = vel_state
|
||||
zero_mask[0, :valid_len] = 1
|
||||
|
||||
inputs.append(zero_obs)
|
||||
actions.append(zero_action)
|
||||
masks.append(zero_mask)
|
||||
|
||||
# shape: B, T, n
|
||||
self.observations = torch.from_numpy(np.concatenate(inputs)).to(device).float()
|
||||
self.actions = torch.from_numpy(np.concatenate(actions)).to(device).float()
|
||||
self.masks = torch.from_numpy(np.concatenate(masks)).to(device).float()
|
||||
|
||||
self.num_data = len(self.observations)
|
||||
|
||||
self.slices = self.get_slices()
|
||||
|
||||
def get_slices(self):
|
||||
slices = []
|
||||
|
||||
min_seq_length = np.inf
|
||||
for i in range(self.num_data):
|
||||
T = self.get_seq_length(i)
|
||||
min_seq_length = min(T, min_seq_length)
|
||||
|
||||
if T - self.window_size < 0:
|
||||
print(
|
||||
f"Ignored short sequence #{i}: len={T}, window={self.window_size}"
|
||||
)
|
||||
else:
|
||||
slices += [
|
||||
(i, start, start + self.window_size)
|
||||
for start in range(T - self.window_size + 1)
|
||||
] # slice indices follow convention [start, end)
|
||||
|
||||
return slices
|
||||
|
||||
def get_seq_length(self, idx):
|
||||
return int(self.masks[idx].sum().item())
|
||||
|
||||
def get_all_actions(self):
|
||||
result = []
|
||||
# mask out invalid actions
|
||||
for i in range(len(self.masks)):
|
||||
T = int(self.masks[i].sum().item())
|
||||
result.append(self.actions[i, :T, :])
|
||||
return torch.cat(result, dim=0)
|
||||
|
||||
def get_all_observations(self):
|
||||
result = []
|
||||
# mask out invalid observations
|
||||
for i in range(len(self.masks)):
|
||||
T = int(self.masks[i].sum().item())
|
||||
result.append(self.observations[i, :T, :])
|
||||
return torch.cat(result, dim=0)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.slices)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
|
||||
i, start, end = self.slices[idx]
|
||||
|
||||
obs = self.observations[i, start:end]
|
||||
act = self.actions[i, start:end]
|
||||
mask = self.masks[i, start:end]
|
||||
|
||||
return obs, act, mask
|
||||
|
||||
|
||||
class Sorting_Img_Dataset(TrajectoryDataset):
|
||||
def __init__(
|
||||
self,
|
||||
data_directory: os.PathLike,
|
||||
device="cpu",
|
||||
obs_dim: int = 20,
|
||||
action_dim: int = 2,
|
||||
max_len_data: int = 256,
|
||||
window_size: int = 1,
|
||||
num_boxes: int = 2,
|
||||
):
|
||||
|
||||
super().__init__(
|
||||
data_directory=data_directory,
|
||||
device=device,
|
||||
obs_dim=obs_dim,
|
||||
action_dim=action_dim,
|
||||
max_len_data=max_len_data,
|
||||
window_size=window_size,
|
||||
)
|
||||
|
||||
logging.info("Loading Sorting Dataset")
|
||||
|
||||
inputs = []
|
||||
actions = []
|
||||
masks = []
|
||||
|
||||
# for root, dirs, files in os.walk(self.data_directory):
|
||||
#
|
||||
# for mode_dir in dirs:
|
||||
|
||||
# state_files = glob.glob(os.path.join(root, mode_dir) + "/env*")
|
||||
# data_dir = os.path.join(sim_framework_path(data_directory), "local")
|
||||
# data_dir = sim_framework_path(data_directory)
|
||||
# state_files = glob.glob(data_dir + "/env*")
|
||||
|
||||
# random.seed(0)
|
||||
#
|
||||
# data_dir = sim_framework_path(data_directory)
|
||||
# state_files = glob.glob(data_dir + '/state/*')
|
||||
#
|
||||
# random.shuffle(state_files)
|
||||
#
|
||||
# if data == "train":
|
||||
# env_state_files = state_files[30:]
|
||||
# elif data == "eval":
|
||||
# env_state_files = state_files[:30]
|
||||
# else:
|
||||
# assert False, "wrong data type"
|
||||
|
||||
if num_boxes == 2:
|
||||
data_dir = sim_framework_path("environments/dataset/data/sorting/2_boxes/")
|
||||
elif num_boxes == 4:
|
||||
data_dir = sim_framework_path("environments/dataset/data/sorting/4_boxes/")
|
||||
elif num_boxes == 6:
|
||||
data_dir = sim_framework_path("environments/dataset/data/sorting/6_boxes/")
|
||||
else:
|
||||
assert False, "check num boxes"
|
||||
|
||||
state_files = np.load(sim_framework_path(data_directory), allow_pickle=True)
|
||||
|
||||
bp_cam_imgs = []
|
||||
inhand_cam_imgs = []
|
||||
|
||||
for file in tqdm(state_files[:100]):
|
||||
with open(os.path.join(data_dir, "state", file), "rb") as f:
|
||||
env_state = pickle.load(f)
|
||||
|
||||
# lengths.append(len(env_state['robot']['des_c_pos']))
|
||||
zero_obs = np.zeros((1, self.max_len_data, self.obs_dim), dtype=np.float32)
|
||||
zero_action = np.zeros(
|
||||
(1, self.max_len_data, self.action_dim), dtype=np.float32
|
||||
)
|
||||
zero_mask = np.zeros((1, self.max_len_data), dtype=np.float32)
|
||||
|
||||
# robot and box posistion
|
||||
robot_des_pos = env_state["robot"]["des_c_pos"][:, :2]
|
||||
robot_c_pos = env_state["robot"]["c_pos"][:, :2]
|
||||
|
||||
file_name = os.path.basename(file).split(".")[0]
|
||||
|
||||
###############################################################
|
||||
bp_images = []
|
||||
bp_imgs = glob.glob(data_dir + "/images/bp-cam/" + file_name + "/*")
|
||||
bp_imgs.sort(key=lambda x: int(os.path.basename(x).split(".")[0]))
|
||||
|
||||
for img in bp_imgs:
|
||||
image = cv2.imread(img).astype(np.float32)
|
||||
image = image.transpose((2, 0, 1)) / 255.0
|
||||
|
||||
image = torch.from_numpy(image).to(self.device).float().unsqueeze(0)
|
||||
|
||||
bp_images.append(image)
|
||||
|
||||
bp_images = torch.concatenate(bp_images, dim=0)
|
||||
################################################################
|
||||
inhand_imgs = glob.glob(data_dir + "/images/inhand-cam/" + file_name + "/*")
|
||||
inhand_imgs.sort(key=lambda x: int(os.path.basename(x).split(".")[0]))
|
||||
inhand_images = []
|
||||
for img in inhand_imgs:
|
||||
image = cv2.imread(img).astype(np.float32)
|
||||
image = image.transpose((2, 0, 1)) / 255.0
|
||||
|
||||
image = torch.from_numpy(image).to(self.device).float().unsqueeze(0)
|
||||
|
||||
inhand_images.append(image)
|
||||
inhand_images = torch.concatenate(inhand_images, dim=0)
|
||||
##################################################################
|
||||
# input_state = np.concatenate((robot_des_pos, robot_c_pos), axis=-1)
|
||||
|
||||
vel_state = robot_des_pos[1:] - robot_des_pos[:-1]
|
||||
|
||||
valid_len = len(vel_state)
|
||||
|
||||
zero_obs[0, :valid_len, :] = robot_des_pos[:-1]
|
||||
zero_action[0, :valid_len, :] = vel_state
|
||||
zero_mask[0, :valid_len] = 1
|
||||
|
||||
bp_cam_imgs.append(bp_images)
|
||||
inhand_cam_imgs.append(inhand_images)
|
||||
|
||||
inputs.append(zero_obs)
|
||||
actions.append(zero_action)
|
||||
masks.append(zero_mask)
|
||||
|
||||
self.bp_cam_imgs = bp_cam_imgs
|
||||
self.inhand_cam_imgs = inhand_cam_imgs
|
||||
|
||||
# shape: B, T, n
|
||||
self.observations = torch.from_numpy(np.concatenate(inputs)).to(device).float()
|
||||
self.actions = torch.from_numpy(np.concatenate(actions)).to(device).float()
|
||||
self.masks = torch.from_numpy(np.concatenate(masks)).to(device).float()
|
||||
|
||||
self.num_data = len(self.actions)
|
||||
|
||||
self.slices = self.get_slices()
|
||||
|
||||
def get_slices(self):
|
||||
slices = []
|
||||
|
||||
min_seq_length = np.inf
|
||||
for i in range(self.num_data):
|
||||
T = self.get_seq_length(i)
|
||||
min_seq_length = min(T, min_seq_length)
|
||||
|
||||
if T - self.window_size < 0:
|
||||
print(
|
||||
f"Ignored short sequence #{i}: len={T}, window={self.window_size}"
|
||||
)
|
||||
else:
|
||||
slices += [
|
||||
(i, start, start + self.window_size)
|
||||
for start in range(T - self.window_size + 1)
|
||||
] # slice indices follow convention [start, end)
|
||||
|
||||
return slices
|
||||
|
||||
def get_seq_length(self, idx):
|
||||
return int(self.masks[idx].sum().item())
|
||||
|
||||
def get_all_actions(self):
|
||||
result = []
|
||||
# mask out invalid actions
|
||||
for i in range(len(self.masks)):
|
||||
T = int(self.masks[i].sum().item())
|
||||
result.append(self.actions[i, :T, :])
|
||||
return torch.cat(result, dim=0)
|
||||
|
||||
def get_all_observations(self):
|
||||
result = []
|
||||
# mask out invalid observations
|
||||
for i in range(len(self.masks)):
|
||||
T = int(self.masks[i].sum().item())
|
||||
result.append(self.observations[i, :T, :])
|
||||
return torch.cat(result, dim=0)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.slices)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
|
||||
i, start, end = self.slices[idx]
|
||||
|
||||
obs = self.observations[i, start:end]
|
||||
act = self.actions[i, start:end]
|
||||
mask = self.masks[i, start:end]
|
||||
|
||||
bp_imgs = self.bp_cam_imgs[i][start:end]
|
||||
inhand_imgs = self.inhand_cam_imgs[i][start:end]
|
||||
|
||||
# bp_imgs = np.zeros((self.window_size, 3, 96, 96), dtype=np.float32)
|
||||
# inhand_imgs = np.zeros((self.window_size, 3, 96, 96), dtype=np.float32)
|
||||
#
|
||||
# for num_frame, img_file in enumerate(bp_img_files):
|
||||
# image = cv2.imread(img_file).astype(np.float32)
|
||||
# bp_imgs[num_frame] = image.transpose((2, 0, 1)) / 255.
|
||||
#
|
||||
# for num_frame, img_file in enumerate(inhand_img_files):
|
||||
# image = cv2.imread(img_file).astype(np.float32)
|
||||
# inhand_imgs[num_frame] = image.transpose((2, 0, 1)) / 255.
|
||||
#
|
||||
# bp_imgs = torch.from_numpy(bp_imgs).to(self.device).float()
|
||||
# inhand_imgs = torch.from_numpy(inhand_imgs).to(self.device).float()
|
||||
|
||||
return bp_imgs, inhand_imgs, obs, act, mask
|
||||
@@ -0,0 +1,398 @@
|
||||
import logging
|
||||
|
||||
import os
|
||||
import glob
|
||||
|
||||
try:
|
||||
import cv2 # not included in pyproject.toml
|
||||
except:
|
||||
print("Installing cv2")
|
||||
os.system("pip install opencv-python")
|
||||
import torch
|
||||
import pickle
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
|
||||
from agent.dataset.d3il_dataset.base_dataset import TrajectoryDataset
|
||||
from agent.dataset.d3il_dataset import sim_framework_path
|
||||
|
||||
from .geo_transform import quat2euler
|
||||
|
||||
|
||||
class Stacking_Dataset(TrajectoryDataset):
|
||||
def __init__(
|
||||
self,
|
||||
data_directory: os.PathLike,
|
||||
# data='train',
|
||||
device="cpu",
|
||||
obs_dim: int = 20,
|
||||
action_dim: int = 2,
|
||||
max_len_data: int = 256,
|
||||
window_size: int = 1,
|
||||
):
|
||||
|
||||
super().__init__(
|
||||
data_directory=data_directory,
|
||||
device=device,
|
||||
obs_dim=obs_dim,
|
||||
action_dim=action_dim,
|
||||
max_len_data=max_len_data,
|
||||
window_size=window_size,
|
||||
)
|
||||
|
||||
logging.info("Loading CubeStacking Dataset")
|
||||
|
||||
inputs = []
|
||||
actions = []
|
||||
masks = []
|
||||
|
||||
# for root, dirs, files in os.walk(self.data_directory):
|
||||
#
|
||||
# for mode_dir in dirs:
|
||||
|
||||
# state_files = glob.glob(os.path.join(root, mode_dir) + "/env*")
|
||||
# data_dir = os.path.join(sim_framework_path(data_directory), "local")
|
||||
# data_dir = sim_framework_path(data_directory)
|
||||
# state_files = glob.glob(data_dir + "/env*")
|
||||
|
||||
# bp_data_dir = sim_framework_path("environments/dataset/data/stacking/all_data_new")
|
||||
# state_files = np.load(sim_framework_path(data_directory), allow_pickle=True)
|
||||
|
||||
# bp_data_dir = sim_framework_path("environments/dataset/data/stacking/single_test")
|
||||
# state_files = os.listdir(bp_data_dir)
|
||||
|
||||
# random.seed(0)
|
||||
#
|
||||
# data_dir = sim_framework_path(data_directory)
|
||||
# state_files = os.listdir(data_dir)
|
||||
#
|
||||
# random.shuffle(state_files)
|
||||
#
|
||||
# if data == "train":
|
||||
# env_state_files = state_files[20:]
|
||||
# elif data == "eval":
|
||||
# env_state_files = state_files[:20]
|
||||
# else:
|
||||
# assert False, "wrong data type"
|
||||
|
||||
data_dir = sim_framework_path("data/stacking/all_data")
|
||||
state_files = np.load(sim_framework_path(data_directory), allow_pickle=True)
|
||||
|
||||
for file in state_files:
|
||||
with open(os.path.join(data_dir, file), "rb") as f:
|
||||
env_state = pickle.load(f)
|
||||
|
||||
# lengths.append(len(env_state['robot']['des_c_pos']))
|
||||
|
||||
zero_obs = np.zeros((1, self.max_len_data, self.obs_dim), dtype=np.float32)
|
||||
zero_action = np.zeros(
|
||||
(1, self.max_len_data, self.action_dim), dtype=np.float32
|
||||
)
|
||||
zero_mask = np.zeros((1, self.max_len_data), dtype=np.float32)
|
||||
|
||||
# robot and box positions
|
||||
robot_des_j_pos = env_state["robot"]["des_j_pos"]
|
||||
robot_des_j_vel = env_state["robot"]["des_j_vel"]
|
||||
|
||||
robot_des_c_pos = env_state["robot"]["des_c_pos"]
|
||||
robot_des_quat = env_state["robot"]["des_c_quat"]
|
||||
|
||||
robot_c_pos = env_state["robot"]["c_pos"]
|
||||
robot_c_quat = env_state["robot"]["c_quat"]
|
||||
|
||||
robot_j_pos = env_state["robot"]["j_pos"]
|
||||
robot_j_vel = env_state["robot"]["j_vel"]
|
||||
|
||||
robot_gripper = np.expand_dims(env_state["robot"]["gripper_width"], -1)
|
||||
# pred_gripper = np.zeros(robot_gripper.shape, dtype=np.float32)
|
||||
# pred_gripper[robot_gripper > 0.075] = 1
|
||||
|
||||
sim_steps = np.expand_dims(np.arange(len(robot_des_j_pos)), -1)
|
||||
|
||||
red_box_pos = env_state["red-box"]["pos"]
|
||||
red_box_quat = np.tan(quat2euler(env_state["red-box"]["quat"])[:, -1:])
|
||||
# red_box_quat = np.concatenate((np.sin(red_box_quat), np.cos(red_box_quat)), axis=-1)
|
||||
|
||||
green_box_pos = env_state["green-box"]["pos"]
|
||||
green_box_quat = np.tan(quat2euler(env_state["green-box"]["quat"])[:, -1:])
|
||||
# green_box_quat = np.concatenate((np.sin(green_box_quat), np.cos(green_box_quat)), axis=-1)
|
||||
|
||||
blue_box_pos = env_state["blue-box"]["pos"]
|
||||
blue_box_quat = np.tan(quat2euler(env_state["blue-box"]["quat"])[:, -1:])
|
||||
# blue_box_quat = np.concatenate((np.sin(blue_box_quat), np.cos(blue_box_quat)), axis=-1)
|
||||
|
||||
# target_box_pos = env_state['target-box']['pos'] #- robot_c_pos
|
||||
|
||||
# input_state = np.concatenate((robot_des_c_pos, robot_des_quat, pred_gripper, red_box_pos, red_box_quat), axis=-1)
|
||||
|
||||
# input_state = np.concatenate((robot_des_j_pos, robot_gripper, blue_box_pos, blue_box_quat), axis=-1)
|
||||
|
||||
input_state = np.concatenate(
|
||||
(
|
||||
robot_des_j_pos,
|
||||
robot_gripper,
|
||||
red_box_pos,
|
||||
red_box_quat,
|
||||
green_box_pos,
|
||||
green_box_quat,
|
||||
blue_box_pos,
|
||||
blue_box_quat,
|
||||
),
|
||||
axis=-1,
|
||||
)
|
||||
|
||||
# input_state = np.concatenate((robot_des_j_pos, robot_des_j_vel, robot_c_pos, robot_c_quat, green_box_pos, green_box_quat,
|
||||
# target_box_pos), axis=-1)
|
||||
|
||||
vel_state = robot_des_j_pos[1:] - robot_des_j_pos[:-1]
|
||||
|
||||
valid_len = len(input_state) - 1
|
||||
|
||||
zero_obs[0, :valid_len, :] = input_state[:-1]
|
||||
zero_action[0, :valid_len, :] = np.concatenate(
|
||||
(vel_state, robot_gripper[1:]), axis=-1
|
||||
)
|
||||
zero_mask[0, :valid_len] = 1
|
||||
|
||||
inputs.append(zero_obs)
|
||||
actions.append(zero_action)
|
||||
masks.append(zero_mask)
|
||||
|
||||
# shape: B, T, n
|
||||
self.observations = torch.from_numpy(np.concatenate(inputs)).to(device).float()
|
||||
self.actions = torch.from_numpy(np.concatenate(actions)).to(device).float()
|
||||
self.masks = torch.from_numpy(np.concatenate(masks)).to(device).float()
|
||||
|
||||
self.num_data = len(self.observations)
|
||||
|
||||
self.slices = self.get_slices()
|
||||
|
||||
def get_slices(self):
|
||||
slices = []
|
||||
|
||||
min_seq_length = np.inf
|
||||
for i in range(self.num_data):
|
||||
T = self.get_seq_length(i)
|
||||
min_seq_length = min(T, min_seq_length)
|
||||
|
||||
if T - self.window_size < 0:
|
||||
print(
|
||||
f"Ignored short sequence #{i}: len={T}, window={self.window_size}"
|
||||
)
|
||||
else:
|
||||
slices += [
|
||||
(i, start, start + self.window_size)
|
||||
for start in range(T - self.window_size + 1)
|
||||
] # slice indices follow convention [start, end)
|
||||
|
||||
return slices
|
||||
|
||||
def get_seq_length(self, idx):
|
||||
return int(self.masks[idx].sum().item())
|
||||
|
||||
def get_all_actions(self):
|
||||
result = []
|
||||
# mask out invalid actions
|
||||
for i in range(len(self.masks)):
|
||||
T = int(self.masks[i].sum().item())
|
||||
result.append(self.actions[i, :T, :])
|
||||
return torch.cat(result, dim=0)
|
||||
|
||||
def get_all_observations(self):
|
||||
result = []
|
||||
# mask out invalid observations
|
||||
for i in range(len(self.masks)):
|
||||
T = int(self.masks[i].sum().item())
|
||||
result.append(self.observations[i, :T, :])
|
||||
return torch.cat(result, dim=0)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.slices)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
|
||||
i, start, end = self.slices[idx]
|
||||
|
||||
obs = self.observations[i, start:end]
|
||||
act = self.actions[i, start:end]
|
||||
mask = self.masks[i, start:end]
|
||||
|
||||
return obs, act, mask
|
||||
|
||||
|
||||
class Stacking_Img_Dataset(TrajectoryDataset):
|
||||
def __init__(
|
||||
self,
|
||||
data_directory: os.PathLike,
|
||||
# data='train',
|
||||
device="cpu",
|
||||
obs_dim: int = 20,
|
||||
action_dim: int = 2,
|
||||
max_len_data: int = 256,
|
||||
window_size: int = 1,
|
||||
):
|
||||
|
||||
super().__init__(
|
||||
data_directory=data_directory,
|
||||
device=device,
|
||||
obs_dim=obs_dim,
|
||||
action_dim=action_dim,
|
||||
max_len_data=max_len_data,
|
||||
window_size=window_size,
|
||||
)
|
||||
|
||||
logging.info("Loading CubeStacking Dataset")
|
||||
|
||||
inputs = []
|
||||
actions = []
|
||||
masks = []
|
||||
|
||||
# TODO: insert data_dir here
|
||||
state_files = np.load(sim_framework_path(data_directory), allow_pickle=True)
|
||||
|
||||
bp_cam_imgs = []
|
||||
inhand_cam_imgs = []
|
||||
|
||||
for file in tqdm(state_files):
|
||||
with open(os.path.join(data_dir, "state", file), "rb") as f:
|
||||
env_state = pickle.load(f)
|
||||
|
||||
# lengths.append(len(env_state['robot']['des_c_pos']))
|
||||
|
||||
zero_obs = np.zeros((1, self.max_len_data, self.obs_dim), dtype=np.float32)
|
||||
zero_action = np.zeros(
|
||||
(1, self.max_len_data, self.action_dim), dtype=np.float32
|
||||
)
|
||||
zero_mask = np.zeros((1, self.max_len_data), dtype=np.float32)
|
||||
|
||||
# robot and box positions
|
||||
robot_des_j_pos = env_state["robot"]["des_j_pos"]
|
||||
robot_des_j_vel = env_state["robot"]["des_j_vel"]
|
||||
|
||||
robot_des_c_pos = env_state["robot"]["des_c_pos"]
|
||||
robot_des_quat = env_state["robot"]["des_c_quat"]
|
||||
|
||||
robot_c_pos = env_state["robot"]["c_pos"]
|
||||
robot_c_quat = env_state["robot"]["c_quat"]
|
||||
|
||||
robot_j_pos = env_state["robot"]["j_pos"]
|
||||
robot_j_vel = env_state["robot"]["j_vel"]
|
||||
|
||||
robot_gripper = np.expand_dims(env_state["robot"]["gripper_width"], -1)
|
||||
# pred_gripper = np.zeros(robot_gripper.shape, dtype=np.float32)
|
||||
# pred_gripper[robot_gripper > 0.075] = 1
|
||||
|
||||
file_name = os.path.basename(file).split(".")[0]
|
||||
|
||||
###############################################################
|
||||
bp_images = []
|
||||
bp_imgs = glob.glob(data_dir + "/images/bp-cam/" + file_name + "/*")
|
||||
bp_imgs.sort(key=lambda x: int(os.path.basename(x).split(".")[0]))
|
||||
|
||||
for img in bp_imgs:
|
||||
image = cv2.imread(img).astype(np.float32)
|
||||
image = image.transpose((2, 0, 1)) / 255.0
|
||||
|
||||
image = torch.from_numpy(image).to(self.device).float().unsqueeze(0)
|
||||
|
||||
bp_images.append(image)
|
||||
|
||||
bp_images = torch.concatenate(bp_images, dim=0)
|
||||
################################################################
|
||||
inhand_imgs = glob.glob(data_dir + "/images/inhand-cam/" + file_name + "/*")
|
||||
inhand_imgs.sort(key=lambda x: int(os.path.basename(x).split(".")[0]))
|
||||
inhand_images = []
|
||||
for img in inhand_imgs:
|
||||
image = cv2.imread(img).astype(np.float32)
|
||||
image = image.transpose((2, 0, 1)) / 255.0
|
||||
|
||||
image = torch.from_numpy(image).to(self.device).float().unsqueeze(0)
|
||||
|
||||
inhand_images.append(image)
|
||||
inhand_images = torch.concatenate(inhand_images, dim=0)
|
||||
##################################################################
|
||||
|
||||
input_state = np.concatenate((robot_des_j_pos, robot_gripper), axis=-1)
|
||||
vel_state = robot_des_j_pos[1:] - robot_des_j_pos[:-1]
|
||||
|
||||
valid_len = len(input_state) - 1
|
||||
|
||||
zero_obs[0, :valid_len, :] = input_state[:-1]
|
||||
zero_action[0, :valid_len, :] = np.concatenate(
|
||||
(vel_state, robot_gripper[1:]), axis=-1
|
||||
)
|
||||
zero_mask[0, :valid_len] = 1
|
||||
|
||||
bp_cam_imgs.append(bp_images)
|
||||
inhand_cam_imgs.append(inhand_images)
|
||||
|
||||
inputs.append(zero_obs)
|
||||
actions.append(zero_action)
|
||||
masks.append(zero_mask)
|
||||
|
||||
self.bp_cam_imgs = bp_cam_imgs
|
||||
self.inhand_cam_imgs = inhand_cam_imgs
|
||||
|
||||
# shape: B, T, n
|
||||
self.observations = torch.from_numpy(np.concatenate(inputs)).to(device).float()
|
||||
self.actions = torch.from_numpy(np.concatenate(actions)).to(device).float()
|
||||
self.masks = torch.from_numpy(np.concatenate(masks)).to(device).float()
|
||||
|
||||
self.num_data = len(self.observations)
|
||||
|
||||
self.slices = self.get_slices()
|
||||
|
||||
def get_slices(self):
|
||||
slices = []
|
||||
|
||||
min_seq_length = np.inf
|
||||
for i in range(self.num_data):
|
||||
T = self.get_seq_length(i)
|
||||
min_seq_length = min(T, min_seq_length)
|
||||
|
||||
if T - self.window_size < 0:
|
||||
print(
|
||||
f"Ignored short sequence #{i}: len={T}, window={self.window_size}"
|
||||
)
|
||||
else:
|
||||
slices += [
|
||||
(i, start, start + self.window_size)
|
||||
for start in range(T - self.window_size + 1)
|
||||
] # slice indices follow convention [start, end)
|
||||
|
||||
return slices
|
||||
|
||||
def get_seq_length(self, idx):
|
||||
return int(self.masks[idx].sum().item())
|
||||
|
||||
def get_all_actions(self):
|
||||
result = []
|
||||
# mask out invalid actions
|
||||
for i in range(len(self.masks)):
|
||||
T = int(self.masks[i].sum().item())
|
||||
result.append(self.actions[i, :T, :])
|
||||
return torch.cat(result, dim=0)
|
||||
|
||||
def get_all_observations(self):
|
||||
result = []
|
||||
# mask out invalid observations
|
||||
for i in range(len(self.masks)):
|
||||
T = int(self.masks[i].sum().item())
|
||||
result.append(self.observations[i, :T, :])
|
||||
return torch.cat(result, dim=0)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.slices)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
|
||||
i, start, end = self.slices[idx]
|
||||
|
||||
obs = self.observations[i, start:end]
|
||||
act = self.actions[i, start:end]
|
||||
mask = self.masks[i, start:end]
|
||||
|
||||
bp_imgs = self.bp_cam_imgs[i][start:end]
|
||||
inhand_imgs = self.inhand_cam_imgs[i][start:end]
|
||||
|
||||
return bp_imgs, inhand_imgs, obs, act, mask
|
||||
Reference in New Issue
Block a user