mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
Deepmind actor-critic net
This commit is contained in:
@@ -165,14 +165,15 @@ def a2c_pixel_atari(name):
|
|||||||
config = Config()
|
config = Config()
|
||||||
config.history_length = 4
|
config.history_length = 4
|
||||||
config.num_workers = 5
|
config.num_workers = 5
|
||||||
task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, frame_size=42,
|
task_fn = lambda: PixelAtari(name, no_op=30, frame_skip=4, frame_size=84,
|
||||||
history_length=config.history_length)
|
history_length=config.history_length)
|
||||||
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers)
|
config.task_fn = lambda: ParallelizedTask(task_fn, config.num_workers)
|
||||||
task = config.task_fn()
|
task = config.task_fn()
|
||||||
config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.0007, eps=1e-5, alpha=0.99)
|
config.optimizer_fn = lambda params: torch.optim.RMSprop(params, lr=0.0007)
|
||||||
# config.optimizer_fn = lambda params: torch.optim.Adam(params, lr=0.0001)
|
# config.optimizer_fn = lambda params: torch.optim.Adam(params, lr=0.0001)
|
||||||
config.network_fn = lambda: OpenAIActorCriticConvNet(
|
# config.network_fn = lambda: OpenAIActorCriticConvNet(
|
||||||
config.history_length, task.task.env.action_space.n, LSTM=False, gpu=True)
|
config.network_fn = lambda: NatureActorCriticConvNet(
|
||||||
|
config.history_length, task.task.env.action_space.n, gpu=True)
|
||||||
config.reward_shift_fn = lambda r: np.sign(r)
|
config.reward_shift_fn = lambda r: np.sign(r)
|
||||||
config.policy_fn = SamplePolicy
|
config.policy_fn = SamplePolicy
|
||||||
config.discount = 0.99
|
config.discount = 0.99
|
||||||
|
|||||||
+25
-1
@@ -116,4 +116,28 @@ class OpenAIConvNet(nn.Module, VanillaNet):
|
|||||||
y = F.elu(self.conv4(y))
|
y = F.elu(self.conv4(y))
|
||||||
y = y.view(y.size(0), -1)
|
y = y.view(y.size(0), -1)
|
||||||
phi = F.elu(self.layer5(y))
|
phi = F.elu(self.layer5(y))
|
||||||
return self.fc6(phi)
|
return self.fc6(phi)
|
||||||
|
|
||||||
|
class NatureActorCriticConvNet(nn.Module, ActorCriticNet):
|
||||||
|
def __init__(self,
|
||||||
|
in_channels,
|
||||||
|
n_actions,
|
||||||
|
gpu=False):
|
||||||
|
super(NatureActorCriticConvNet, 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)
|
||||||
|
self.conv3 = nn.Conv2d(64, 64, kernel_size=3, stride=1)
|
||||||
|
self.fc4 = nn.Linear(7 * 7 * 64, 512)
|
||||||
|
|
||||||
|
self.fc_actor = nn.Linear(512, n_actions)
|
||||||
|
self.fc_critic = nn.Linear(512, 1)
|
||||||
|
BasicNet.__init__(self, gpu=gpu)
|
||||||
|
|
||||||
|
def forward(self, x, _):
|
||||||
|
x = self.to_torch_variable(x)
|
||||||
|
x = F.relu(self.conv1(x))
|
||||||
|
x = F.relu(self.conv2(x))
|
||||||
|
x = F.relu(self.conv3(x))
|
||||||
|
x = x.view(x.size(0), -1)
|
||||||
|
phi = F.relu(self.fc4(x))
|
||||||
|
return phi
|
||||||
|
|||||||
Reference in New Issue
Block a user