wrapping all environments in timelimit to stop runs longer than 100s, since some others were erroring

This commit is contained in:
Arec
2022-02-21 17:54:29 -08:00
parent e7f8385628
commit e7f4f6a871
7 changed files with 49 additions and 35 deletions

View File

@@ -14,7 +14,7 @@ from src.options.envs import OptionsEnv
from src.safe_options.collisions import feasible
from intersim.envs import IntersimpleLidarFlatIncrementingAgent
from src.util.wrappers import OptionsTimeLimit, Setobs, TransformObservation
from src.util.wrappers import Setobs, TransformObservation
@dataclass
class Buffer:
@@ -251,10 +251,10 @@ obs_max = np.array([
[50, np.pi, 20, 20, np.pi, 1e-1],
]).reshape(-1)
def NormalizedSafeOptionsEvalEnv(max_episode_steps=float('inf'), safe_actions_collision_method=None, abort_unsafe_collision_method=None, **kwargs):
return OptionsTimeLimit(SafeOptionsEnv(Setobs(
def NormalizedSafeOptionsEvalEnv(safe_actions_collision_method=None, abort_unsafe_collision_method=None, **kwargs):
return SafeOptionsEnv(Setobs(
TransformObservation(IntersimpleLidarFlatIncrementingAgent(
n_rays=5,
**kwargs,
), lambda obs: (obs - obs_min) / (obs_max - obs_min + 1e-10))
), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5), (10, 5)], safe_actions_collision_method=safe_actions_collision_method, abort_unsafe_collision_method=abort_unsafe_collision_method), max_episode_steps=max_episode_steps)
), options=[(0, 5), (1, 5), (2, 5), (4, 5), (6, 5), (8, 5), (10, 5)], safe_actions_collision_method=safe_actions_collision_method, abort_unsafe_collision_method=abort_unsafe_collision_method)