mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
Support shared stats for continuous A3C
This commit is contained in:
+12
-4
@@ -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()
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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
@@ -1,5 +1,5 @@
|
||||
from config import *
|
||||
from shifter import *
|
||||
from normalizer import *
|
||||
try:
|
||||
from tf_logger import Logger
|
||||
except:
|
||||
|
||||
@@ -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'])
|
||||
@@ -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_
|
||||
Reference in New Issue
Block a user