mirror of
https://github.com/wassname/DeepRL.git
synced 2026-10-03 12:00:18 +08:00
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:
1 parent
63c5d5ea54
commit
9b7d21a0f2
1 file changed
+16
@@ -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
|
||||
|
||||
Reference in new issue
Block a user