test-rerun/rsl_rl/storage/replay_buffer_multi.py

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