calc episodic infos for some fancy envs

This commit is contained in:
2023-01-26 17:27:34 +01:00
parent f37c8caaa4
commit 6f1837bda5
3 changed files with 7 additions and 3 deletions
+5 -3
View File
@@ -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):