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))