Minor update

This commit is contained in:
Shangtong Zhang
2018-02-02 15:19:55 -07:00
parent ad9189e25a
commit 049e56e976
5 changed files with 10 additions and 8 deletions
+1 -1
View File
@@ -97,7 +97,7 @@ class A2CAgent:
prob, log_prob, value, actions, returns, advantages = map(lambda x: torch.cat(x, dim=0), zip(*processed_rollout))
policy_loss = -log_prob.gather(1, Variable(actions)) * Variable(advantages)
policy_loss += config.entropy_weight * torch.sum(prob * log_prob, dim=1, keepdim=True)
value_loss = 0.5 * (Variable(returns) - value).pow(2)
value_loss = config.value_loss_weight * 0.5 * (Variable(returns) - value).pow(2)
self.optimizer.zero_grad()
(policy_loss + value_loss).sum().backward()
+1 -1
View File
@@ -19,7 +19,7 @@ def dqn_pixel_atari(name):
config.task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, normalized_state=False)
action_dim = config.task_fn().action_dim
config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.00025, alpha=0.95, eps=0.01)
config.network_fn = lambda optimizer_fn: NatureConvNet(config.history_length, action_dim, optimizer_fn)
config.network_fn = lambda: NatureConvNet(config.history_length, action_dim)
config.policy_fn = lambda: GreedyPolicy(epsilon=1.0, final_step=1000000, min_epsilon=0.1)
config.replay_fn = lambda: Replay(memory_size=1000000, batch_size=32, dtype=np.uint8)
config.discount = 0.99
+5 -4
View File
@@ -163,7 +163,7 @@ def a3c_pixel_atari(name):
def a2c_pixel_atari(name):
config = Config()
config.history_length = 1
config.num_workers = 16
config.num_workers = 8
task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, frame_size=42, max_steps=10000)
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers)
task = config.task_fn()
@@ -173,11 +173,12 @@ def a2c_pixel_atari(name):
config.reward_shift_fn = lambda r: np.sign(r)
config.policy_fn = SamplePolicy
config.discount = 0.99
config.gae_tau = 0.97
config.gae_tau = 1.0
config.entropy_weight = 0.01
config.rollout_length = 20
config.test_interval = 1000
config.test_repetitions = 10
config.value_loss_weight = 0.5
config.logger = Logger('./log', logger)
run_episodes(A2CAgent(config))
@@ -321,7 +322,7 @@ if __name__ == '__main__':
# dqn_cart_pole()
# async_cart_pole()
# a3c_cart_pole()
a2c_cart_pole()
# a2c_cart_pole()
# a3c_continuous()
# p3o_continuous()
# d3pg_continuous()
@@ -330,7 +331,7 @@ if __name__ == '__main__':
# dqn_pixel_atari('PongNoFrameskip-v4')
# async_pixel_atari('PongNoFrameskip-v4')
# a3c_pixel_atari('PongNoFrameskip-v4')
# a2c_pixel_atari('PongNoFrameskip-v4')
a2c_pixel_atari('PongNoFrameskip-v4')
# dqn_pixel_atari('BreakoutNoFrameskip-v4')
# async_pixel_atari('BreakoutNoFrameskip-v4')
+2 -2
View File
@@ -8,7 +8,7 @@ from .base_network import *
# Network for pixel Atari game with value based methods
class NatureConvNet(nn.Module, VanillaNet):
def __init__(self, in_channels, n_actions, optimizer_fn=None, gpu=True):
def __init__(self, in_channels, n_actions, gpu=True):
super(NatureConvNet, self).__init__()
self.conv1 = nn.Conv2d(in_channels, 32, kernel_size=8, stride=4)
self.conv2 = nn.Conv2d(32, 64, kernel_size=4, stride=2)
@@ -28,7 +28,7 @@ class NatureConvNet(nn.Module, VanillaNet):
# Network for pixel Atari game with dueling architecture
class DuelingNatureConvNet(nn.Module, DuelingNet):
def __init__(self, in_channels, n_actions, optimizer_fn=None, gpu=True):
def __init__(self, in_channels, n_actions, gpu=True):
super(DuelingNatureConvNet, self).__init__()
self.conv1 = nn.Conv2d(in_channels, 32, kernel_size=8, stride=4)
self.conv2 = nn.Conv2d(32, 64, kernel_size=4, stride=2)
+1
View File
@@ -52,3 +52,4 @@ class Config:
self.success_threshold = float('inf')
self.render_episode_freq = 0
self.rollout_length = None
self.value_loss_weight = 1.0