mirror of
https://github.com/wassname/DeepRL.git
synced 2026-10-08 11:39:05 +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
|
||||
|
||||
Reference in new issue
Block a user