start contextual dmp wrapper
This commit is contained in:
@@ -4,8 +4,8 @@ import numpy as np
|
||||
from _collections import defaultdict
|
||||
|
||||
|
||||
def make_env(env_id, rank, seed=0):
|
||||
env = gym.make(env_id)
|
||||
def make_env(env_id, rank, seed=0, **env_kwargs):
|
||||
env = gym.make(env_id, **env_kwargs)
|
||||
env.seed(seed + rank)
|
||||
return lambda: env
|
||||
|
||||
@@ -45,9 +45,9 @@ class AlrMpEnvSampler:
|
||||
An asynchronous sampler for non contextual MPWrapper environments. A sampler object can be called with a set of
|
||||
parameters and returns the corresponding final obs, rewards, dones and info dicts.
|
||||
"""
|
||||
def __init__(self, env_id, num_envs, seed=0):
|
||||
def __init__(self, env_id, num_envs, seed=0, **env_kwargs):
|
||||
self.num_envs = num_envs
|
||||
self.env = AsyncVectorEnv([make_env(env_id, seed, i) for i in range(num_envs)])
|
||||
self.env = AsyncVectorEnv([make_env(env_id, seed, i, **env_kwargs) for i in range(num_envs)])
|
||||
|
||||
def __call__(self, params):
|
||||
params = np.atleast_2d(params)
|
||||
@@ -67,6 +67,36 @@ class AlrMpEnvSampler:
|
||||
_flatten_list(vals['done'])[:n_samples], _flatten_list(vals['info'])[:n_samples]
|
||||
|
||||
|
||||
class AlrContextualMpEnvSampler:
|
||||
"""
|
||||
An asynchronous sampler for non contextual MPWrapper environments. A sampler object can be called with a set of
|
||||
parameters and returns the corresponding final obs, rewards, dones and info dicts.
|
||||
"""
|
||||
def __init__(self, env_id, num_envs, seed=0, **env_kwargs):
|
||||
self.num_envs = num_envs
|
||||
self.env = AsyncVectorEnv([make_env(env_id, seed, i, **env_kwargs) for i in range(num_envs)])
|
||||
|
||||
def __call__(self, dist, n_samples):
|
||||
|
||||
repeat = int(np.ceil(n_samples / self.env.num_envs))
|
||||
vals = defaultdict(list)
|
||||
for i in range(repeat):
|
||||
obs = self.env.reset()
|
||||
|
||||
new_contexts = obs[-2]
|
||||
new_samples = dist.sample(new_contexts)
|
||||
|
||||
obs, reward, done, info = self.env.step(p)
|
||||
vals['obs'].append(obs)
|
||||
vals['reward'].append(reward)
|
||||
vals['done'].append(done)
|
||||
vals['info'].append(info)
|
||||
|
||||
# do not return values above threshold
|
||||
return np.vstack(vals['obs'])[:n_samples], np.hstack(vals['reward'])[:n_samples],\
|
||||
_flatten_list(vals['done'])[:n_samples], _flatten_list(vals['info'])[:n_samples]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
env_name = "alr_envs:ALRBallInACupSimpleDMP-v0"
|
||||
n_cpu = 8
|
||||
|
||||
@@ -36,9 +36,10 @@ class DmpWrapper(MPWrapper):
|
||||
dt = env.dt if hasattr(env, "dt") else dt
|
||||
assert dt is not None
|
||||
start_pos = start_pos if start_pos is not None else env.start_pos if hasattr(env, "start_pos") else None
|
||||
assert start_pos is not None
|
||||
# TODO: assert start_pos is not None # start_pos will be set in initialize, do we need this here?
|
||||
if learn_goal:
|
||||
final_pos = np.zeros_like(start_pos) # arbitrary, will be learned
|
||||
# final_pos = np.zeros_like(start_pos) # arbitrary, will be learned
|
||||
final_pos = np.zeros((1, num_dof)) # arbitrary, will be learned
|
||||
else:
|
||||
final_pos = final_pos if final_pos is not None else start_pos if return_to_start else None
|
||||
assert final_pos is not None
|
||||
@@ -62,7 +63,10 @@ class DmpWrapper(MPWrapper):
|
||||
dmp = dmps.DMP(num_dof=num_dof, basis_generator=basis_generator, phase_generator=phase_generator,
|
||||
num_time_steps=int(duration / dt), dt=dt)
|
||||
|
||||
dmp.dmp_start_pos = start_pos.reshape((1, num_dof))
|
||||
# dmp.dmp_start_pos = start_pos.reshape((1, num_dof))
|
||||
# in a contextual environment, the start_pos may be not fixed, set in mp_rollout?
|
||||
# TODO: Should we set start_pos in init at all? It's only used after calling rollout anyway...
|
||||
dmp.dmp_start_pos = start_pos.reshape((1, num_dof)) if start_pos is not None else np.zeros((1, num_dof))
|
||||
|
||||
weights = np.zeros((num_basis, num_dof))
|
||||
goal_pos = np.zeros(num_dof) if self.learn_goal else final_pos
|
||||
@@ -87,6 +91,8 @@ class DmpWrapper(MPWrapper):
|
||||
return goal_pos * self.goal_scale, weight_matrix * self.weights_scale
|
||||
|
||||
def mp_rollout(self, action):
|
||||
if self.mp.start_pos is None:
|
||||
self.mp.start_pos = self.env.start_pos
|
||||
goal_pos, weight_matrix = self.goal_and_weights(action)
|
||||
self.mp.set_weights(weight_matrix, goal_pos)
|
||||
return self.mp.reference_trajectory(self.t)
|
||||
|
||||
@@ -61,6 +61,9 @@ class MPWrapper(gym.Wrapper, ABC):
|
||||
def configure(self, context):
|
||||
self.env.configure(context)
|
||||
|
||||
def reset(self):
|
||||
return self.env.reset()
|
||||
|
||||
def step(self, action: np.ndarray):
|
||||
""" This function generates a trajectory based on a DMP and then does the usual loop over reset and step"""
|
||||
trajectory, velocity = self.mp_rollout(action)
|
||||
@@ -78,8 +81,9 @@ class MPWrapper(gym.Wrapper, ABC):
|
||||
# TODO: @Max Why do we need this configure, states should be part of the model
|
||||
# TODO: Ask Onur if the context distribution needs to be outside the environment
|
||||
# TODO: For now create a new env with each context
|
||||
# TODO: Explicitly call reset before step to obtain context from obs?
|
||||
# self.env.configure(context)
|
||||
obs = self.env.reset()
|
||||
# obs = self.env.reset()
|
||||
info = {}
|
||||
|
||||
for t, pos_vel in enumerate(zip(trajectory, velocity)):
|
||||
|
||||
Reference in New Issue
Block a user