mirror of
https://github.com/wassname/DeepRL.git
synced 2026-08-22 11:40:47 +08:00
42 lines
1.4 KiB
Python
42 lines
1.4 KiB
Python
#######################################################################
|
|
# Copyright (C) 2017 Shangtong Zhang(zhangshangtong.cpp@gmail.com) #
|
|
# Permission given to modify the code as long as you keep this #
|
|
# declaration at the top #
|
|
#######################################################################
|
|
|
|
import numpy as np
|
|
|
|
class Replay:
|
|
def __init__(self, memory_size, batch_size):
|
|
self.memory_size = memory_size
|
|
self.batch_size = batch_size
|
|
self.data = []
|
|
self.pos = 0
|
|
|
|
def feed(self, experience):
|
|
if self.pos >= len(self.data):
|
|
self.data.append(experience)
|
|
else:
|
|
self.data[self.pos] = experience
|
|
self.pos = (self.pos + 1) % self.memory_size
|
|
|
|
def feed_batch(self, experience):
|
|
experience = zip(*experience)
|
|
for exp in experience:
|
|
self.feed(exp)
|
|
|
|
def sample(self, batch_size=None):
|
|
if batch_size is None:
|
|
batch_size = self.batch_size
|
|
|
|
sampled_indices = [np.random.randint(0, len(self.data)) for _ in range(batch_size)]
|
|
sampled_data = [self.data[ind] for ind in sampled_indices]
|
|
batch_data = list(map(lambda x: np.asarray(x), zip(*sampled_data)))
|
|
return batch_data
|
|
|
|
def size(self):
|
|
return len(self.data)
|
|
|
|
def empty(self):
|
|
return not len(self.data)
|