From 71df3c740947a0f6702e66d899b8e5b837989bab Mon Sep 17 00:00:00 2001 From: Antonin RAFFIN Date: Thu, 23 Apr 2020 14:56:05 +0200 Subject: [PATCH] Add docstrings and missing types --- torchy_baselines/common/identity_env.py | 61 +++++++++++-------- torchy_baselines/common/policies.py | 1 - torchy_baselines/common/type_aliases.py | 4 +- .../common/vec_env/vec_transpose.py | 21 ++++++- 4 files changed, 58 insertions(+), 29 deletions(-) diff --git a/torchy_baselines/common/identity_env.py b/torchy_baselines/common/identity_env.py index 526f5d3..3975d7e 100644 --- a/torchy_baselines/common/identity_env.py +++ b/torchy_baselines/common/identity_env.py @@ -1,10 +1,13 @@ -from typing import List +from typing import List, Union import numpy as np from gym import Env from gym.spaces import Discrete, MultiDiscrete, MultiBinary, Box +from torchy_baselines.common.type_aliases import GymStepReturn, GymObs + + class IdentityEnv(Env): def __init__(self, dim, ep_length=100): """ @@ -20,30 +23,32 @@ class IdentityEnv(Env): self.dim = dim self.reset() - def reset(self): + def reset(self) -> GymObs: self.current_step = 0 self._choose_next_state() return self.state - def step(self, action): + def step(self, action: Union[int, np.ndarray]) -> GymStepReturn: reward = self._get_reward(action) self._choose_next_state() self.current_step += 1 done = self.current_step >= self.ep_length return self.state, reward, done, {} - def _choose_next_state(self): + def _choose_next_state(self) -> None: self.state = self.action_space.sample() - def _get_reward(self, action): - return 1 if np.all(self.state == action) else 0 + def _get_reward(self, action: Union[int, np.ndarray]) -> float: + return 1.0 if np.all(self.state == action) else 0.0 - def render(self, mode='human'): + def render(self, mode: str = 'human') -> None: pass class IdentityEnvBox(IdentityEnv): - def __init__(self, low=-1, high=1, eps=0.05, ep_length=100): + def __init__(self, low: float = -1.0, + high: float = 1.0, eps: float = 0.05, + ep_length: int = 100): """ Identity environment for testing purposes @@ -58,27 +63,27 @@ class IdentityEnvBox(IdentityEnv): self.eps = eps self.reset() - def reset(self): + def reset(self) -> np.ndarray: self.current_step = 0 self._choose_next_state() return self.state - def step(self, action): + def step(self, action: np.ndarray) -> GymStepReturn: reward = self._get_reward(action) self._choose_next_state() self.current_step += 1 done = self.current_step >= self.ep_length return self.state, reward, done, {} - def _choose_next_state(self): + def _choose_next_state(self) -> None: self.state = self.observation_space.sample() - def _get_reward(self, action): - return 1 if (self.state - self.eps) <= action <= (self.state + self.eps) else 0 + def _get_reward(self, action: np.ndarray) -> float: + return 1.0 if (self.state - self.eps) <= action <= (self.state + self.eps) else 0.0 class IdentityEnvMultiDiscrete(IdentityEnv): - def __init__(self, dim, ep_length=100): + def __init__(self, dim: int, ep_length: int = 100): """ Identity environment for testing purposes @@ -92,7 +97,7 @@ class IdentityEnvMultiDiscrete(IdentityEnv): class IdentityEnvMultiBinary(IdentityEnv): - def __init__(self, dim, ep_length=100): + def __init__(self, dim: int, ep_length: int = 100): """ Identity environment for testing purposes @@ -105,35 +110,39 @@ class IdentityEnvMultiBinary(IdentityEnv): self.reset() - class FakeImageEnv(Env): + """ + Fake image environment for testing purposes, it mimics Atari games. + + :param action_dim: (int) Number of discrete actions + :param screen_height: (int) Height of the image + :param screen_width: (int) Width of the image + :param n_channels: (int) Number of color channels + :param discrete: (bool) + """ def __init__(self, action_dim: int = 6, screen_height: int = 210, screen_width: int = 160, n_channels: int = 3, discrete: bool = True): - """ - Fake atari environment for testing purposes. - """ - self.observation_space = Box(low=0, high=255, shape=(screen_height, screen_width, n_channels), dtype=np.uint8) + + self.observation_space = Box(low=0, high=255, shape=(screen_height, screen_width, + n_channels), dtype=np.uint8) if discrete: self.action_space = Discrete(action_dim) else: self.action_space = Box(low=-1, high=1, shape=(5,), dtype=np.float32) self.ep_length = 10 - def reset(self): + def reset(self) -> np.ndarray: self.current_step = 0 return self.observation_space.sample() - def step(self, action: int): + def step(self, action: Union[np.ndarray, int]) -> GymStepReturn: reward = 0.0 self.current_step += 1 done = self.current_step >= self.ep_length return self.observation_space.sample(), reward, done, {} - def render(self, mode='human'): + def render(self, mode: str = 'human') -> None: pass - - def get_action_meanings(self) -> List[str]: - return ['NOOP'] diff --git a/torchy_baselines/common/policies.py b/torchy_baselines/common/policies.py index fc3f7a6..6dd787c 100644 --- a/torchy_baselines/common/policies.py +++ b/torchy_baselines/common/policies.py @@ -59,7 +59,6 @@ class NatureCNN(BaseFeaturesExtractor): def __init__(self, observation_space: gym.spaces.Box, features_dim: int = 512): super(NatureCNN, self).__init__(observation_space, features_dim) - # TODO: custom init? # We assume CxWxH images (channels first) # Re-ordering will be done by pre-preprocessing or wrapper assert is_image_space(observation_space), ('You should use NatureCNN ' diff --git a/torchy_baselines/common/type_aliases.py b/torchy_baselines/common/type_aliases.py index db6b453..1026361 100644 --- a/torchy_baselines/common/type_aliases.py +++ b/torchy_baselines/common/type_aliases.py @@ -1,7 +1,7 @@ """ Common aliases for type hint """ -from typing import Union, Dict, Any, NamedTuple, Optional, List, Callable +from typing import Union, Dict, Any, NamedTuple, Optional, List, Callable, Tuple import numpy as np import torch as th @@ -12,6 +12,8 @@ from torchy_baselines.common.callbacks import BaseCallback GymEnv = Union[gym.Env, VecEnv] +GymObs = Union[Tuple, Dict[str, Any], np.ndarray, int] +GymStepReturn = Tuple[GymObs, float, bool, Dict] TensorDict = Dict[str, th.Tensor] OptimizerStateDict = Dict[str, Any] MaybeCallback = Union[None, Callable, List[BaseCallback], BaseCallback] diff --git a/torchy_baselines/common/vec_env/vec_transpose.py b/torchy_baselines/common/vec_env/vec_transpose.py index f4ef501..3b13c39 100644 --- a/torchy_baselines/common/vec_env/vec_transpose.py +++ b/torchy_baselines/common/vec_env/vec_transpose.py @@ -1,15 +1,22 @@ import warnings +import typing import numpy as np from gym import spaces from torchy_baselines.common.vec_env.base_vec_env import VecEnv, VecEnvWrapper from torchy_baselines.common.preprocessing import is_image_space +if typing.TYPE_CHECKING: + from torchy_baselines.common.type_aliases import GymStepReturn + + class VecTransposeImage(VecEnvWrapper): """ Re-order channels, from WxHxC to CxWxH. + It is required for PyTorch convolution layers. + :param venv: (VecEnv) """ def __init__(self, venv: VecEnv): @@ -20,6 +27,12 @@ class VecTransposeImage(VecEnvWrapper): @staticmethod def transpose_space(observation_space: spaces.Box) -> spaces.Box: + """ + Transpose an observation space (re-order channels). + + :param observation_space: (spaces.Box) + :return: (spaces.Box) + """ assert is_image_space(observation_space), 'The observation space must be an image' width, height, channels = observation_space.shape new_shape = (channels, width, height) @@ -27,11 +40,17 @@ class VecTransposeImage(VecEnvWrapper): @staticmethod def transpose_image(image: np.ndarray) -> np.ndarray: + """ + Transpose an image or batch of images (re-order channels). + + :param image: (np.ndarray) + :return: (np.ndarray) + """ if len(image.shape) == 3: return np.transpose(image, (2, 0, 1)) return np.transpose(image, (0, 3, 1, 2)) - def step_wait(self): + def step_wait(self) -> 'GymStepReturn': observations, rewards, dones, infos = self.venv.step_wait() return self.transpose_image(observations), rewards, dones, infos