diff --git a/agent/DDPG_agent.py b/agent/DDPG_agent.py index aa649b1..59ef543 100644 --- a/agent/DDPG_agent.py +++ b/agent/DDPG_agent.py @@ -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() diff --git a/utils/normalizer.py b/utils/normalizer.py index 8070679..72b59c7 100644 --- a/utils/normalizer.py +++ b/utils/normalizer.py @@ -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)