| 275 | |
| 276 | |
| 277 | class FrameStack(gym.Wrapper): |
| 278 | def __init__(self, env, k): |
| 279 | gym.Wrapper.__init__(self, env) |
| 280 | self._k = k |
| 281 | self._frames = deque([], maxlen=k) |
| 282 | self._qpos = deque([], maxlen=k) |
| 283 | shp = env.observation_space.shape |
| 284 | self.observation_space = gym.spaces.Box( |
| 285 | low=0, |
| 286 | high=1, |
| 287 | shape=((shp[0] * k,) + shp[1:]), |
| 288 | dtype=env.observation_space.dtype |
| 289 | ) |
| 290 | self._max_episode_steps = env._max_episode_steps |
| 291 | |
| 292 | def reset(self): |
| 293 | obs = self.env.reset() |
| 294 | qpos = self.env.get_qpos() |
| 295 | for _ in range(self._k): |
| 296 | self._frames.append(obs) |
| 297 | self._qpos.append(qpos) |
| 298 | return self._get_obs(), self._get_qpos() |
| 299 | |
| 300 | def step(self, action): |
| 301 | obs, reward, done, info = self.env.step(action) |
| 302 | |
| 303 | self._frames.append(obs) |
| 304 | self._qpos.append(info['qpos']) |
| 305 | return self._get_obs(), reward, done, self._get_qpos() |
| 306 | |
| 307 | def _get_obs(self): |
| 308 | assert len(self._frames) == self._k |
| 309 | return np.concatenate(list(self._frames), axis=0) |
| 310 | |
| 311 | def _get_qpos(self): |
| 312 | assert len(self._qpos) == self._k |
| 313 | #return np.concatenate(list(self._qpos), axis=0) |
| 314 | return np.stack(list(self._qpos)) |
nothing calls this directly
no outgoing calls
no test coverage detected