| 26 | |
| 27 | ### The experience replay memory ### |
| 28 | class ReplayBuffer: |
| 29 | def __init__(self, obs_dim, act_dim, size): |
| 30 | self.obs1_buf = np.zeros([size, obs_dim], dtype=np.float32) |
| 31 | self.obs2_buf = np.zeros([size, obs_dim], dtype=np.float32) |
| 32 | self.acts_buf = np.zeros(size, dtype=np.uint8) |
| 33 | self.rews_buf = np.zeros(size, dtype=np.float32) |
| 34 | self.done_buf = np.zeros(size, dtype=np.uint8) |
| 35 | self.ptr, self.size, self.max_size = 0, 0, size |
| 36 | |
| 37 | def store(self, obs, act, rew, next_obs, done): |
| 38 | self.obs1_buf[self.ptr] = obs |
| 39 | self.obs2_buf[self.ptr] = next_obs |
| 40 | self.acts_buf[self.ptr] = act |
| 41 | self.rews_buf[self.ptr] = rew |
| 42 | self.done_buf[self.ptr] = done |
| 43 | self.ptr = (self.ptr+1) % self.max_size |
| 44 | self.size = min(self.size+1, self.max_size) |
| 45 | |
| 46 | def sample_batch(self, batch_size=32): |
| 47 | idxs = np.random.randint(0, self.size, size=batch_size) |
| 48 | return dict(s=self.obs1_buf[idxs], |
| 49 | s2=self.obs2_buf[idxs], |
| 50 | a=self.acts_buf[idxs], |
| 51 | r=self.rews_buf[idxs], |
| 52 | d=self.done_buf[idxs]) |
| 53 | |
| 54 | |
| 55 | |