Support shared stats for continuous A3C

This commit is contained in:
Shangtong Zhang
2017-10-04 20:58:42 -06:00
parent 3935dfc428
commit c773da7d08
7 changed files with 105 additions and 53 deletions
+12 -4
View File
@@ -25,10 +25,10 @@ def train(id, config, learning_network, target_network):
config.logger.debug('worker %d, episode %d, return %f, avg return %f, episode steps %d, total steps %d' % (
id, episode, rewards[-1], np.mean(rewards[-100:]), steps, config.total_steps.value))
def evaluate(config, task, learning_network):
def evaluate(config, task, learning_network, extra):
test_rewards = []
test_points = []
worker = config.worker(config, learning_network, None)
worker = config.worker(config, learning_network, extra)
# config.logger = Logger('./evaluation_log', gym.logger)
while True:
steps = config.total_steps.value
@@ -69,8 +69,16 @@ class AsyncAgent:
target_network.load_state_dict(learning_network.state_dict())
os.environ['OMP_NUM_THREADS'] = '1'
args = [(i, config, learning_network, target_network) for i in range(config.num_workers)]
args.append((config, task, learning_network))
if config.worker == NStepQLearning or config.worker == OneStepQLearning or config.worker == OneStepSarsa:
extra = target_network
elif config.worker == ContinuousAdvantageActorCritic:
state_normalizer = StaticNormalizer(task.state_dim)
reward_normalizer = StaticNormalizer(1)
extra = [state_normalizer, reward_normalizer]
else:
extra = None
args = [(i, config, learning_network, extra) for i in range(config.num_workers)]
args.append((config, task, learning_network, extra))
procs = [mp.Process(target=train, args=args[i]) for i in range(config.num_workers)]
procs.append(mp.Process(target=evaluate, args=args[-1]))
for p in procs: p.start()
+21 -6
View File
@@ -7,9 +7,10 @@ import numpy as np
import torch
from torch.autograd import Variable
import torch.nn as nn
from utils import *
class ContinuousAdvantageActorCritic:
def __init__(self, config, learning_network, target_network):
def __init__(self, config, learning_network, extra):
self.config = config
self.actor_opt = config.actor_optimizer_fn(learning_network.actor.parameters())
self.critic_opt = config.critic_optimizer_fn(learning_network.critic.parameters())
@@ -20,15 +21,21 @@ class ContinuousAdvantageActorCritic:
self.learning_network = learning_network
self.counter = 0
self.shared_state_normalizer = extra[0]
self.state_normalizer = StaticNormalizer(self.task.state_dim)
self.shared_reward_normalizer = extra[1]
self.reward_normalizer = StaticNormalizer(1)
def episode(self, deterministic=False):
config = self.config
self.state_normalizer.offline_stats.load(self.shared_state_normalizer.offline_stats)
self.reward_normalizer.offline_stats.load(self.shared_reward_normalizer.offline_stats)
state = self.task.reset()
state = config.state_shift_fn(state)
state = self.state_normalizer(state)
steps = 0
total_reward = 0
pending = []
while not config.stop_signal.value and \
(not config.max_episode_length or steps < config.max_episode_length):
while not config.stop_signal.value:
mean, std, log_std = self.worker_network.actor.predict(np.stack([state]))
value = self.worker_network.critic.predict(np.stack([state]))
action = self.policy.sample(mean.data.numpy().flatten(),
@@ -36,7 +43,9 @@ class ContinuousAdvantageActorCritic:
False)
action = self.config.action_shift_fn(action)
next_state, reward, terminal, _ = self.task.step(action)
next_state = config.state_shift_fn(next_state)
terminal = (terminal or (config.max_episode_length and steps >= config.max_episode_length))
next_state = self.state_normalizer(next_state)
# next_state = config.state_shift_fn(next_state)
# if deterministic:
# self.config.logger.scalar_summary('reward', reward, self.counter)
@@ -49,7 +58,7 @@ class ContinuousAdvantageActorCritic:
steps += 1
total_reward += reward
reward = config.reward_shift_fn(reward)
reward = np.asscalar(self.reward_normalizer(np.array([reward])))
if deterministic:
if terminal:
@@ -106,4 +115,10 @@ class ContinuousAdvantageActorCritic:
break
state = next_state
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
+8 -13
View File
@@ -66,13 +66,12 @@ def a3c_cart_pole():
def a3c_pendulum():
config = Config()
config.task_fn = lambda: Pendulum()
# config.reward_shift_fn = lambda reward: reward / 10
task = config.task_fn()
config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.0001)
config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
config.actor_network_fn = lambda: GaussianActorNet(task.state_dim, task.action_dim)
config.critic_network_fn = lambda: GaussianCriticNet(task.state_dim)
config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn)
config.network_fn = lambda: DisjointActorCriticNet(
lambda: GaussianActorNet(task.state_dim, task.action_dim),
lambda: GaussianCriticNet(task.state_dim))
config.policy_fn = lambda: GaussianPolicy()
config.worker = ContinuousAdvantageActorCritic
config.discount = 0.99
@@ -81,7 +80,6 @@ def a3c_pendulum():
config.update_interval = 5
config.test_interval = 1
config.test_repetitions = 5
# config.entropy_weight = 0.0001
config.entropy_weight = 0
config.gradient_clip = 40
config.logger = Logger('./log', gym.logger)
@@ -91,14 +89,12 @@ def a3c_pendulum():
def a3c_walker():
config = Config()
config.task_fn = lambda: BipedalWalker()
shifter = Shifter()
config.state_shift_fn = lambda state: shifter(state)
task = config.task_fn()
config.actor_optimizer_fn = lambda params: torch.optim.Adam(params, 0.0001)
config.critic_optimizer_fn = lambda params: torch.optim.Adam(params, 0.001)
config.actor_network_fn = lambda: GaussianActorNet(task.state_dim, task.action_dim)
config.critic_network_fn = lambda: GaussianCriticNet(task.state_dim)
config.network_fn = lambda: DisjointActorCriticNet(config.actor_network_fn, config.critic_network_fn)
config.network_fn = lambda: DisjointActorCriticNet(
lambda: GaussianActorNet(task.state_dim, task.action_dim),
lambda: GaussianCriticNet(task.state_dim))
config.policy_fn = lambda: GaussianPolicy()
config.worker = ContinuousAdvantageActorCritic
config.discount = 0.99
@@ -107,9 +103,8 @@ def a3c_walker():
config.update_interval = 20
config.test_interval = 1
config.test_repetitions = 5
# config.entropy_weight = 0.01
config.entropy_weight = 0
config.gradient_clip = 30
config.gradient_clip = 40
config.logger = Logger('./log', gym.logger)
agent = AsyncAgent(config)
agent.run()
@@ -343,7 +338,7 @@ if __name__ == '__main__':
# async_cart_pole()
# a3c_cart_pole()
# a3c_pendulum()
a3c_walker()
# a3c_walker()
# ddpg_pendulum()
# ddpg_walker()
# ppo_pendulum()
+1 -1
View File
@@ -51,7 +51,7 @@ class VanillaNet(BasicNet):
# Base class for actor critic method
class ActorCriticNet(BasicNet):
def predict(self, x, _):
def predict(self, x):
phi = self.forward(x, True)
pre_prob = self.fc_actor(phi)
prob = F.softmax(pre_prob)
+1 -1
View File
@@ -1,5 +1,5 @@
from config import *
from shifter import *
from normalizer import *
try:
from tf_logger import Logger
except:
+62
View File
@@ -0,0 +1,62 @@
import torch
class StaticNormalizer:
def __init__(self, o_size):
self.offline_stats = SharedStats(o_size)
self.online_stats = SharedStats(o_size)
def __call__(self, o_):
o = torch.FloatTensor(o_)
self.online_stats.feed(o)
std = (self.offline_stats.v + 1e-6) ** .5
o = (o - self.offline_stats.m) / std
return o.numpy().reshape(o_.shape)
class SharedStats:
def __init__(self, o_size):
self.m = torch.zeros(o_size)
self.v = torch.zeros(o_size)
self.n = torch.zeros(1)
self.m.share_memory_()
self.v.share_memory_()
self.n.share_memory_()
def feed(self, o):
n = self.n[0]
new_m = self.m * (n / (n + 1)) + o / (n + 1)
self.v.copy_(self.v * (n / (n + 1)) + (o - self.m) * (o - new_m) / (n + 1))
self.m.copy_(new_m)
self.n.add_(1)
def zero(self):
self.m.zero_()
self.v.zero_()
self.n.zero_()
def load(self, stats):
self.m.copy_(stats.m)
self.v.copy_(stats.v)
self.n.copy_(stats.n)
def merge(self, B):
A = self
n_A = self.n[0]
n_B = B.n[0]
n = n_A + n_B
delta = B.m - A.m
m = A.m + delta * n_B / n
v = A.v * n_A + B.v * n_B + delta * delta * n_A * n_B / n
v /= n
self.m.copy_(m)
self.v.copy_(v)
self.n.add_(B.n)
def state_dict(self):
return {'m': self.m.numpy(),
'v': self.v.numpy(),
'n': self.n.numpy()}
def load_state_dict(self, saved):
self.m = torch.FloatTensor(saved['m'])
self.v = torch.FloatTensor(saved['v'])
self.n = torch.FloatTensor(saved['n'])
-28
View File
@@ -1,28 +0,0 @@
# Adapted from https://github.com/kvfrans/parallel-trpo/blob/master/utils.py
class Shifter:
def __init__(self, filter_mean=True):
self.m = 0
self.v = 0
self.n = 0.
self.filter_mean = filter_mean
def state_dict(self):
return {'m': self.m,
'v': self.v,
'n': self.n}
def load_state_dict(self, saved):
self.m = saved['m']
self.v = saved['v']
self.n = saved['n']
def __call__(self, o):
self.m = self.m * (self.n / (self.n + 1)) + o * 1 / (1 + self.n)
self.v = self.v * (self.n / (self.n + 1)) + (o - self.m) ** 2 * 1 / (1 + self.n)
self.std = (self.v + 1e-6) ** .5 # std
self.n += 1
if self.filter_mean:
o_ = (o - self.m) / self.std
else:
o_ = o / self.std
return o_