save normalizer and replay with ddpg agent

This commit is contained in:
Mike C
2018-01-22 10:48:19 +08:00
parent 9318818a97
commit 6a92f52b4f
2 changed files with 27 additions and 1 deletions
+15 -1
View File
@@ -36,9 +36,23 @@ class DDPGAgent:
target_param.data.copy_(target_param.data * (1.0 - self.config.target_network_mix) +
param.data * self.config.target_network_mix)
def state_dict(self):
return {
'worker_network': self.worker_network.state_dict(),
'replay': self.replay.state_dict(),
'state_normalizer': self.state_normalizer.state_dict(),
'reward_normalizer': self.reward_normalizer.state_dict()
}
def load_state_dict(self, saved):
self.worker_network.load_state_dict(saved['worker_network'])
self.replay.load_state_dict(saved['replay'])
self.state_normalizer.load_state_dict(saved['state_normalizer'])
self.reward_normalizer.load_state_dict(saved['reward_normalizer'])
def save(self, file_name):
with open(file_name, 'wb') as f:
torch.save(self.worker_network.state_dict(), f)
torch.save(self.state_dict(), f)
def episode(self, deterministic=False, video_recorder=None):
self.random_process.reset_states()
+12
View File
@@ -25,6 +25,12 @@ class Normalizer:
o = o.reshape(o_.shape)
return o
def state_dict(self):
return self.stats.state_dict()
def load_state_dict(self, saved):
self.stats.load_state_dict(saved)
class StaticNormalizer:
def __init__(self, o_size):
self.offline_stats = SharedStats(o_size)
@@ -47,6 +53,12 @@ class StaticNormalizer:
o = o.reshape(o_.shape)
return o
def state_dict(self):
return self.offline_stats.state_dict()
def load_state_dict(self, saved):
self.offline_stats.load_state_dict(saved)
class SharedStats:
def __init__(self, o_size):
self.m = torch.zeros(o_size)