calc episodic infos for some fancy envs
This commit is contained in:
@@ -6,7 +6,7 @@ from gym import spaces
|
||||
|
||||
from stable_baselines3.common.buffers import RolloutBuffer
|
||||
from stable_baselines3.common.vec_env import VecNormalize
|
||||
from stable_baselines3.common.vec_env import VecEnv
|
||||
from stable_baselines3.common.vec_env import VecEnv, DummyVecEnv
|
||||
from stable_baselines3.common.callbacks import BaseCallback
|
||||
from stable_baselines3.common.utils import obs_as_tensor
|
||||
|
||||
@@ -228,6 +228,10 @@ class GaussianRolloutCollectorAuxclass():
|
||||
actions, self.action_space.low, self.action_space.high)
|
||||
|
||||
new_obs, rewards, dones, infos = env.step(clipped_actions)
|
||||
if 'episode_end' in infos[0]:
|
||||
for i in range(len(infos)):
|
||||
if infos[i]['episode_end'] and 'episode' not in infos:
|
||||
infos[i]['episode'] = {'r': rewards[i]}
|
||||
if len(infos) and 'r' not in infos[0]:
|
||||
for i in range(len(infos)):
|
||||
if 'r' not in infos[i]:
|
||||
@@ -286,8 +290,6 @@ class GaussianRolloutCollectorAuxclass():
|
||||
:param infos: List of additional information about the transition.
|
||||
:param dones: Termination signals
|
||||
"""
|
||||
import pdb
|
||||
pdb.set_trace()
|
||||
if dones is None:
|
||||
dones = np.array([False] * len(infos))
|
||||
for idx, info in enumerate(infos):
|
||||
|
||||
Reference in New Issue
Block a user