This commit is contained in:
Fabian
2022-07-07 10:47:04 +02:00
parent 4a3134d7be
commit fc00cf8a87
24 changed files with 235 additions and 302 deletions
+7 -5
View File
@@ -67,7 +67,9 @@ class BlackBoxWrapper(gym.ObservationWrapper):
def observation(self, observation):
# return context space if we are
return observation[self.env.context_mask] if self.return_context_observation else observation
obs = observation[self.env.context_mask] if self.return_context_observation else observation
# cast dtype because metaworld returns incorrect that throws gym error
return obs.astype(self.observation_space.dtype)
def get_trajectory(self, action: np.ndarray) -> Tuple:
clipped_params = np.clip(action, self.traj_gen_action_space.low, self.traj_gen_action_space.high)
@@ -147,7 +149,7 @@ class BlackBoxWrapper(gym.ObservationWrapper):
infos[k] = elems
if self.render_kwargs:
self.render(**self.render_kwargs)
self.env.render(**self.render_kwargs)
if done or self.replanning_schedule(self.current_pos, self.current_vel, obs, c_action,
t + 1 + self.current_traj_steps):
@@ -170,13 +172,13 @@ class BlackBoxWrapper(gym.ObservationWrapper):
def render(self, **kwargs):
"""Only set render options here, such that they can be used during the rollout.
This only needs to be called once"""
self.render_kwargs = kwargs or self.render_kwargs
self.render_kwargs = kwargs
# self.env.render(mode=self.render_mode, **self.render_kwargs)
self.env.render(**self.render_kwargs)
# self.env.render(**self.render_kwargs)
def reset(self, *, seed: Optional[int] = None, return_info: bool = False, options: Optional[dict] = None):
self.current_traj_steps = 0
return super(BlackBoxWrapper, self).reset(seed=seed, return_info=return_info, options=options)
return super(BlackBoxWrapper, self).reset()
def plot_trajs(self, des_trajs, des_vels):
import matplotlib.pyplot as plt
@@ -10,13 +10,12 @@ class MetaWorldController(BaseController):
Unlike the other Controllers, this is a special tracking_controller for MetaWorld environments.
They use a position delta for the xyz coordinates and a raw position for the gripper opening.
:param env: A position environment
"""
def get_action(self, des_pos, des_vel, c_pos, c_vel):
gripper_pos = des_pos[-1]
cur_pos = env.current_pos[:-1]
cur_pos = c_pos[:-1]
xyz_pos = des_pos[:-1]
assert xyz_pos.shape == cur_pos.shape, \
@@ -18,4 +18,4 @@ def get_controller(controller_type: str, **kwargs):
return MetaWorldController()
else:
raise ValueError(f"Specified controller type {controller_type} not supported, "
f"please choose one of {ALL_TYPES}.")
f"please choose one of {ALL_TYPES}.")