This commit is contained in:
Maximilian Huettenrauch
2021-02-11 12:32:32 +01:00
parent c81378b9e7
commit 13a292f0e0
10 changed files with 116 additions and 155 deletions
+47 -51
View File
@@ -7,9 +7,54 @@ import multiprocessing as mp
import sys
def _worker(index, env_fn, pipe, parent_pipe, shared_memory, error_queue):
assert shared_memory is None
env = env_fn()
parent_pipe.close()
try:
while True:
command, data = pipe.recv()
if command == 'reset':
observation = env.reset()
pipe.send((observation, True))
elif command == 'step':
observation, reward, done, info = env.step(data)
if done:
observation = env.reset()
pipe.send(((observation, reward, done, info), True))
elif command == 'rollout':
rewards = []
infos = []
for p, c in zip(*data):
reward, info = env.rollout(p, c)
rewards.append(reward)
infos.append(info)
pipe.send(((rewards, infos), (True, ) * len(rewards)))
elif command == 'seed':
env.seed(data)
pipe.send((None, True))
elif command == 'close':
env.close()
pipe.send((None, True))
break
elif command == 'idle':
pipe.send((None, True))
elif command == '_check_observation_space':
pipe.send((data == env.observation_space, True))
else:
raise RuntimeError('Received unknown command `{0}`. Must '
'be one of {`reset`, `step`, `seed`, `close`, '
'`_check_observation_space`}.'.format(command))
except (KeyboardInterrupt, Exception):
error_queue.put((index,) + sys.exc_info()[:2])
pipe.send((None, False))
finally:
env.close()
class DmpAsyncVectorEnv(gym.vector.AsyncVectorEnv):
def __init__(self, env_fns, n_samples, observation_space=None, action_space=None,
shared_memory=True, copy=True, context=None, daemon=True, worker=None):
shared_memory=False, copy=True, context="spawn", daemon=True, worker=_worker):
super(DmpAsyncVectorEnv, self).__init__(env_fns,
observation_space=observation_space,
action_space=action_space,
@@ -91,7 +136,7 @@ class DmpAsyncVectorEnv(gym.vector.AsyncVectorEnv):
self._raise_if_errors(successes)
self._state = AsyncState.DEFAULT
observations_list, rewards, dones, infos = [_flatten_list(r) for r in zip(*results)]
rewards, infos = [_flatten_list(r) for r in zip(*results)]
# for now, we ignore the observations and only return the rewards
@@ -109,55 +154,6 @@ class DmpAsyncVectorEnv(gym.vector.AsyncVectorEnv):
return self.rollout_wait()
def _worker(index, env_fn, pipe, parent_pipe, shared_memory, error_queue):
assert shared_memory is None
env = env_fn()
parent_pipe.close()
try:
while True:
command, data = pipe.recv()
if command == 'reset':
observation = env.reset()
pipe.send((observation, True))
elif command == 'step':
observation, reward, done, info = env.step(data)
if done:
observation = env.reset()
pipe.send(((observation, reward, done, info), True))
elif command == 'rollout':
observations = []
rewards = []
dones = []
infos = []
for p, c in zip(*data):
observation, reward, done, info = env.rollout(p, c)
observations.append(observation)
rewards.append(reward)
dones.append(done)
infos.append(info)
pipe.send(((observations, rewards, dones, infos), (True, ) * len(rewards)))
elif command == 'seed':
env.seed(data)
pipe.send((None, True))
elif command == 'close':
env.close()
pipe.send((None, True))
break
elif command == 'idle':
pipe.send((None, True))
elif command == '_check_observation_space':
pipe.send((data == env.observation_space, True))
else:
raise RuntimeError('Received unknown command `{0}`. Must '
'be one of {`reset`, `step`, `seed`, `close`, '
'`_check_observation_space`}.'.format(command))
except (KeyboardInterrupt, Exception):
error_queue.put((index,) + sys.exc_info()[:2])
pipe.send((None, False))
finally:
env.close()
def _flatten_obs(obs):
assert isinstance(obs, (list, tuple))
assert len(obs) > 0
+3 -8
View File
@@ -69,15 +69,11 @@ class DmpEnvWrapper(gym.Wrapper):
def __call__(self, params, contexts=None):
params = np.atleast_2d(params)
observations = []
rewards = []
dones = []
infos = []
for p, c in zip(params, contexts):
observation, reward, done, info = self.rollout(p, c)
observations.append(observation)
reward, info = self.rollout(p, c)
rewards.append(reward)
dones.append(done)
infos.append(info)
return np.array(rewards), infos
@@ -116,9 +112,8 @@ class DmpEnvWrapper(gym.Wrapper):
rews = []
infos = []
self.env.configure(context)
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(pos_vel[0], pos_vel[1])
@@ -132,4 +127,4 @@ class DmpEnvWrapper(gym.Wrapper):
reward = np.sum(rews)
return obs, reward, done, info
return reward, info