support for contexts, policy classes, pd controller example, breaking changes etc
This commit is contained in:
@@ -24,14 +24,14 @@ class DmpAsyncVectorEnv(gym.vector.AsyncVectorEnv):
|
||||
n=n_samples,
|
||||
fn=np.zeros)
|
||||
|
||||
def __call__(self, params):
|
||||
return self.rollout(params)
|
||||
def __call__(self, params, contexts=None):
|
||||
return self.rollout(params, contexts)
|
||||
|
||||
def rollout_async(self, actions):
|
||||
def rollout_async(self, params, contexts):
|
||||
"""
|
||||
Parameters
|
||||
----------
|
||||
actions : iterable of samples from `action_space`
|
||||
params : iterable of samples from `action_space`
|
||||
List of actions.
|
||||
"""
|
||||
self._assert_is_running()
|
||||
@@ -40,11 +40,17 @@ class DmpAsyncVectorEnv(gym.vector.AsyncVectorEnv):
|
||||
'for a pending call to `{0}` to complete.'.format(
|
||||
self._state.value), self._state.value)
|
||||
|
||||
actions = np.atleast_2d(actions)
|
||||
split_actions = np.array_split(actions, np.minimum(len(actions), self.num_envs))
|
||||
for pipe, action in zip(self.parent_pipes, split_actions):
|
||||
pipe.send(('rollout', action))
|
||||
for pipe in self.parent_pipes[len(split_actions):]:
|
||||
params = np.atleast_2d(params)
|
||||
split_params = np.array_split(params, np.minimum(len(params), self.num_envs))
|
||||
if contexts is None:
|
||||
split_contexts = np.array_split([None, ] * len(params), np.minimum(len(params), self.num_envs))
|
||||
else:
|
||||
split_contexts = np.array_split(contexts, np.minimum(len(contexts), self.num_envs))
|
||||
|
||||
assert np.all([len(p) == len(c) for p, c in zip(split_params, split_contexts)])
|
||||
for pipe, param, context in zip(self.parent_pipes, split_params, split_contexts):
|
||||
pipe.send(('rollout', (param, context)))
|
||||
for pipe in self.parent_pipes[len(split_params):]:
|
||||
pipe.send(('idle', None))
|
||||
self._state = AsyncState.WAITING_ROLLOUT
|
||||
|
||||
@@ -98,8 +104,8 @@ class DmpAsyncVectorEnv(gym.vector.AsyncVectorEnv):
|
||||
|
||||
return np.array(rewards), infos
|
||||
|
||||
def rollout(self, actions):
|
||||
self.rollout_async(actions)
|
||||
def rollout(self, actions, contexts):
|
||||
self.rollout_async(actions, contexts)
|
||||
return self.rollout_wait()
|
||||
|
||||
|
||||
@@ -123,8 +129,8 @@ def _worker(index, env_fn, pipe, parent_pipe, shared_memory, error_queue):
|
||||
rewards = []
|
||||
dones = []
|
||||
infos = []
|
||||
for d in data:
|
||||
observation, reward, done, info = env.rollout(d)
|
||||
for p, c in zip(*data):
|
||||
observation, reward, done, info = env.rollout(p, c)
|
||||
observations.append(observation)
|
||||
rewards.append(reward)
|
||||
dones.append(done)
|
||||
|
||||
@@ -5,7 +5,19 @@ import numpy as np
|
||||
import gym
|
||||
|
||||
|
||||
class DmpEnvWrapperBase(gym.Wrapper):
|
||||
def get_policy_class(policy_type):
|
||||
if policy_type == "motor":
|
||||
from alr_envs.utils.policies import PDController
|
||||
return PDController
|
||||
elif policy_type == "velocity":
|
||||
from alr_envs.utils.policies import VelController
|
||||
return VelController
|
||||
elif policy_type == "position":
|
||||
from alr_envs.utils.policies import PosController
|
||||
return PosController
|
||||
|
||||
|
||||
class DmpEnvWrapper(gym.Wrapper):
|
||||
def __init__(self,
|
||||
env,
|
||||
num_dof,
|
||||
@@ -17,8 +29,9 @@ class DmpEnvWrapperBase(gym.Wrapper):
|
||||
dt=0.01,
|
||||
learn_goal=False,
|
||||
post_traj_time=0.,
|
||||
policy=None):
|
||||
super(DmpEnvWrapperBase, self).__init__(env)
|
||||
policy_type=None,
|
||||
weights_scale=1.):
|
||||
super(DmpEnvWrapper, self).__init__(env)
|
||||
self.num_dof = num_dof
|
||||
self.num_basis = num_basis
|
||||
self.dim = num_dof * num_basis
|
||||
@@ -49,17 +62,19 @@ class DmpEnvWrapperBase(gym.Wrapper):
|
||||
dmp_goal_pos = final_pos
|
||||
|
||||
self.dmp.set_weights(dmp_weights, dmp_goal_pos)
|
||||
self.weights_scale = weights_scale
|
||||
|
||||
self.policy = policy
|
||||
policy_class = get_policy_class(policy_type)
|
||||
self.policy = policy_class(env)
|
||||
|
||||
def __call__(self, params):
|
||||
def __call__(self, params, contexts=None):
|
||||
params = np.atleast_2d(params)
|
||||
observations = []
|
||||
rewards = []
|
||||
dones = []
|
||||
infos = []
|
||||
for p in params:
|
||||
observation, reward, done, info = self.rollout(p)
|
||||
for p, c in zip(params, contexts):
|
||||
observation, reward, done, info = self.rollout(p, c)
|
||||
observations.append(observation)
|
||||
rewards.append(reward)
|
||||
dones.append(done)
|
||||
@@ -81,82 +96,11 @@ class DmpEnvWrapperBase(gym.Wrapper):
|
||||
goal_pos = None
|
||||
weight_matrix = np.reshape(params, [self.num_basis, self.num_dof])
|
||||
|
||||
return goal_pos, weight_matrix
|
||||
return goal_pos, weight_matrix * self.weights_scale
|
||||
|
||||
def rollout(self, params, render=False):
|
||||
def rollout(self, params, context=None, render=False):
|
||||
""" This function generates a trajectory based on a DMP and then does the usual loop over reset and step"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class DmpEnvWrapperPos(DmpEnvWrapperBase):
|
||||
"""
|
||||
Wrapper for gym environments which creates a trajectory in joint angle space
|
||||
"""
|
||||
def rollout(self, action, render=False):
|
||||
goal_pos, weight_matrix = self.goal_and_weights(action)
|
||||
if hasattr(self.env, "weight_matrix_scale"):
|
||||
weight_matrix = weight_matrix * self.env.weight_matrix_scale
|
||||
self.dmp.set_weights(weight_matrix, goal_pos)
|
||||
trajectory, _ = self.dmp.reference_trajectory(self.t)
|
||||
|
||||
if self.post_traj_steps > 0:
|
||||
trajectory = np.vstack([trajectory, np.tile(trajectory[-1, :], [self.post_traj_steps, 1])])
|
||||
|
||||
self._trajectory = trajectory
|
||||
|
||||
rews = []
|
||||
|
||||
self.env.reset()
|
||||
|
||||
for t, traj in enumerate(trajectory):
|
||||
obs, rew, done, info = self.env.step(traj)
|
||||
rews.append(rew)
|
||||
if render:
|
||||
self.env.render(mode="human")
|
||||
if done:
|
||||
break
|
||||
|
||||
reward = np.sum(rews)
|
||||
|
||||
return obs, reward, done, info
|
||||
|
||||
|
||||
class DmpEnvWrapperVel(DmpEnvWrapperBase):
|
||||
"""
|
||||
Wrapper for gym environments which creates a trajectory in joint velocity space
|
||||
"""
|
||||
def rollout(self, action, render=False):
|
||||
goal_pos, weight_matrix = self.goal_and_weights(action)
|
||||
if hasattr(self.env, "weight_matrix_scale"):
|
||||
weight_matrix = weight_matrix * self.env.weight_matrix_scale
|
||||
self.dmp.set_weights(weight_matrix, goal_pos)
|
||||
_, velocities = self.dmp.reference_trajectory(self.t)
|
||||
|
||||
rews = []
|
||||
infos = []
|
||||
|
||||
self.env.reset()
|
||||
|
||||
for t, vel in enumerate(velocities):
|
||||
obs, rew, done, info = self.env.step(vel)
|
||||
rews.append(rew)
|
||||
infos.append(info)
|
||||
if render:
|
||||
self.env.render(mode="human")
|
||||
if done:
|
||||
break
|
||||
|
||||
reward = np.sum(rews)
|
||||
|
||||
return obs, reward, done, info
|
||||
|
||||
|
||||
class DmpEnvWrapperPD(DmpEnvWrapperBase):
|
||||
"""
|
||||
Wrapper for gym environments which creates a trajectory in joint velocity space
|
||||
"""
|
||||
def rollout(self, action, render=False):
|
||||
goal_pos, weight_matrix = self.goal_and_weights(action)
|
||||
goal_pos, weight_matrix = self.goal_and_weights(params)
|
||||
if hasattr(self.env, "weight_matrix_scale"):
|
||||
weight_matrix = weight_matrix * self.env.weight_matrix_scale
|
||||
self.dmp.set_weights(weight_matrix, goal_pos)
|
||||
@@ -173,9 +117,11 @@ class DmpEnvWrapperPD(DmpEnvWrapperBase):
|
||||
infos = []
|
||||
|
||||
self.env.reset()
|
||||
if context is not None:
|
||||
self.env.configure(context)
|
||||
|
||||
for t, pos_vel in enumerate(zip(trajectory, velocity)):
|
||||
ac = self.policy.get_action(self.env, pos_vel[0], pos_vel[1])
|
||||
ac = self.policy.get_action(pos_vel[0], pos_vel[1])
|
||||
obs, rew, done, info = self.env.step(ac)
|
||||
rews.append(rew)
|
||||
infos.append(info)
|
||||
|
||||
@@ -1,15 +1,37 @@
|
||||
class PDController:
|
||||
def __init__(self, p_gains, d_gains):
|
||||
self.p_gains = p_gains
|
||||
self.d_gains = d_gains
|
||||
from alr_envs.mujoco.alr_mujoco_env import AlrMujocoEnv
|
||||
|
||||
def get_action(self, env, des_pos, des_vel):
|
||||
|
||||
class BaseController:
|
||||
def __init__(self, env: AlrMujocoEnv):
|
||||
self.env = env
|
||||
|
||||
def get_action(self, des_pos, des_vel):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class PosController(BaseController):
|
||||
def get_action(self, des_pos, des_vel):
|
||||
return des_pos
|
||||
|
||||
|
||||
class VelController(BaseController):
|
||||
def get_action(self, des_pos, des_vel):
|
||||
return des_vel
|
||||
|
||||
|
||||
class PDController(BaseController):
|
||||
def __init__(self, env):
|
||||
self.p_gains = env.p_gains
|
||||
self.d_gains = env.d_gains
|
||||
super(PDController, self).__init__(env)
|
||||
|
||||
def get_action(self, des_pos, des_vel):
|
||||
# TODO: make standardized ALRenv such that all of them have current_pos/vel attributes
|
||||
cur_pos = env.current_pos
|
||||
cur_vel = env.current_vel
|
||||
cur_pos = self.env.current_pos
|
||||
cur_vel = self.env.current_vel
|
||||
if len(des_pos) != len(cur_pos):
|
||||
des_pos = env.extend_des_pos(des_pos)
|
||||
des_pos = self.env.extend_des_pos(des_pos)
|
||||
if len(des_vel) != len(cur_vel):
|
||||
des_vel = env.extend_des_vel(des_vel)
|
||||
des_vel = self.env.extend_des_vel(des_vel)
|
||||
trq = self.p_gains * (des_pos - cur_pos) + self.d_gains * (des_vel - cur_vel)
|
||||
return trq
|
||||
|
||||
Reference in New Issue
Block a user