mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
Minor update
This commit is contained in:
+1
-1
@@ -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
@@ -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
|
||||
|
||||
@@ -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')
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user