load and save for SharedReplay

Not sure if your interested in adding this, it just supports saving and loading SharedReplay. For example if you want to load previous DDPG experience.
This commit is contained in:
Mike Clark authored and GitHub committed 2018-02-02 15:11:44 +08:00
1 parent 63c5d5ea54
commit 9b7d21a0f2
1 file changed
+16
+16
View File
@@ -152,6 +152,22 @@ class SharedReplay:
with self.buffer_lock:
return self.sample_()
def state_dict(self):
return dict((key, getattr(self, key)) for key in ['actions', 'states', 'rewards', 'next_states', 'terminals', 'pos'])
def load_state_dict(self, state):
for key in ['actions', 'states', 'rewards', 'next_states', 'terminals', 'pos']:
val = state[key]
setattr(self, key, val)
def save(self, file_name):
with open(file_name, 'wb') as f:
torch.save(self.state_dict(), f)
def load(self, file_name):
state = torch.load(file_name)
self.load_state_dict(state)
class HighDimActionReplay:
def __init__(self, memory_size, batch_size, dtype=np.float32):
self.memory_size = memory_size