Major update

This commit is contained in:
Shangtong Zhang
2017-11-11 20:52:57 -07:00
parent 8c197e183c
commit c1fdb1bd2f
17 changed files with 48 additions and 199 deletions
+2 -5
View File
@@ -7,6 +7,7 @@ import numpy as np
import torch
from torch.autograd import Variable
import torch.nn as nn
from utils import *
class AdvantageActorCritic:
def __init__(self, config, learning_network, target_network):
@@ -68,11 +69,7 @@ class AdvantageActorCritic:
self.optimizer.zero_grad()
loss.backward()
nn.utils.clip_grad_norm(self.worker_network.parameters(), config.gradient_clip)
for param, worker_param in zip(
self.learning_network.parameters(), self.worker_network.parameters()):
if param.grad is not None:
break
param._grad = worker_param.grad
sync_grad(self.learning_network, self.worker_network)
self.optimizer.step()
self.worker_network.load_state_dict(self.learning_network.state_dict())
self.worker_network.reset(terminal)
+1 -5
View File
@@ -92,11 +92,7 @@ class ContinuousAdvantageActorCritic:
actor_loss.backward()
critic_loss.backward()
nn.utils.clip_grad_norm(self.worker_network.parameters(), config.gradient_clip)
for param, worker_param in zip(
self.learning_network.parameters(), self.worker_network.parameters()):
if param.grad is not None:
break
param._grad = worker_param.grad
sync_grad(self.learning_network, self.worker_network)
self.actor_opt.step()
self.critic_opt.step()
self.worker_network.load_state_dict(self.learning_network.state_dict())
+11 -12
View File
@@ -30,11 +30,9 @@ class DeterministicPolicyGradient:
self.random_process = config.random_process_fn()
self.criterion = nn.MSELoss()
# self.state_normalizer = Normalizer(self.task.state_dim)
self.shared_state_normalizer = extra[0]
self.shared_state_normalizer, self.shared_reward_normalizer, self.replay = extra
self.state_normalizer = StaticNormalizer(self.task.state_dim)
# self.replay = config.replay_fn()
self.replay = extra[-1]
self.reward_normalizer = StaticNormalizer(1)
def soft_update(self, target, src):
for target_param, param in zip(target.parameters(), src.parameters()):
@@ -63,6 +61,7 @@ class DeterministicPolicyGradient:
done = (done or (config.max_episode_length and steps >= config.max_episode_length))
next_state = self.state_normalizer(next_state)
total_reward += reward
reward = self.reward_normalizer(reward)
if not deterministic:
self.replay.feed([state, action, reward, next_state, int(done)])
@@ -92,10 +91,7 @@ class DeterministicPolicyGradient:
self.critic_opt.zero_grad()
critic_loss.backward()
with config.network_lock:
for param, worker_param in zip(self.shared_network.critic.parameters(), critic.parameters()):
if param.grad is not None:
break
param._grad = worker_param.grad
sync_grad(self.shared_network.critic, critic)
self.critic_opt.step()
actions = actor.predict(states, False)
@@ -107,14 +103,17 @@ class DeterministicPolicyGradient:
self.actor_opt.zero_grad()
actions.backward(-var_actions.grad.data)
with config.network_lock:
for param, worker_param in zip(self.shared_network.actor.parameters(), actor.parameters()):
if param.grad is not None:
break
param._grad = worker_param.grad
sync_grad(self.shared_network.actor, actor)
self.actor_opt.step()
self.worker_network.load_state_dict(self.shared_network.state_dict())
self.soft_update(self.target_network, self.worker_network)
self.shared_state_normalizer.offline_stats.merge(self.state_normalizer.online_stats)
self.state_normalizer.online_stats.zero()
self.shared_reward_normalizer.offline_stats.merge(self.reward_normalizer.online_stats)
self.reward_normalizer.online_stats.zero()
return steps, total_reward
+2 -5
View File
@@ -7,6 +7,7 @@ import numpy as np
import torch
from torch.autograd import Variable
import torch.nn as nn
from utils import *
class NStepQLearning:
def __init__(self, config, learning_network, target_network):
@@ -63,11 +64,7 @@ class NStepQLearning:
self.optimizer.zero_grad()
loss.backward()
nn.utils.clip_grad_norm(self.worker_network.parameters(), config.gradient_clip)
for param, worker_param in zip(
self.learning_network.parameters(), self.worker_network.parameters()):
if param.grad is not None:
break
param._grad = worker_param.grad
sync_grad(self.learning_network, self.worker_network)
self.optimizer.step()
self.worker_network.load_state_dict(self.learning_network.state_dict())
self.worker_network.reset(terminal)
+2 -5
View File
@@ -7,6 +7,7 @@ import numpy as np
import torch
from torch.autograd import Variable
import torch.nn as nn
from utils import *
class OneStepQLearning:
def __init__(self, config, learning_network, target_network):
@@ -60,11 +61,7 @@ class OneStepQLearning:
self.optimizer.zero_grad()
loss.backward()
nn.utils.clip_grad_norm(self.worker_network.parameters(), config.gradient_clip)
for param, worker_param in zip(
self.learning_network.parameters(), self.worker_network.parameters()):
if param.grad is not None:
break
param._grad = worker_param.grad
sync_grad(self.learning_network, self.worker_network)
self.optimizer.step()
self.worker_network.load_state_dict(self.learning_network.state_dict())
self.worker_network.reset(terminal)
+2 -5
View File
@@ -7,6 +7,7 @@ import numpy as np
import torch
from torch.autograd import Variable
import torch.nn as nn
from utils import *
class OneStepSarsa:
def __init__(self, config, learning_network, target_network):
@@ -65,11 +66,7 @@ class OneStepSarsa:
self.optimizer.zero_grad()
loss.backward()
nn.utils.clip_grad_norm(self.worker_network.parameters(), config.gradient_clip)
for param, worker_param in zip(
self.learning_network.parameters(), self.worker_network.parameters()):
if param.grad is not None:
break
param._grad = worker_param.grad
sync_grad(self.learning_network, self.worker_network)
self.optimizer.step()
self.worker_network.load_state_dict(self.learning_network.state_dict())
self.worker_network.reset(terminal)
+1 -2
View File
@@ -152,8 +152,7 @@ class ProximalPolicyOptimization:
self.shared_network.zero_grad()
self.actor_opt.zero_grad()
self.critic_opt.zero_grad()
for param, worker_param in zip(self.shared_network.parameters(), self.worker_network.parameters()):
param._grad = worker_param.grad.clone()
sync_grad(self.shared_network, self.worker_network)
self.actor_opt.step()
self.critic_opt.step()