Added manual control over cuda usage

This commit is contained in:
bakirillov committed 2019-10-18 14:02:21 +03:00
1 parent 482e550797
commit 8e28f9704e
1 file changed
+6 -6
+6 -6
View File
@@ -114,11 +114,11 @@ class AddCoords(nn.Module):
class CoordConv1d(conv.Conv1d):
def __init__(self, in_channels, out_channels, kernel_size, stride=1,
padding=0, dilation=1, groups=1, bias=True, with_r=False):
padding=0, dilation=1, groups=1, bias=True, with_r=False, use_cuda=True):
super(CoordConv1d, self).__init__(in_channels, out_channels, kernel_size,
stride, padding, dilation, groups, bias)
self.rank = 1
self.addcoords = AddCoords(self.rank, with_r)
self.addcoords = AddCoords(self.rank, with_r, use_cuda=use_cuda)
self.conv = nn.Conv1d(in_channels + self.rank + int(with_r), out_channels,
kernel_size, stride, padding, dilation, groups, bias)
@@ -136,11 +136,11 @@ class CoordConv1d(conv.Conv1d):
class CoordConv2d(conv.Conv2d):
def __init__(self, in_channels, out_channels, kernel_size, stride=1,
padding=0, dilation=1, groups=1, bias=True, with_r=False):
padding=0, dilation=1, groups=1, bias=True, with_r=False, use_cuda=True):
super(CoordConv2d, self).__init__(in_channels, out_channels, kernel_size,
stride, padding, dilation, groups, bias)
self.rank = 2
self.addcoords = AddCoords(self.rank, with_r)
self.addcoords = AddCoords(self.rank, with_r, use_cuda=use_cuda)
self.conv = nn.Conv2d(in_channels + self.rank + int(with_r), out_channels,
kernel_size, stride, padding, dilation, groups, bias)
@@ -158,11 +158,11 @@ class CoordConv2d(conv.Conv2d):
class CoordConv3d(conv.Conv3d):
def __init__(self, in_channels, out_channels, kernel_size, stride=1,
padding=0, dilation=1, groups=1, bias=True, with_r=False):
padding=0, dilation=1, groups=1, bias=True, with_r=False, use_cuda=True):
super(CoordConv3d, self).__init__(in_channels, out_channels, kernel_size,
stride, padding, dilation, groups, bias)
self.rank = 3
self.addcoords = AddCoords(self.rank, with_r)
self.addcoords = AddCoords(self.rank, with_r, use_cuda=use_cuda)
self.conv = nn.Conv3d(in_channels + self.rank + int(with_r), out_channels,
kernel_size, stride, padding, dilation, groups, bias)