context wip
This commit is contained in:
@@ -56,6 +56,7 @@ class AlrMpEnvSampler:
|
||||
|
||||
vals = defaultdict(list)
|
||||
for p in split_params:
|
||||
self.env.reset()
|
||||
obs, reward, done, info = self.env.step(p)
|
||||
vals['obs'].append(obs)
|
||||
vals['reward'].append(reward)
|
||||
@@ -82,8 +83,9 @@ class AlrContextualMpEnvSampler:
|
||||
vals = defaultdict(list)
|
||||
for i in range(repeat):
|
||||
new_contexts = self.env.reset()
|
||||
|
||||
new_samples = dist.sample(new_contexts)
|
||||
vals['new_contexts'].append(new_contexts)
|
||||
new_samples, new_contexts = dist.sample(new_contexts)
|
||||
vals['new_samples'].append(new_samples)
|
||||
|
||||
obs, reward, done, info = self.env.step(new_samples)
|
||||
vals['obs'].append(obs)
|
||||
@@ -92,7 +94,8 @@ class AlrContextualMpEnvSampler:
|
||||
vals['info'].append(info)
|
||||
|
||||
# do not return values above threshold
|
||||
return np.vstack(vals['obs'])[:n_samples], np.hstack(vals['reward'])[:n_samples],\
|
||||
return np.vstack(vals['new_samples'])[:n_samples], np.vstack(vals['new_contexts'])[:n_samples], \
|
||||
np.vstack(vals['obs'])[:n_samples], np.hstack(vals['reward'])[:n_samples], \
|
||||
_flatten_list(vals['done'])[:n_samples], _flatten_list(vals['info'])[:n_samples]
|
||||
|
||||
|
||||
|
||||
@@ -98,7 +98,7 @@ class DmpWrapper(MPWrapper):
|
||||
|
||||
def mp_rollout(self, action):
|
||||
# if self.mp.start_pos is None:
|
||||
self.mp.dmp_start_pos = self.env.init_qpos # start_pos
|
||||
self.mp.dmp_start_pos = self.env.init_qpos.reshape((1, self.num_dof)) # 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)
|
||||
|
||||
@@ -22,7 +22,7 @@ class MPWrapper(gym.Wrapper, ABC):
|
||||
):
|
||||
super().__init__(env)
|
||||
|
||||
# self.num_dof = num_dof
|
||||
self.num_dof = num_dof
|
||||
# self.num_basis = num_basis
|
||||
# self.duration = duration # seconds
|
||||
|
||||
@@ -50,6 +50,7 @@ class MPWrapper(gym.Wrapper, ABC):
|
||||
# for p, c in zip(params, contexts):
|
||||
for p in params:
|
||||
# self.configure(c)
|
||||
# context = self.reset()
|
||||
ob, reward, done, info = self.step(p)
|
||||
obs.append(ob)
|
||||
rewards.append(reward)
|
||||
|
||||
Reference in New Issue
Block a user