39 lines
1.4 KiB
Python
39 lines
1.4 KiB
Python
import torch
|
|
import numpy as np
|
|
|
|
|
|
class ReplayBufferMulti:
|
|
"""Fixed-size buffer to store experience tuples."""
|
|
|
|
def __init__(self, obs_dim, buffer_size, num_amp_frames, device):
|
|
"""Initialize a ReplayBuffer object.
|
|
Arguments:
|
|
buffer_size (int): maximum size of buffer
|
|
"""
|
|
self.states = torch.zeros(buffer_size, num_amp_frames, obs_dim).to(device)
|
|
self.num_amp_frames = num_amp_frames
|
|
self.buffer_size = buffer_size
|
|
self.device = device
|
|
|
|
self.step = 0
|
|
self.num_samples = 0
|
|
|
|
def insert(self, states):
|
|
"""Add new states to memory."""
|
|
num_states = states.shape[0]
|
|
start_idx = self.step
|
|
end_idx = self.step + num_states
|
|
if end_idx > self.buffer_size:
|
|
self.states[self.step:self.buffer_size] = states[:self.buffer_size - self.step]
|
|
self.states[:end_idx - self.buffer_size] = states[self.buffer_size - self.step:]
|
|
else:
|
|
self.states[start_idx:end_idx] = states
|
|
|
|
self.num_samples = min(self.buffer_size, max(end_idx, self.num_samples))
|
|
self.step = (self.step + num_states) % self.buffer_size
|
|
|
|
def feed_forward_generator(self, num_mini_batch, mini_batch_size):
|
|
for _ in range(num_mini_batch):
|
|
sample_idxs = np.random.choice(self.num_samples, size=mini_batch_size)
|
|
yield (self.states[sample_idxs].to(self.device))
|