mirror of
https://github.com/wassname/DeepRL.git
synced 2026-10-07 11:13:51 +08:00
Specify a gpu for a network
This commit is contained in:
1 parent
fef04b6bf6
commit
e91ec3be45
12 files changed
+93
-103
No files matched your search
+21
-13
@@ -13,25 +13,33 @@ import numpy as np
|
||||
# Base class for all kinds of network
|
||||
class BasicNet:
|
||||
def __init__(self, gpu, LSTM=False):
|
||||
self.gpu = gpu and torch.cuda.is_available()
|
||||
if not torch.cuda.is_available():
|
||||
gpu = -1
|
||||
self.gpu = gpu
|
||||
self.LSTM = LSTM
|
||||
if self.gpu:
|
||||
self.cuda()
|
||||
self.FloatTensor = torch.cuda.FloatTensor
|
||||
self.LongTensor = torch.cuda.LongTensor
|
||||
else:
|
||||
self.FloatTensor = torch.FloatTensor
|
||||
self.LongTensor = torch.LongTensor
|
||||
if self.gpu >= 0:
|
||||
self.cuda(self.gpu)
|
||||
|
||||
def to_torch_variable(self, x, dtype='float32'):
|
||||
def supported_dtype(self, x, torch_type):
|
||||
if torch_type == torch.FloatTensor:
|
||||
return np.asarray(x, dtype=np.float32)
|
||||
if torch_type == torch.LongTensor:
|
||||
return np.asarray(x, dtype=np.int64)
|
||||
|
||||
def variable(self, x, dtype=torch.FloatTensor):
|
||||
if isinstance(x, Variable):
|
||||
return x
|
||||
if not isinstance(x, torch.FloatTensor):
|
||||
x = torch.from_numpy(np.asarray(x, dtype=dtype))
|
||||
if self.gpu:
|
||||
x = x.cuda()
|
||||
x = dtype(torch.from_numpy(self.supported_dtype(x, dtype)))
|
||||
if self.gpu >= 0:
|
||||
x = x.cuda(self.gpu)
|
||||
return Variable(x)
|
||||
|
||||
def tensor(self, x, dtype=torch.FloatTensor):
|
||||
x = dtype(torch.from_numpy(self.supported_dtype(x, dtype)))
|
||||
if self.gpu >= 0:
|
||||
x = x.cuda(self.gpu)
|
||||
return x
|
||||
|
||||
def reset(self, terminal):
|
||||
if not self.LSTM:
|
||||
return
|
||||
|
||||
@@ -12,7 +12,7 @@ class DeterministicActorNet(nn.Module, BasicNet):
|
||||
action_dim,
|
||||
action_gate,
|
||||
action_scale,
|
||||
gpu=False,
|
||||
gpu=-1,
|
||||
batch_norm=False,
|
||||
non_linear=F.relu,
|
||||
hidden_size=64):
|
||||
@@ -43,7 +43,7 @@ class DeterministicActorNet(nn.Module, BasicNet):
|
||||
nn.init.constant(self.layer2.bias.data, 0)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.to_torch_variable(x)
|
||||
x = self.variable(x)
|
||||
x = self.non_linear(self.layer1(x))
|
||||
if self.batch_norm:
|
||||
x = self.bn1(x)
|
||||
@@ -64,7 +64,7 @@ class DeterministicCriticNet(nn.Module, BasicNet):
|
||||
def __init__(self,
|
||||
state_dim,
|
||||
action_dim,
|
||||
gpu=False,
|
||||
gpu=-1,
|
||||
batch_norm=False,
|
||||
non_linear=F.relu,
|
||||
hidden_size=64):
|
||||
@@ -93,8 +93,8 @@ class DeterministicCriticNet(nn.Module, BasicNet):
|
||||
nn.init.constant(self.layer2.bias.data, 0)
|
||||
|
||||
def forward(self, x, action):
|
||||
x = self.to_torch_variable(x)
|
||||
action = self.to_torch_variable(action)
|
||||
x = self.variable(x)
|
||||
action = self.variable(action)
|
||||
x = self.non_linear(self.layer1(x))
|
||||
if self.batch_norm:
|
||||
x = self.bn1(x)
|
||||
@@ -113,7 +113,7 @@ class GaussianActorNet(nn.Module, BasicNet):
|
||||
action_dim,
|
||||
action_scale=1.0,
|
||||
action_gate=None,
|
||||
gpu=False,
|
||||
gpu=-1,
|
||||
unit_std=True,
|
||||
hidden_size=64):
|
||||
super(GaussianActorNet, self).__init__()
|
||||
@@ -133,7 +133,7 @@ class GaussianActorNet(nn.Module, BasicNet):
|
||||
BasicNet.__init__(self, gpu, False)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.to_torch_variable(x)
|
||||
x = self.variable(x)
|
||||
phi = F.tanh(self.fc1(x))
|
||||
phi = F.tanh(self.fc2(phi))
|
||||
mean = self.action_mean(phi)
|
||||
@@ -161,7 +161,7 @@ class GaussianActorNet(nn.Module, BasicNet):
|
||||
class GaussianCriticNet(nn.Module, BasicNet):
|
||||
def __init__(self,
|
||||
state_dim,
|
||||
gpu=False,
|
||||
gpu=-1,
|
||||
hidden_size=64):
|
||||
super(GaussianCriticNet, self).__init__()
|
||||
self.fc1 = nn.Linear(state_dim, hidden_size)
|
||||
@@ -170,7 +170,7 @@ class GaussianCriticNet(nn.Module, BasicNet):
|
||||
BasicNet.__init__(self, gpu, False)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.to_torch_variable(x)
|
||||
x = self.variable(x)
|
||||
phi = F.tanh(self.fc1(x))
|
||||
phi = F.tanh(self.fc2(phi))
|
||||
value = self.fc_value(phi)
|
||||
|
||||
+12
-12
@@ -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, gpu=True):
|
||||
def __init__(self, in_channels, n_actions, gpu=0):
|
||||
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)
|
||||
@@ -18,7 +18,7 @@ class NatureConvNet(nn.Module, VanillaNet):
|
||||
BasicNet.__init__(self, gpu)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.to_torch_variable(x)
|
||||
x = self.variable(x)
|
||||
y = F.relu(self.conv1(x))
|
||||
y = F.relu(self.conv2(y))
|
||||
y = F.relu(self.conv3(y))
|
||||
@@ -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, gpu=True):
|
||||
def __init__(self, in_channels, n_actions, gpu=0):
|
||||
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)
|
||||
@@ -39,7 +39,7 @@ class DuelingNatureConvNet(nn.Module, DuelingNet):
|
||||
BasicNet.__init__(self, gpu)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.to_torch_variable(x)
|
||||
x = self.variable(x)
|
||||
y = F.relu(self.conv1(x))
|
||||
y = F.relu(self.conv2(y))
|
||||
y = F.relu(self.conv3(y))
|
||||
@@ -52,7 +52,7 @@ class OpenAIActorCriticConvNet(nn.Module, ActorCriticNet):
|
||||
in_channels,
|
||||
n_actions,
|
||||
LSTM=False,
|
||||
gpu=False):
|
||||
gpu=-1):
|
||||
super(OpenAIActorCriticConvNet, self).__init__()
|
||||
self.conv1 = nn.Conv2d(in_channels, 32, 3, stride=2, padding=1)
|
||||
self.conv2 = nn.Conv2d(32, 32, 3, stride=2, padding=1)
|
||||
@@ -71,11 +71,11 @@ class OpenAIActorCriticConvNet(nn.Module, ActorCriticNet):
|
||||
self.fc_critic = nn.Linear(hidden_units, 1)
|
||||
BasicNet.__init__(self, gpu=gpu, LSTM=LSTM)
|
||||
if LSTM:
|
||||
self.h = self.to_torch_variable(np.zeros((1, hidden_units)))
|
||||
self.c = self.to_torch_variable(np.zeros((1, hidden_units)))
|
||||
self.h = self.variable(np.zeros((1, hidden_units)))
|
||||
self.c = self.variable(np.zeros((1, hidden_units)))
|
||||
|
||||
def forward(self, x, update_LSTM=True):
|
||||
x = self.to_torch_variable(x)
|
||||
x = self.variable(x)
|
||||
y = F.elu(self.conv1(x))
|
||||
y = F.elu(self.conv2(y))
|
||||
y = F.elu(self.conv3(y))
|
||||
@@ -95,7 +95,7 @@ class OpenAIConvNet(nn.Module, VanillaNet):
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
n_actions,
|
||||
gpu=False):
|
||||
gpu=0):
|
||||
super(OpenAIConvNet, self).__init__()
|
||||
self.conv1 = nn.Conv2d(in_channels, 32, 3, stride=2, padding=1)
|
||||
self.conv2 = nn.Conv2d(32, 32, 3, stride=2, padding=1)
|
||||
@@ -109,7 +109,7 @@ class OpenAIConvNet(nn.Module, VanillaNet):
|
||||
BasicNet.__init__(self, gpu=gpu, LSTM=False)
|
||||
|
||||
def forward(self, x, update_LSTM=True):
|
||||
x = self.to_torch_variable(x)
|
||||
x = self.variable(x)
|
||||
y = F.elu(self.conv1(x))
|
||||
y = F.elu(self.conv2(y))
|
||||
y = F.elu(self.conv3(y))
|
||||
@@ -122,7 +122,7 @@ class NatureActorCriticConvNet(nn.Module, ActorCriticNet):
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
n_actions,
|
||||
gpu=False):
|
||||
gpu=-1):
|
||||
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)
|
||||
@@ -134,7 +134,7 @@ class NatureActorCriticConvNet(nn.Module, ActorCriticNet):
|
||||
BasicNet.__init__(self, gpu=gpu)
|
||||
|
||||
def forward(self, x, _):
|
||||
x = self.to_torch_variable(x)
|
||||
x = self.variable(x)
|
||||
x = F.relu(self.conv1(x))
|
||||
x = F.relu(self.conv2(x))
|
||||
x = F.relu(self.conv3(x))
|
||||
|
||||
@@ -8,7 +8,7 @@ from .base_network import *
|
||||
|
||||
# Network for CartPole with value based methods
|
||||
class FCNet(nn.Module, VanillaNet):
|
||||
def __init__(self, dims, gpu=True):
|
||||
def __init__(self, dims, gpu=0):
|
||||
super(FCNet, self).__init__()
|
||||
self.fc1 = nn.Linear(dims[0], dims[1])
|
||||
self.fc2 = nn.Linear(dims[1], dims[2])
|
||||
@@ -16,8 +16,7 @@ class FCNet(nn.Module, VanillaNet):
|
||||
BasicNet.__init__(self, gpu)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.to_torch_variable(x)
|
||||
x = x.view(x.size(0), -1)
|
||||
x = self.variable(x)
|
||||
y = F.relu(self.fc1(x))
|
||||
y = F.relu(self.fc2(y))
|
||||
y = self.fc3(y)
|
||||
@@ -25,7 +24,7 @@ class FCNet(nn.Module, VanillaNet):
|
||||
|
||||
# Network for CartPole with dueling architecture
|
||||
class DuelingFCNet(nn.Module, DuelingNet):
|
||||
def __init__(self, dims, gpu=True):
|
||||
def __init__(self, dims, gpu=0):
|
||||
super(DuelingFCNet, self).__init__()
|
||||
self.fc1 = nn.Linear(dims[0], dims[1])
|
||||
self.fc2 = nn.Linear(dims[1], dims[2])
|
||||
@@ -34,8 +33,7 @@ class DuelingFCNet(nn.Module, DuelingNet):
|
||||
BasicNet.__init__(self, gpu)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.to_torch_variable(x)
|
||||
x = x.view(x.size(0), -1)
|
||||
x = self.variable(x)
|
||||
y = F.relu(self.fc1(x))
|
||||
phi = F.relu(self.fc2(y))
|
||||
return phi
|
||||
@@ -53,8 +51,7 @@ class ActorCriticFCNet(nn.Module, ActorCriticNet):
|
||||
BasicNet.__init__(self, False)
|
||||
|
||||
def forward(self, x, update_LSTM=True):
|
||||
x = self.to_torch_variable(x)
|
||||
x = x.view(x.size(0), -1)
|
||||
x = self.variable(x)
|
||||
x = F.relu(self.fc1(x))
|
||||
phi = self.fc2(x)
|
||||
return phi
|
||||
Reference in new issue
Block a user