RAP / dataset_process /utils /spinnet /patchnet.py
YuePanEdward's picture
Squash history: release superseded example-data blobs
be88765
Raw History Blame Contribute Delete
8.06 kB
import torch.nn as nn
import torch
from .utils import common
class BaseNet(nn.Module):
def forward_one(self, x):
raise NotImplementedError()
def forward(self, imgs):
res = self.forward_one(imgs)
return res
class Cyclindrical_ConvNet(BaseNet):
def __init__(self, inchan=3, dilated=True, dilation=1, bn=True, bn_affine=False):
BaseNet.__init__(self)
self.inchan = inchan
self.curchan = inchan
self.dilated = dilated
self.dilation = dilation
self.bn = bn
self.bn_affine = bn_affine
self.ops = nn.ModuleList([])
def _make_bn_2d(self, outd):
return nn.BatchNorm2d(outd, affine=self.bn_affine)
def _make_bn_3d(self, outd):
return nn.BatchNorm3d(outd, affine=self.bn_affine)
def _add_conv_2d(self, outd, k=3, stride=1, dilation=1, bn=True, relu=True):
d = self.dilation * dilation
self.dilation *= stride
self.ops.append(nn.Conv2d(self.curchan, outd, kernel_size=(k, k), dilation=d))
if bn and self.bn: self.ops.append( self._make_bn_2d(outd) )
if relu: self.ops.append( nn.ReLU(inplace=True) )
self.curchan = outd
def _add_conv_3d(self, outd, k, stride=1, dilation=1, bn=True, relu=True):
d = self.dilation * dilation
self.dilation *= stride
self.ops.append(nn.Conv3d(self.curchan, outd, kernel_size=(k[0], k[1], k[2]), dilation=d))
if bn and self.bn: self.ops.append( self._make_bn_3d(outd) )
if relu: self.ops.append( nn.ReLU(inplace=True) )
self.curchan = outd
def forward_one(self, x):
assert self.ops, "You need to add convolutions first"
for n,op in enumerate(self.ops):
k_exist = hasattr(op, 'kernel_size')
if k_exist:
if len(op.kernel_size) == 3:
x = common.pad_image_3d(x, op.kernel_size[1] + (op.kernel_size[1]-1)*(op.dilation[0]-1))
else:
if len(x.shape) == 5:
x = x.squeeze(2)
mid_feat = x
x = common.pad_image(x, op.kernel_size[0] + (op.kernel_size[0]-1)*(op.dilation[0]-1))
x = op(x)
try:
mid_feat
except NameError:
return x
else:
return x, mid_feat
class Cylindrical_Net (Cyclindrical_ConvNet):
"""
Compute a 32D descriptor for cylindrical feature maps
"""
def __init__(self, inchan=16, dim=32, **kw ):
Cyclindrical_ConvNet.__init__(self, inchan=inchan, **kw)
add_conv_2d = lambda n, **kw: self._add_conv_2d(n, **kw)
add_conv_3d = lambda n, **kw: self._add_conv_3d(n, **kw)
add_conv_3d(64, k=[3, 3, 3])
add_conv_2d(64)
add_conv_2d(128)
add_conv_2d(128)
add_conv_2d(64)
add_conv_2d(64)
add_conv_2d(32)
add_conv_2d(32, bn=False, relu=False)
self.out_dim = dim
class Cylindrical_UNet(nn.Module):
"""
Compute a 32D descriptor for cylindrical feature maps with U-Net-like architecture
"""
def __init__(self, inchan=16, dim=32):
super(Cylindrical_UNet, self).__init__()
# Initial Conv3D Block
self.conv3d = nn.Sequential(
nn.Conv3d(inchan, 32, kernel_size=(3, 3, 3), stride=1, dilation=1),
nn.BatchNorm3d(32),
nn.ReLU(inplace=True),
)
# U-Net Encoder
self.encoder1 = self.make_conv_block(32, 32) # Encoder Level 1
self.encoder2 = self.make_conv_block(32, 64) # Encoder Level 2
self.encoder3 = self.make_conv_block(64, 128) # Encoder Level 3
# U-Net Bottleneck
self.bottleneck = self.make_conv_block(128, 128)
# U-Net Decoder
self.decoder3 = self.make_conv_block(128 + 128, 64) # Concat with Encoder Level 3
self.decoder2 = self.make_conv_block(64 + 64, 32) # Concat with Encoder Level 2
self.decoder1 = self.make_conv_block(32 + 32, 32) # Concat with Encoder Level 1
# Final Output Layer
self.output_layer = nn.Sequential(
nn.Conv2d(32, dim, kernel_size=3, stride=1, dilation=1),
nn.BatchNorm2d(dim),
nn.ReLU(inplace=True),
)
def make_conv_block(self, in_channels, out_channels, kernel_size=3, stride=1, dilation=1, bn=True, relu=True):
layers = [nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, dilation=dilation)]
if bn:
layers.append(nn.BatchNorm2d(out_channels))
if relu:
layers.append(nn.ReLU(inplace=True))
return nn.Sequential(*layers)
def forward(self, x):
# Conv3D Feature Extraction
x = self.conv3d(common.pad_image_3d(x, kernel_size=3))
x = x.squeeze(2) # Squeeze 3D output to 2D
# U-Net Encoder
enc1 = self.encoder1(common.pad_image(x, kernel_size=3)) # Level 1
enc2 = self.encoder2(common.pad_image(enc1, kernel_size=3)) # Level 2
enc3 = self.encoder3(common.pad_image(enc2, kernel_size=3)) # Level 3
# U-Net Bottleneck
bottleneck = self.bottleneck(common.pad_image(enc3, kernel_size=3))
# U-Net Decoder with Concatenation-based Skip Connections
dec3 = self.decoder3(common.pad_image(torch.cat([bottleneck, enc3], dim=1), kernel_size=3)) # Concat with Encoder Level 3
dec2 = self.decoder2(common.pad_image(torch.cat([dec3, enc2], dim=1), kernel_size=3)) # Concat with Encoder Level 2
dec1 = self.decoder1(common.pad_image(torch.cat([dec2, enc1], dim=1), kernel_size=3)) # Concat with Encoder Level 1
# Final Output
output = self.output_layer(common.pad_image(dec1, kernel_size=3))
return output, None
class CostBlock(BaseNet):
def __init__(self, inchan=32, dilated=True, dilation=1, bn=True, bn_affine=False):
BaseNet.__init__(self)
self.inchan = inchan
self.curchan = inchan
self.dilated = dilated
self.dilation = dilation
self.bn = bn
self.bn_affine = bn_affine
self.ops = nn.ModuleList([])
def _make_bn_2d(self, outd):
return nn.BatchNorm2d(outd, affine=self.bn_affine)
def _make_bn_3d(self, outd):
return nn.BatchNorm3d(outd, affine=self.bn_affine)
def _add_conv_2d(self, outd, k=3, stride=1, dilation=1, bn=True, relu=True):
d = self.dilation * dilation
self.dilation *= stride
self.ops.append(nn.Conv2d(self.curchan, outd, kernel_size=(k, k), dilation=d))
if bn and self.bn: self.ops.append( self._make_bn_2d(outd) )
if relu: self.ops.append( nn.ReLU(inplace=True) )
self.curchan = outd
def _add_conv_3d(self, outd, k, stride=1, dilation=1, bn=True, relu=True):
d = self.dilation * dilation
self.dilation *= stride
self.ops.append(nn.Conv3d(self.curchan, outd, kernel_size=(k[0], k[1], k[2]), dilation=d))
if bn and self.bn: self.ops.append( self._make_bn_3d(outd) )
if relu: self.ops.append( nn.ReLU(inplace=True) )
self.curchan = outd
def forward_one(self, x):
assert self.ops, "You need to add convolutions first"
for n,op in enumerate(self.ops):
x = op(x)
return x
class CostNet(CostBlock):
"""
Cost aggregation
"""
def __init__(self, inchan=32, dim=1, **kw ):
CostBlock.__init__(self, inchan=inchan, **kw)
add_conv_2d = lambda n, **kw: self._add_conv_2d(n, **kw)
add_conv_3d = lambda n, **kw: self._add_conv_3d(n, **kw)
add_conv_3d(32, k=[3, 3, 3])
add_conv_3d(64, k=[3, 3, 3])
add_conv_3d(64, k=[3, 1, 3])
add_conv_3d(128, k=[3, 1, 3])
add_conv_3d(128, k=[3, 1, 3])
add_conv_3d(64, k=[3, 1, 3])
add_conv_3d(64, k=[3, 1, 3])
add_conv_3d(32, k=[3, 1, 3])
add_conv_3d(32, k=[3, 1, 3])
add_conv_3d(dim, k=[2, 1, 2], bn=False, relu=False)
self.out_dim = dim