| 126 | |
| 127 | |
| 128 | class MaxAndSkipEnv(gym.Wrapper): |
| 129 | def __init__(self, env, skip=4): |
| 130 | """Return only every `skip`-th frame""" |
| 131 | gym.Wrapper.__init__(self, env) |
| 132 | # most recent raw observations (for max pooling across time steps) |
| 133 | self._obs_buffer = np.zeros((2,)+env.observation_space.shape, dtype=np.uint8) |
| 134 | self._skip = skip |
| 135 | self.max_frame = np.zeros(env.observation_space.shape, dtype=np.uint8) |
| 136 | |
| 137 | def step(self, action): |
| 138 | """Repeat action, sum reward, and max over last observations.""" |
| 139 | total_reward = 0.0 |
| 140 | done = None |
| 141 | for i in range(self._skip): |
| 142 | obs, reward, done, info = self.env.step(action) |
| 143 | if i == self._skip - 2: self._obs_buffer[0] = obs |
| 144 | if i == self._skip - 1: self._obs_buffer[1] = obs |
| 145 | total_reward += reward |
| 146 | if done: |
| 147 | break |
| 148 | # Note that the observation on the done=True frame |
| 149 | # doesn't matter |
| 150 | self.max_frame = self._obs_buffer.max(axis=0) |
| 151 | |
| 152 | return self.max_frame, total_reward, done, info |
| 153 | |
| 154 | def reset(self, **kwargs): |
| 155 | return self.env.reset(**kwargs) |
| 156 | |
| 157 | def render(self, mode='human', **kwargs): |
| 158 | img = self.max_frame |
| 159 | img = cv2.resize(img, (400, 400), interpolation=cv2.INTER_AREA).astype(np.uint8) |
| 160 | if mode == 'rgb_array': |
| 161 | return img |
| 162 | elif mode == 'human': |
| 163 | from gym.envs.classic_control import rendering |
| 164 | if self.viewer is None: |
| 165 | self.viewer = rendering.SimpleImageViewer() |
| 166 | self.viewer.imshow(img) |
| 167 | return self.viewer.isopen |
| 168 | |
| 169 | |
| 170 | class WarpFrame(gym.ObservationWrapper): |