mirror of
https://github.com/wassname/CoordConv.git
synced 2026-10-04 12:10:39 +08:00
Added manual control over cuda usage
This commit is contained in:
1 parent
482e550797
commit
8e28f9704e
1 file changed
+6
-6
+6
-6
@@ -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)
|
||||
|
||||
|
||||
Reference in new issue
Block a user