isaac commited on
Commit ·
af2b273
0
Parent(s):
RADAR ZeroGPU Space
Browse files- .gitattributes +1 -0
- .gitignore +17 -0
- RADAR_inference/dynamic_network_architectures/__init__.py +0 -0
- RADAR_inference/dynamic_network_architectures/architectures/__init__.py +0 -0
- RADAR_inference/dynamic_network_architectures/architectures/resnet.py +236 -0
- RADAR_inference/dynamic_network_architectures/architectures/resnet_vl.py +225 -0
- RADAR_inference/dynamic_network_architectures/architectures/unet.py +218 -0
- RADAR_inference/dynamic_network_architectures/architectures/unet_lightdecoder.py +218 -0
- RADAR_inference/dynamic_network_architectures/architectures/vgg.py +85 -0
- RADAR_inference/dynamic_network_architectures/building_blocks/__init__.py +0 -0
- RADAR_inference/dynamic_network_architectures/building_blocks/helper.py +242 -0
- RADAR_inference/dynamic_network_architectures/building_blocks/plain_conv_encoder.py +105 -0
- RADAR_inference/dynamic_network_architectures/building_blocks/regularization.py +86 -0
- RADAR_inference/dynamic_network_architectures/building_blocks/residual.py +371 -0
- RADAR_inference/dynamic_network_architectures/building_blocks/residual_encoders.py +172 -0
- RADAR_inference/dynamic_network_architectures/building_blocks/simple_conv_blocks.py +167 -0
- RADAR_inference/dynamic_network_architectures/building_blocks/unet_decoder.py +154 -0
- RADAR_inference/dynamic_network_architectures/building_blocks/unet_decoder_light.py +154 -0
- RADAR_inference/dynamic_network_architectures/building_blocks/unet_residual_decoder.py +155 -0
- RADAR_inference/dynamic_network_architectures/initialization/__init__.py +0 -0
- RADAR_inference/dynamic_network_architectures/initialization/weight_init.py +34 -0
- RADAR_inference/dynamic_network_architectures/med.py +1502 -0
- RADAR_inference/dynamic_network_architectures/vision_branch.py +160 -0
- RADAR_inference/inference_demo.py +630 -0
- README.md +21 -0
- app.py +133 -0
- ckpt/infer_text_embedding_radar.pt +3 -0
- requirements.txt +45 -0
.gitattributes
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
*.pt filter=lfs diff=lfs merge=lfs -text
|
.gitignore
ADDED
|
@@ -0,0 +1,17 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Model weights and volumes are fetched at runtime, never committed.
|
| 2 |
+
*.pth
|
| 3 |
+
*.ptc
|
| 4 |
+
*.nii
|
| 5 |
+
*.nii.gz
|
| 6 |
+
*.hdr
|
| 7 |
+
*.img
|
| 8 |
+
RADAR_infer_results_*.csv
|
| 9 |
+
# Upstream scaffold, removed after vendoring.
|
| 10 |
+
radar-upstream/
|
| 11 |
+
damo-radar.zip
|
| 12 |
+
# Python noise.
|
| 13 |
+
__pycache__/
|
| 14 |
+
*.pyc
|
| 15 |
+
.venv/
|
| 16 |
+
# Local knowledge-graph tooling; never pushed to the Space.
|
| 17 |
+
graphify-out/
|
RADAR_inference/dynamic_network_architectures/__init__.py
ADDED
|
File without changes
|
RADAR_inference/dynamic_network_architectures/architectures/__init__.py
ADDED
|
File without changes
|
RADAR_inference/dynamic_network_architectures/architectures/resnet.py
ADDED
|
@@ -0,0 +1,236 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from dynamic_network_architectures.building_blocks.residual_encoders import ResidualEncoder, BottleneckD, BasicBlockD
|
| 3 |
+
from dynamic_network_architectures.building_blocks.helper import get_matching_pool_op, get_default_network_config
|
| 4 |
+
from dynamic_network_architectures.building_blocks.simple_conv_blocks import ConvDropoutNormReLU
|
| 5 |
+
from torch import nn
|
| 6 |
+
|
| 7 |
+
_ResNet_CONFIGS = {
|
| 8 |
+
'18': {'features_per_stage': (64, 128, 256, 512), 'n_blocks_per_stage': (2, 2, 2, 2), 'strides': (1, 2, 2, 2),
|
| 9 |
+
'block': BasicBlockD, 'bottleneck_channels': None, 'disable_default_stem': True, 'stem_channels': None},
|
| 10 |
+
'34': {'features_per_stage': (64, 128, 256, 512), 'n_blocks_per_stage': (3, 4, 6, 3), 'strides': (1, 2, 2, 2),
|
| 11 |
+
'block': BasicBlockD, 'bottleneck_channels': None, 'disable_default_stem': True, 'stem_channels': None},
|
| 12 |
+
'50': {'features_per_stage': (64, 128, 256, 512), 'n_blocks_per_stage': (4, 6, 10, 5), 'strides': (1, 2, 2, 2),
|
| 13 |
+
'block': BasicBlockD, 'bottleneck_channels': None, 'disable_default_stem': True, 'stem_channels': None},
|
| 14 |
+
'152': {'features_per_stage': (64, 128, 256, 512), 'n_blocks_per_stage': (4, 13, 55, 4), 'strides': (1, 2, 2, 2),
|
| 15 |
+
'block': BasicBlockD, 'bottleneck_channels': None, 'disable_default_stem': True, 'stem_channels': None},
|
| 16 |
+
'50_bn': {'features_per_stage': (256, 512, 1024, 2048), 'n_blocks_per_stage': (3, 4, 6, 3), 'strides': (1, 2, 2, 2),
|
| 17 |
+
'block': BottleneckD, 'bottleneck_channels': (64, 128, 256, 512), 'disable_default_stem': True,
|
| 18 |
+
'stem_channels': 64},
|
| 19 |
+
'152_bn': {'features_per_stage': (256, 512, 1024, 2048), 'n_blocks_per_stage': (3, 8, 36, 3),
|
| 20 |
+
'strides': (1, 2, 2, 2),
|
| 21 |
+
'block': BottleneckD, 'bottleneck_channels': (64, 128, 256, 512), 'disable_default_stem': True,
|
| 22 |
+
'stem_channels': 64},
|
| 23 |
+
'18_cifar': {'features_per_stage': (64, 128, 256, 512), 'n_blocks_per_stage': (2, 2, 2, 2), 'strides': (1, 2, 2, 2),
|
| 24 |
+
'block': BasicBlockD, 'bottleneck_channels': None, 'disable_default_stem': False,
|
| 25 |
+
'stem_channels': None},
|
| 26 |
+
'34_cifar': {'features_per_stage': (64, 128, 256, 512), 'n_blocks_per_stage': (3, 4, 6, 3), 'strides': (1, 2, 2, 2),
|
| 27 |
+
'block': BasicBlockD, 'bottleneck_channels': None, 'disable_default_stem': False,
|
| 28 |
+
'stem_channels': None},
|
| 29 |
+
'50_cifar': {'features_per_stage': (64, 128, 256, 512), 'n_blocks_per_stage': (4, 6, 10, 5),
|
| 30 |
+
'strides': (1, 2, 2, 2),
|
| 31 |
+
'block': BasicBlockD, 'bottleneck_channels': None, 'disable_default_stem': False,
|
| 32 |
+
'stem_channels': None},
|
| 33 |
+
'152_cifar': {'features_per_stage': (64, 128, 256, 512), 'n_blocks_per_stage': (4, 13, 55, 4),
|
| 34 |
+
'strides': (1, 2, 2, 2),
|
| 35 |
+
'block': BasicBlockD, 'bottleneck_channels': None, 'disable_default_stem': False,
|
| 36 |
+
'stem_channels': None},
|
| 37 |
+
'50_cifar_bn': {'features_per_stage': (256, 512, 1024, 2048), 'n_blocks_per_stage': (3, 4, 6, 3),
|
| 38 |
+
'strides': (1, 2, 2, 2),
|
| 39 |
+
'block': BottleneckD, 'bottleneck_channels': (64, 128, 256, 512), 'disable_default_stem': False,
|
| 40 |
+
'stem_channels': 64},
|
| 41 |
+
'152_cifar_bn': {'features_per_stage': (256, 512, 1024, 2048), 'n_blocks_per_stage': (3, 8, 36, 3),
|
| 42 |
+
'strides': (1, 2, 2, 2),
|
| 43 |
+
'block': BottleneckD, 'bottleneck_channels': (64, 128, 256, 512), 'disable_default_stem': False,
|
| 44 |
+
'stem_channels': 64},
|
| 45 |
+
}
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
class ResNetD(nn.Module):
|
| 49 |
+
def __init__(self, n_classes: int, n_input_channel: int = 3, config='18', input_dimension=2,
|
| 50 |
+
final_layer_dropout=0.0, stochastic_depth_p=0.0, squeeze_excitation=False,
|
| 51 |
+
squeeze_excitation_rd_ratio=1./16):
|
| 52 |
+
"""
|
| 53 |
+
Implements ResNetD (https://arxiv.org/pdf/1812.01187.pdf).
|
| 54 |
+
Args:
|
| 55 |
+
n_classes: Number of classes
|
| 56 |
+
n_input_channel: Number of input channels (e.g. 3 for RGB)
|
| 57 |
+
config: Configuration of the ResNet
|
| 58 |
+
input_dimension: Number of dimensions of the data (1, 2 or 3)
|
| 59 |
+
final_layer_dropout: Probability of dropout before the final classifier
|
| 60 |
+
stochastic_depth_p: Stochastic Depth probability
|
| 61 |
+
squeeze_excitation: Whether Squeeze and Excitation should be applied
|
| 62 |
+
squeeze_excitation_rd_ratio: Squeeze and Excitation Reduction Ratio
|
| 63 |
+
Returns:
|
| 64 |
+
ResNet Model
|
| 65 |
+
"""
|
| 66 |
+
super().__init__()
|
| 67 |
+
self.input_channels = n_input_channel
|
| 68 |
+
self.cfg = _ResNet_CONFIGS[config]
|
| 69 |
+
self.ops = get_default_network_config(dimension=input_dimension)
|
| 70 |
+
self.final_layer_dropout_p = final_layer_dropout
|
| 71 |
+
|
| 72 |
+
if self.cfg['disable_default_stem']:
|
| 73 |
+
stem_features = self.cfg['stem_channels'] if self.cfg['stem_channels'] is not None else \
|
| 74 |
+
self.cfg['features_per_stage'][0]
|
| 75 |
+
self.stem = self._build_imagenet_stem_D(stem_features)
|
| 76 |
+
encoder_input_features = stem_features
|
| 77 |
+
else:
|
| 78 |
+
encoder_input_features = n_input_channel
|
| 79 |
+
self.stem = None
|
| 80 |
+
|
| 81 |
+
self.encoder = ResidualEncoder(encoder_input_features, n_stages=len(self.cfg['features_per_stage']),
|
| 82 |
+
features_per_stage=self.cfg['features_per_stage'], conv_op=self.ops['conv_op'],
|
| 83 |
+
kernel_sizes=3, strides=self.cfg['strides'],
|
| 84 |
+
n_blocks_per_stage=self.cfg['n_blocks_per_stage'], conv_bias=False,
|
| 85 |
+
norm_op=self.ops['norm_op'], norm_op_kwargs=None, dropout_op=None,
|
| 86 |
+
dropout_op_kwargs=None, nonlin=nn.ReLU,
|
| 87 |
+
nonlin_kwargs={'inplace': True}, block=self.cfg['block'],
|
| 88 |
+
bottleneck_channels=self.cfg['bottleneck_channels'], return_skips=False,
|
| 89 |
+
disable_default_stem=self.cfg['disable_default_stem'],
|
| 90 |
+
stem_channels=self.cfg['stem_channels'],
|
| 91 |
+
stochastic_depth_p=stochastic_depth_p,
|
| 92 |
+
squeeze_excitation=squeeze_excitation,
|
| 93 |
+
squeeze_excitation_reduction_ratio=squeeze_excitation_rd_ratio)
|
| 94 |
+
|
| 95 |
+
self.gap = get_matching_pool_op(conv_op=self.ops['conv_op'], adaptive=True, pool_type='avg')(1)
|
| 96 |
+
self.classifier = nn.Linear(self.cfg['features_per_stage'][-1], n_classes, True)
|
| 97 |
+
self.final_layer_dropout = self.ops['dropout_op'](p=self.final_layer_dropout_p)
|
| 98 |
+
|
| 99 |
+
def forward(self, x):
|
| 100 |
+
if self.stem is not None:
|
| 101 |
+
x = self.stem(x)
|
| 102 |
+
x = self.encoder(x)
|
| 103 |
+
x = self.gap(x)
|
| 104 |
+
x = self.final_layer_dropout(x).squeeze()
|
| 105 |
+
|
| 106 |
+
return self.classifier(x)
|
| 107 |
+
|
| 108 |
+
def _build_imagenet_stem_D(self, stem_features):
|
| 109 |
+
"""
|
| 110 |
+
https://arxiv.org/pdf/1812.01187.pdf
|
| 111 |
+
|
| 112 |
+
use 3 3x3(x3) convs instead of one 7x7. Stride is located in first conv.
|
| 113 |
+
|
| 114 |
+
Fig2 b) describes this
|
| 115 |
+
:return:
|
| 116 |
+
"""
|
| 117 |
+
c1 = ConvDropoutNormReLU(self.ops['conv_op'], self.input_channels, stem_features, 3, 2, False,
|
| 118 |
+
self.ops['norm_op'], None, None, None, nn.ReLU, {'inplace': True})
|
| 119 |
+
c2 = ConvDropoutNormReLU(self.ops['conv_op'], stem_features, stem_features, 3, 1, False,
|
| 120 |
+
self.ops['norm_op'], None, None, None, nn.ReLU, {'inplace': True})
|
| 121 |
+
c3 = ConvDropoutNormReLU(self.ops['conv_op'], stem_features, stem_features, 3, 1, False,
|
| 122 |
+
self.ops['norm_op'], None, None, None, nn.ReLU, {'inplace': True})
|
| 123 |
+
pl = get_matching_pool_op(conv_op=self.ops['conv_op'], adaptive=False, pool_type='max')(2)
|
| 124 |
+
stem = nn.Sequential(c1, c2, c3, pl)
|
| 125 |
+
return stem
|
| 126 |
+
|
| 127 |
+
|
| 128 |
+
class ResNet18_CIFAR(ResNetD):
|
| 129 |
+
def __init__(self, n_classes: int, n_input_channels: int = 3, input_dimension: int = 2,
|
| 130 |
+
final_layer_dropout: float = 0.0, stochastic_depth_p: float = 0.0, squeeze_excitation: bool = False,
|
| 131 |
+
squeeze_excitation_rd_ratio: float = 1./16):
|
| 132 |
+
super().__init__(n_classes, n_input_channels, config='18_cifar', input_dimension=input_dimension,
|
| 133 |
+
final_layer_dropout=final_layer_dropout, stochastic_depth_p=stochastic_depth_p,
|
| 134 |
+
squeeze_excitation=squeeze_excitation, squeeze_excitation_rd_ratio=squeeze_excitation_rd_ratio)
|
| 135 |
+
|
| 136 |
+
class ResNet34_CIFAR(ResNetD):
|
| 137 |
+
def __init__(self, n_classes: int, n_input_channels: int = 3, input_dimension: int = 2,
|
| 138 |
+
final_layer_dropout: float = 0.0, stochastic_depth_p: float = 0.0, squeeze_excitation: bool = False,
|
| 139 |
+
squeeze_excitation_rd_ratio: float = 1./16):
|
| 140 |
+
super().__init__(n_classes, n_input_channels, config='34_cifar', input_dimension=input_dimension,
|
| 141 |
+
final_layer_dropout=final_layer_dropout, stochastic_depth_p=stochastic_depth_p,
|
| 142 |
+
squeeze_excitation=squeeze_excitation, squeeze_excitation_rd_ratio=squeeze_excitation_rd_ratio)
|
| 143 |
+
|
| 144 |
+
class ResNet50_CIFAR(ResNetD):
|
| 145 |
+
def __init__(self, n_classes: int, n_input_channels: int = 3, input_dimension: int = 2,
|
| 146 |
+
final_layer_dropout: float = 0.0, stochastic_depth_p: float = 0.0, squeeze_excitation: bool = False,
|
| 147 |
+
squeeze_excitation_rd_ratio: float = 1./16):
|
| 148 |
+
super().__init__(n_classes, n_input_channels, config='50_cifar', input_dimension=input_dimension,
|
| 149 |
+
final_layer_dropout=final_layer_dropout, stochastic_depth_p=stochastic_depth_p,
|
| 150 |
+
squeeze_excitation=squeeze_excitation, squeeze_excitation_rd_ratio=squeeze_excitation_rd_ratio)
|
| 151 |
+
|
| 152 |
+
class ResNet152_CIFAR(ResNetD):
|
| 153 |
+
def __init__(self, n_classes: int, n_input_channels: int = 3, input_dimension: int = 2,
|
| 154 |
+
final_layer_dropout: float = 0.0, stochastic_depth_p: float = 0.0, squeeze_excitation: bool = False,
|
| 155 |
+
squeeze_excitation_rd_ratio: float = 1./16):
|
| 156 |
+
super().__init__(n_classes, n_input_channels, config='152_cifar', input_dimension=input_dimension,
|
| 157 |
+
final_layer_dropout=final_layer_dropout, stochastic_depth_p=stochastic_depth_p,
|
| 158 |
+
squeeze_excitation=squeeze_excitation, squeeze_excitation_rd_ratio=squeeze_excitation_rd_ratio)
|
| 159 |
+
|
| 160 |
+
class ResNet50bn_CIFAR(ResNetD):
|
| 161 |
+
def __init__(self, n_classes: int, n_input_channels: int = 3, input_dimension: int = 2,
|
| 162 |
+
final_layer_dropout: float = 0.0, stochastic_depth_p: float = 0.0, squeeze_excitation: bool = False,
|
| 163 |
+
squeeze_excitation_rd_ratio: float = 1./16):
|
| 164 |
+
super().__init__(n_classes, n_input_channels, config='50_cifar_bn', input_dimension=input_dimension,
|
| 165 |
+
final_layer_dropout=final_layer_dropout, stochastic_depth_p=stochastic_depth_p,
|
| 166 |
+
squeeze_excitation=squeeze_excitation, squeeze_excitation_rd_ratio=squeeze_excitation_rd_ratio)
|
| 167 |
+
|
| 168 |
+
class ResNet152bn_CIFAR(ResNetD):
|
| 169 |
+
def __init__(self, n_classes: int, n_input_channels: int = 3, input_dimension: int = 2,
|
| 170 |
+
final_layer_dropout: float = 0.0, stochastic_depth_p: float = 0.0, squeeze_excitation: bool = False,
|
| 171 |
+
squeeze_excitation_rd_ratio: float = 1./16):
|
| 172 |
+
super().__init__(n_classes, n_input_channels, config='152_cifar_bn', input_dimension=input_dimension,
|
| 173 |
+
final_layer_dropout=final_layer_dropout, stochastic_depth_p=stochastic_depth_p,
|
| 174 |
+
squeeze_excitation=squeeze_excitation, squeeze_excitation_rd_ratio=squeeze_excitation_rd_ratio)
|
| 175 |
+
|
| 176 |
+
class ResNet18(ResNetD):
|
| 177 |
+
def __init__(self, n_classes: int, n_input_channels: int = 3, input_dimension: int = 2,
|
| 178 |
+
final_layer_dropout: float = 0.0, stochastic_depth_p: float = 0.0, squeeze_excitation: bool = False,
|
| 179 |
+
squeeze_excitation_rd_ratio: float = 1./16):
|
| 180 |
+
super().__init__(n_classes, n_input_channels, config='18', input_dimension=input_dimension,
|
| 181 |
+
final_layer_dropout=final_layer_dropout, stochastic_depth_p=stochastic_depth_p,
|
| 182 |
+
squeeze_excitation=squeeze_excitation, squeeze_excitation_rd_ratio=squeeze_excitation_rd_ratio)
|
| 183 |
+
|
| 184 |
+
class ResNet34(ResNetD):
|
| 185 |
+
def __init__(self, n_classes: int, n_input_channels: int = 3, input_dimension: int = 2,
|
| 186 |
+
final_layer_dropout: float = 0.0, stochastic_depth_p: float = 0.0, squeeze_excitation: bool = False,
|
| 187 |
+
squeeze_excitation_rd_ratio: float = 1./16):
|
| 188 |
+
super().__init__(n_classes, n_input_channels, config='34', input_dimension=input_dimension,
|
| 189 |
+
final_layer_dropout=final_layer_dropout, stochastic_depth_p=stochastic_depth_p,
|
| 190 |
+
squeeze_excitation=squeeze_excitation, squeeze_excitation_rd_ratio=squeeze_excitation_rd_ratio)
|
| 191 |
+
|
| 192 |
+
class ResNet50(ResNetD):
|
| 193 |
+
def __init__(self, n_classes: int, n_input_channels: int = 3, input_dimension: int = 2,
|
| 194 |
+
final_layer_dropout: float = 0.0, stochastic_depth_p: float = 0.0, squeeze_excitation: bool = False,
|
| 195 |
+
squeeze_excitation_rd_ratio: float = 1./16):
|
| 196 |
+
super().__init__(n_classes, n_input_channels, config='50', input_dimension=input_dimension,
|
| 197 |
+
final_layer_dropout=final_layer_dropout, stochastic_depth_p=stochastic_depth_p,
|
| 198 |
+
squeeze_excitation=squeeze_excitation, squeeze_excitation_rd_ratio=squeeze_excitation_rd_ratio)
|
| 199 |
+
|
| 200 |
+
class ResNet152(ResNetD):
|
| 201 |
+
def __init__(self, n_classes: int, n_input_channels: int = 3, input_dimension: int = 2,
|
| 202 |
+
final_layer_dropout: float = 0.0, stochastic_depth_p: float = 0.0, squeeze_excitation: bool = False,
|
| 203 |
+
squeeze_excitation_rd_ratio: float = 1./16):
|
| 204 |
+
super().__init__(n_classes, n_input_channels, config='152', input_dimension=input_dimension,
|
| 205 |
+
final_layer_dropout=final_layer_dropout, stochastic_depth_p=stochastic_depth_p,
|
| 206 |
+
squeeze_excitation=squeeze_excitation, squeeze_excitation_rd_ratio=squeeze_excitation_rd_ratio)
|
| 207 |
+
|
| 208 |
+
class ResNet50bn(ResNetD):
|
| 209 |
+
def __init__(self, n_classes: int, n_input_channels: int = 3, input_dimension: int = 2,
|
| 210 |
+
final_layer_dropout: float = 0.0, stochastic_depth_p: float = 0.0, squeeze_excitation: bool = False,
|
| 211 |
+
squeeze_excitation_rd_ratio: float = 1./16):
|
| 212 |
+
super().__init__(n_classes, n_input_channels, config='50_bn', input_dimension=input_dimension,
|
| 213 |
+
final_layer_dropout=final_layer_dropout, stochastic_depth_p=stochastic_depth_p,
|
| 214 |
+
squeeze_excitation=squeeze_excitation, squeeze_excitation_rd_ratio=squeeze_excitation_rd_ratio)
|
| 215 |
+
|
| 216 |
+
class ResNet152bn(ResNetD):
|
| 217 |
+
def __init__(self, n_classes: int, n_input_channels: int = 3, input_dimension: int = 2,
|
| 218 |
+
final_layer_dropout: float = 0.0, stochastic_depth_p: float = 0.0, squeeze_excitation: bool = False,
|
| 219 |
+
squeeze_excitation_rd_ratio: float = 1./16):
|
| 220 |
+
super().__init__(n_classes, n_input_channels, config='152_bn', input_dimension=input_dimension,
|
| 221 |
+
final_layer_dropout=final_layer_dropout, stochastic_depth_p=stochastic_depth_p,
|
| 222 |
+
squeeze_excitation=squeeze_excitation, squeeze_excitation_rd_ratio=squeeze_excitation_rd_ratio)
|
| 223 |
+
|
| 224 |
+
|
| 225 |
+
if __name__ == '__main__':
|
| 226 |
+
data = torch.rand((1, 3, 224, 224))
|
| 227 |
+
|
| 228 |
+
model = ResNet50bn(10, 3)
|
| 229 |
+
import hiddenlayer as hl
|
| 230 |
+
|
| 231 |
+
g = hl.build_graph(model, data,
|
| 232 |
+
transforms=None)
|
| 233 |
+
g.save("network_architecture.pdf")
|
| 234 |
+
del g
|
| 235 |
+
|
| 236 |
+
#print(model.compute_conv_feature_map_size((32, 32)))
|
RADAR_inference/dynamic_network_architectures/architectures/resnet_vl.py
ADDED
|
@@ -0,0 +1,225 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
import torch.nn.functional as F
|
| 4 |
+
from torch.autograd import Variable
|
| 5 |
+
import math
|
| 6 |
+
from functools import partial
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
def conv3x3x3(in_planes, out_planes, stride=1, dilation=1):
|
| 10 |
+
# 3x3x3 convolution with padding
|
| 11 |
+
return nn.Conv3d(
|
| 12 |
+
in_planes,
|
| 13 |
+
out_planes,
|
| 14 |
+
kernel_size=3,
|
| 15 |
+
dilation=dilation,
|
| 16 |
+
stride=stride,
|
| 17 |
+
padding=dilation,
|
| 18 |
+
bias=False)
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
def downsample_basic_block(x, planes, stride, no_cuda=False):
|
| 22 |
+
out = F.avg_pool3d(x, kernel_size=1, stride=stride)
|
| 23 |
+
zero_pads = torch.Tensor(
|
| 24 |
+
out.size(0), planes - out.size(1), out.size(2), out.size(3),
|
| 25 |
+
out.size(4)).zero_()
|
| 26 |
+
if not no_cuda:
|
| 27 |
+
# if isinstance(out.data, torch.cuda.FloatTensor):
|
| 28 |
+
zero_pads = zero_pads.cuda()
|
| 29 |
+
|
| 30 |
+
out = Variable(torch.cat([out.data, zero_pads], dim=1))
|
| 31 |
+
|
| 32 |
+
return out
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
class BasicBlock(nn.Module):
|
| 36 |
+
expansion = 1
|
| 37 |
+
|
| 38 |
+
def __init__(self, inplanes, planes, stride=1, dilation=1, downsample=None):
|
| 39 |
+
super(BasicBlock, self).__init__()
|
| 40 |
+
self.conv1 = conv3x3x3(inplanes, planes, stride=stride, dilation=dilation)
|
| 41 |
+
self.bn1 = nn.BatchNorm3d(planes)
|
| 42 |
+
self.relu = nn.ReLU(inplace=True)
|
| 43 |
+
self.conv2 = conv3x3x3(planes, planes, dilation=dilation)
|
| 44 |
+
self.bn2 = nn.BatchNorm3d(planes)
|
| 45 |
+
self.downsample = downsample
|
| 46 |
+
self.stride = stride
|
| 47 |
+
self.dilation = dilation
|
| 48 |
+
|
| 49 |
+
def forward(self, x):
|
| 50 |
+
residual = x
|
| 51 |
+
|
| 52 |
+
out = self.conv1(x)
|
| 53 |
+
out = self.bn1(out)
|
| 54 |
+
out = self.relu(out)
|
| 55 |
+
out = self.conv2(out)
|
| 56 |
+
out = self.bn2(out)
|
| 57 |
+
|
| 58 |
+
if self.downsample is not None:
|
| 59 |
+
residual = self.downsample(x)
|
| 60 |
+
|
| 61 |
+
out += residual
|
| 62 |
+
out = self.relu(out)
|
| 63 |
+
|
| 64 |
+
return out
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
class Bottleneck(nn.Module):
|
| 68 |
+
expansion = 4
|
| 69 |
+
|
| 70 |
+
def __init__(self, inplanes, planes, stride=1, dilation=1, downsample=None):
|
| 71 |
+
super(Bottleneck, self).__init__()
|
| 72 |
+
self.conv1 = nn.Conv3d(inplanes, planes, kernel_size=1, bias=False)
|
| 73 |
+
self.bn1 = nn.BatchNorm3d(planes)
|
| 74 |
+
self.conv2 = nn.Conv3d(
|
| 75 |
+
planes, planes, kernel_size=3, stride=stride, dilation=dilation, padding=dilation, bias=False)
|
| 76 |
+
self.bn2 = nn.BatchNorm3d(planes)
|
| 77 |
+
self.conv3 = nn.Conv3d(planes, planes * 4, kernel_size=1, bias=False)
|
| 78 |
+
self.bn3 = nn.BatchNorm3d(planes * 4)
|
| 79 |
+
self.relu = nn.ReLU(inplace=True)
|
| 80 |
+
self.downsample = downsample
|
| 81 |
+
self.stride = stride
|
| 82 |
+
self.dilation = dilation
|
| 83 |
+
|
| 84 |
+
def forward(self, x):
|
| 85 |
+
residual = x
|
| 86 |
+
|
| 87 |
+
out = self.conv1(x)
|
| 88 |
+
out = self.bn1(out)
|
| 89 |
+
out = self.relu(out)
|
| 90 |
+
|
| 91 |
+
out = self.conv2(out)
|
| 92 |
+
out = self.bn2(out)
|
| 93 |
+
out = self.relu(out)
|
| 94 |
+
|
| 95 |
+
out = self.conv3(out)
|
| 96 |
+
out = self.bn3(out)
|
| 97 |
+
|
| 98 |
+
if self.downsample is not None:
|
| 99 |
+
residual = self.downsample(x)
|
| 100 |
+
|
| 101 |
+
out += residual
|
| 102 |
+
out = self.relu(out)
|
| 103 |
+
|
| 104 |
+
return out
|
| 105 |
+
|
| 106 |
+
class ResNet(nn.Module):
|
| 107 |
+
|
| 108 |
+
def __init__(self,
|
| 109 |
+
block,
|
| 110 |
+
layers,
|
| 111 |
+
shortcut_type='B',
|
| 112 |
+
no_cuda = False):
|
| 113 |
+
self.inplanes = 64
|
| 114 |
+
self.no_cuda = no_cuda
|
| 115 |
+
super(ResNet, self).__init__()
|
| 116 |
+
self.conv1 = nn.Conv3d(
|
| 117 |
+
1,
|
| 118 |
+
64,
|
| 119 |
+
kernel_size=7,
|
| 120 |
+
stride=(1, 2, 2),
|
| 121 |
+
padding=(3, 3, 3),
|
| 122 |
+
bias=False)
|
| 123 |
+
|
| 124 |
+
self.bn1 = nn.BatchNorm3d(64)
|
| 125 |
+
self.relu = nn.ReLU(inplace=True)
|
| 126 |
+
self.maxpool = nn.MaxPool3d(kernel_size=(3, 3, 3), stride=2, padding=1)
|
| 127 |
+
self.layer1 = self._make_layer(block, 64, layers[0], shortcut_type)
|
| 128 |
+
self.layer2 = self._make_layer(
|
| 129 |
+
block, 128, layers[1], shortcut_type, stride=2)
|
| 130 |
+
self.layer3 = self._make_layer(
|
| 131 |
+
block, 256, layers[2], shortcut_type, stride=2, dilation=2)
|
| 132 |
+
self.layer4 = self._make_layer(
|
| 133 |
+
block, 512, layers[3], shortcut_type, stride=2, dilation=4)
|
| 134 |
+
|
| 135 |
+
for m in self.modules():
|
| 136 |
+
if isinstance(m, nn.Conv3d):
|
| 137 |
+
m.weight = nn.init.kaiming_normal(m.weight, mode='fan_out')
|
| 138 |
+
elif isinstance(m, nn.BatchNorm3d):
|
| 139 |
+
m.weight.data.fill_(1)
|
| 140 |
+
m.bias.data.zero_()
|
| 141 |
+
|
| 142 |
+
def _make_layer(self, block, planes, blocks, shortcut_type, stride=1, dilation=1):
|
| 143 |
+
downsample = None
|
| 144 |
+
if stride != 1 or self.inplanes != planes * block.expansion:
|
| 145 |
+
if shortcut_type == 'A':
|
| 146 |
+
downsample = partial(
|
| 147 |
+
downsample_basic_block,
|
| 148 |
+
planes=planes * block.expansion,
|
| 149 |
+
stride=stride,
|
| 150 |
+
no_cuda=self.no_cuda)
|
| 151 |
+
else:
|
| 152 |
+
downsample = nn.Sequential(
|
| 153 |
+
nn.Conv3d(
|
| 154 |
+
self.inplanes,
|
| 155 |
+
planes * block.expansion,
|
| 156 |
+
kernel_size=1,
|
| 157 |
+
stride=stride,
|
| 158 |
+
bias=False), nn.BatchNorm3d(planes * block.expansion))
|
| 159 |
+
|
| 160 |
+
layers = []
|
| 161 |
+
layers.append(block(self.inplanes, planes, stride=stride, dilation=dilation, downsample=downsample))
|
| 162 |
+
self.inplanes = planes * block.expansion
|
| 163 |
+
for i in range(1, blocks):
|
| 164 |
+
layers.append(block(self.inplanes, planes, dilation=dilation))
|
| 165 |
+
|
| 166 |
+
return nn.Sequential(*layers)
|
| 167 |
+
|
| 168 |
+
def forward(self, x):
|
| 169 |
+
x = self.conv1(x)
|
| 170 |
+
x = self.bn1(x)
|
| 171 |
+
x6 = self.relu(x)
|
| 172 |
+
x5 = self.maxpool(x6)
|
| 173 |
+
x4 = self.layer1(x5)
|
| 174 |
+
x3 = self.layer2(x4)
|
| 175 |
+
x2 = self.layer3(x3)
|
| 176 |
+
x1 = self.layer4(x2)
|
| 177 |
+
|
| 178 |
+
return [x6, x4, x3, x2, x1]
|
| 179 |
+
|
| 180 |
+
def resnet10(**kwargs):
|
| 181 |
+
"""Constructs a ResNet-18 model.
|
| 182 |
+
"""
|
| 183 |
+
model = ResNet(BasicBlock, [1, 1, 1, 1], **kwargs)
|
| 184 |
+
return model
|
| 185 |
+
|
| 186 |
+
def resnet18(**kwargs):
|
| 187 |
+
"""Constructs a ResNet-18 model.
|
| 188 |
+
"""
|
| 189 |
+
model = ResNet(BasicBlock, [2, 2, 2, 2], **kwargs)
|
| 190 |
+
return model
|
| 191 |
+
|
| 192 |
+
|
| 193 |
+
def resnet34(**kwargs):
|
| 194 |
+
"""Constructs a ResNet-34 model.
|
| 195 |
+
"""
|
| 196 |
+
model = ResNet(BasicBlock, [3, 4, 6, 3], **kwargs)
|
| 197 |
+
return model
|
| 198 |
+
|
| 199 |
+
|
| 200 |
+
def resnet50(**kwargs):
|
| 201 |
+
"""Constructs a ResNet-50 model.
|
| 202 |
+
"""
|
| 203 |
+
model = ResNet(Bottleneck, [3, 4, 6, 3], **kwargs)
|
| 204 |
+
return model
|
| 205 |
+
|
| 206 |
+
|
| 207 |
+
def resnet101(**kwargs):
|
| 208 |
+
"""Constructs a ResNet-101 model.
|
| 209 |
+
"""
|
| 210 |
+
model = ResNet(Bottleneck, [3, 4, 23, 3], **kwargs)
|
| 211 |
+
return model
|
| 212 |
+
|
| 213 |
+
|
| 214 |
+
def resnet152(**kwargs):
|
| 215 |
+
"""Constructs a ResNet-101 model.
|
| 216 |
+
"""
|
| 217 |
+
model = ResNet(Bottleneck, [3, 8, 36, 3], **kwargs)
|
| 218 |
+
return model
|
| 219 |
+
|
| 220 |
+
|
| 221 |
+
def resnet200(**kwargs):
|
| 222 |
+
"""Constructs a ResNet-101 model.
|
| 223 |
+
"""
|
| 224 |
+
model = ResNet(Bottleneck, [3, 24, 36, 3], **kwargs)
|
| 225 |
+
return model
|
RADAR_inference/dynamic_network_architectures/architectures/unet.py
ADDED
|
@@ -0,0 +1,218 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import Union, Type, List, Tuple
|
| 2 |
+
|
| 3 |
+
import torch
|
| 4 |
+
from dynamic_network_architectures.building_blocks.helper import convert_conv_op_to_dim
|
| 5 |
+
from dynamic_network_architectures.building_blocks.plain_conv_encoder import PlainConvEncoder
|
| 6 |
+
from dynamic_network_architectures.building_blocks.residual import BasicBlockD, BottleneckD
|
| 7 |
+
from dynamic_network_architectures.building_blocks.residual_encoders import ResidualEncoder
|
| 8 |
+
from dynamic_network_architectures.building_blocks.unet_decoder import UNetDecoder
|
| 9 |
+
from dynamic_network_architectures.building_blocks.unet_residual_decoder import UNetResDecoder
|
| 10 |
+
from dynamic_network_architectures.initialization.weight_init import InitWeights_He
|
| 11 |
+
from dynamic_network_architectures.initialization.weight_init import init_last_bn_before_add_to_0
|
| 12 |
+
from torch import nn
|
| 13 |
+
from torch.nn.modules.conv import _ConvNd
|
| 14 |
+
from torch.nn.modules.dropout import _DropoutNd
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
class PlainConvUNet(nn.Module):
|
| 18 |
+
def __init__(self,
|
| 19 |
+
input_channels: int,
|
| 20 |
+
n_stages: int,
|
| 21 |
+
features_per_stage: Union[int, List[int], Tuple[int, ...]],
|
| 22 |
+
conv_op: Type[_ConvNd],
|
| 23 |
+
kernel_sizes: Union[int, List[int], Tuple[int, ...]],
|
| 24 |
+
strides: Union[int, List[int], Tuple[int, ...]],
|
| 25 |
+
n_conv_per_stage: Union[int, List[int], Tuple[int, ...]],
|
| 26 |
+
num_classes: int,
|
| 27 |
+
n_conv_per_stage_decoder: Union[int, Tuple[int, ...], List[int]],
|
| 28 |
+
conv_bias: bool = False,
|
| 29 |
+
norm_op: Union[None, Type[nn.Module]] = None,
|
| 30 |
+
norm_op_kwargs: dict = None,
|
| 31 |
+
dropout_op: Union[None, Type[_DropoutNd]] = None,
|
| 32 |
+
dropout_op_kwargs: dict = None,
|
| 33 |
+
nonlin: Union[None, Type[torch.nn.Module]] = None,
|
| 34 |
+
nonlin_kwargs: dict = None,
|
| 35 |
+
deep_supervision: bool = False,
|
| 36 |
+
nonlin_first: bool = False
|
| 37 |
+
):
|
| 38 |
+
"""
|
| 39 |
+
nonlin_first: if True you get conv -> nonlin -> norm. Else it's conv -> norm -> nonlin
|
| 40 |
+
"""
|
| 41 |
+
super().__init__()
|
| 42 |
+
if isinstance(n_conv_per_stage, int):
|
| 43 |
+
n_conv_per_stage = [n_conv_per_stage] * n_stages
|
| 44 |
+
if isinstance(n_conv_per_stage_decoder, int):
|
| 45 |
+
n_conv_per_stage_decoder = [n_conv_per_stage_decoder] * (n_stages - 1)
|
| 46 |
+
assert len(n_conv_per_stage) == n_stages, "n_conv_per_stage must have as many entries as we have " \
|
| 47 |
+
f"resolution stages. here: {n_stages}. " \
|
| 48 |
+
f"n_conv_per_stage: {n_conv_per_stage}"
|
| 49 |
+
assert len(n_conv_per_stage_decoder) == (n_stages - 1), "n_conv_per_stage_decoder must have one less entries " \
|
| 50 |
+
f"as we have resolution stages. here: {n_stages} " \
|
| 51 |
+
f"stages, so it should have {n_stages - 1} entries. " \
|
| 52 |
+
f"n_conv_per_stage_decoder: {n_conv_per_stage_decoder}"
|
| 53 |
+
self.encoder = PlainConvEncoder(input_channels, n_stages, features_per_stage, conv_op, kernel_sizes, strides,
|
| 54 |
+
n_conv_per_stage, conv_bias, norm_op, norm_op_kwargs, dropout_op,
|
| 55 |
+
dropout_op_kwargs, nonlin, nonlin_kwargs, return_skips=True,
|
| 56 |
+
nonlin_first=nonlin_first)
|
| 57 |
+
self.decoder = UNetDecoder(self.encoder, num_classes, n_conv_per_stage_decoder, deep_supervision,
|
| 58 |
+
nonlin_first=nonlin_first)
|
| 59 |
+
|
| 60 |
+
def forward(self, x):
|
| 61 |
+
skips = self.encoder(x)
|
| 62 |
+
return self.decoder(skips)
|
| 63 |
+
|
| 64 |
+
def compute_conv_feature_map_size(self, input_size):
|
| 65 |
+
assert len(input_size) == convert_conv_op_to_dim(self.encoder.conv_op), "just give the image size without color/feature channels or " \
|
| 66 |
+
"batch channel. Do not give input_size=(b, c, x, y(, z)). " \
|
| 67 |
+
"Give input_size=(x, y(, z))!"
|
| 68 |
+
return self.encoder.compute_conv_feature_map_size(input_size) + self.decoder.compute_conv_feature_map_size(input_size)
|
| 69 |
+
|
| 70 |
+
@staticmethod
|
| 71 |
+
def initialize(module):
|
| 72 |
+
InitWeights_He(1e-2)(module)
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
class ResidualEncoderUNet(nn.Module):
|
| 76 |
+
def __init__(self,
|
| 77 |
+
input_channels: int,
|
| 78 |
+
n_stages: int,
|
| 79 |
+
features_per_stage: Union[int, List[int], Tuple[int, ...]],
|
| 80 |
+
conv_op: Type[_ConvNd],
|
| 81 |
+
kernel_sizes: Union[int, List[int], Tuple[int, ...]],
|
| 82 |
+
strides: Union[int, List[int], Tuple[int, ...]],
|
| 83 |
+
n_blocks_per_stage: Union[int, List[int], Tuple[int, ...]],
|
| 84 |
+
num_classes: int,
|
| 85 |
+
n_conv_per_stage_decoder: Union[int, Tuple[int, ...], List[int]],
|
| 86 |
+
conv_bias: bool = False,
|
| 87 |
+
norm_op: Union[None, Type[nn.Module]] = None,
|
| 88 |
+
norm_op_kwargs: dict = None,
|
| 89 |
+
dropout_op: Union[None, Type[_DropoutNd]] = None,
|
| 90 |
+
dropout_op_kwargs: dict = None,
|
| 91 |
+
nonlin: Union[None, Type[torch.nn.Module]] = None,
|
| 92 |
+
nonlin_kwargs: dict = None,
|
| 93 |
+
deep_supervision: bool = False,
|
| 94 |
+
block: Union[Type[BasicBlockD], Type[BottleneckD]] = BasicBlockD,
|
| 95 |
+
bottleneck_channels: Union[int, List[int], Tuple[int, ...]] = None,
|
| 96 |
+
stem_channels: int = None
|
| 97 |
+
):
|
| 98 |
+
super().__init__()
|
| 99 |
+
if isinstance(n_blocks_per_stage, int):
|
| 100 |
+
n_blocks_per_stage = [n_blocks_per_stage] * n_stages
|
| 101 |
+
if isinstance(n_conv_per_stage_decoder, int):
|
| 102 |
+
n_conv_per_stage_decoder = [n_conv_per_stage_decoder] * (n_stages - 1)
|
| 103 |
+
assert len(n_blocks_per_stage) == n_stages, "n_blocks_per_stage must have as many entries as we have " \
|
| 104 |
+
f"resolution stages. here: {n_stages}. " \
|
| 105 |
+
f"n_blocks_per_stage: {n_blocks_per_stage}"
|
| 106 |
+
assert len(n_conv_per_stage_decoder) == (n_stages - 1), "n_conv_per_stage_decoder must have one less entries " \
|
| 107 |
+
f"as we have resolution stages. here: {n_stages} " \
|
| 108 |
+
f"stages, so it should have {n_stages - 1} entries. " \
|
| 109 |
+
f"n_conv_per_stage_decoder: {n_conv_per_stage_decoder}"
|
| 110 |
+
self.encoder = ResidualEncoder(input_channels, n_stages, features_per_stage, conv_op, kernel_sizes, strides,
|
| 111 |
+
n_blocks_per_stage, conv_bias, norm_op, norm_op_kwargs, dropout_op,
|
| 112 |
+
dropout_op_kwargs, nonlin, nonlin_kwargs, block, bottleneck_channels,
|
| 113 |
+
return_skips=True, disable_default_stem=False, stem_channels=stem_channels)
|
| 114 |
+
self.decoder = UNetDecoder(self.encoder, num_classes, n_conv_per_stage_decoder, deep_supervision)
|
| 115 |
+
|
| 116 |
+
def forward(self, x):
|
| 117 |
+
skips = self.encoder(x)
|
| 118 |
+
return self.decoder(skips)
|
| 119 |
+
|
| 120 |
+
def compute_conv_feature_map_size(self, input_size):
|
| 121 |
+
assert len(input_size) == convert_conv_op_to_dim(self.encoder.conv_op), "just give the image size without color/feature channels or " \
|
| 122 |
+
"batch channel. Do not give input_size=(b, c, x, y(, z)). " \
|
| 123 |
+
"Give input_size=(x, y(, z))!"
|
| 124 |
+
return self.encoder.compute_conv_feature_map_size(input_size) + self.decoder.compute_conv_feature_map_size(input_size)
|
| 125 |
+
|
| 126 |
+
@staticmethod
|
| 127 |
+
def initialize(module):
|
| 128 |
+
InitWeights_He(1e-2)(module)
|
| 129 |
+
init_last_bn_before_add_to_0(module)
|
| 130 |
+
|
| 131 |
+
|
| 132 |
+
class ResidualUNet(nn.Module):
|
| 133 |
+
def __init__(self,
|
| 134 |
+
input_channels: int,
|
| 135 |
+
n_stages: int,
|
| 136 |
+
features_per_stage: Union[int, List[int], Tuple[int, ...]],
|
| 137 |
+
conv_op: Type[_ConvNd],
|
| 138 |
+
kernel_sizes: Union[int, List[int], Tuple[int, ...]],
|
| 139 |
+
strides: Union[int, List[int], Tuple[int, ...]],
|
| 140 |
+
n_blocks_per_stage: Union[int, List[int], Tuple[int, ...]],
|
| 141 |
+
num_classes: int,
|
| 142 |
+
n_conv_per_stage_decoder: Union[int, Tuple[int, ...], List[int]],
|
| 143 |
+
conv_bias: bool = False,
|
| 144 |
+
norm_op: Union[None, Type[nn.Module]] = None,
|
| 145 |
+
norm_op_kwargs: dict = None,
|
| 146 |
+
dropout_op: Union[None, Type[_DropoutNd]] = None,
|
| 147 |
+
dropout_op_kwargs: dict = None,
|
| 148 |
+
nonlin: Union[None, Type[torch.nn.Module]] = None,
|
| 149 |
+
nonlin_kwargs: dict = None,
|
| 150 |
+
deep_supervision: bool = False,
|
| 151 |
+
block: Union[Type[BasicBlockD], Type[BottleneckD]] = BasicBlockD,
|
| 152 |
+
bottleneck_channels: Union[int, List[int], Tuple[int, ...]] = None,
|
| 153 |
+
stem_channels: int = None
|
| 154 |
+
):
|
| 155 |
+
super().__init__()
|
| 156 |
+
if isinstance(n_blocks_per_stage, int):
|
| 157 |
+
n_blocks_per_stage = [n_blocks_per_stage] * n_stages
|
| 158 |
+
if isinstance(n_conv_per_stage_decoder, int):
|
| 159 |
+
n_conv_per_stage_decoder = [n_conv_per_stage_decoder] * (n_stages - 1)
|
| 160 |
+
assert len(n_blocks_per_stage) == n_stages, "n_blocks_per_stage must have as many entries as we have " \
|
| 161 |
+
f"resolution stages. here: {n_stages}. " \
|
| 162 |
+
f"n_blocks_per_stage: {n_blocks_per_stage}"
|
| 163 |
+
assert len(n_conv_per_stage_decoder) == (n_stages - 1), "n_conv_per_stage_decoder must have one less entries " \
|
| 164 |
+
f"as we have resolution stages. here: {n_stages} " \
|
| 165 |
+
f"stages, so it should have {n_stages - 1} entries. " \
|
| 166 |
+
f"n_conv_per_stage_decoder: {n_conv_per_stage_decoder}"
|
| 167 |
+
self.encoder = ResidualEncoder(input_channels, n_stages, features_per_stage, conv_op, kernel_sizes, strides,
|
| 168 |
+
n_blocks_per_stage, conv_bias, norm_op, norm_op_kwargs, dropout_op,
|
| 169 |
+
dropout_op_kwargs, nonlin, nonlin_kwargs, block, bottleneck_channels,
|
| 170 |
+
return_skips=True, disable_default_stem=False, stem_channels=stem_channels)
|
| 171 |
+
self.decoder = UNetResDecoder(self.encoder, num_classes, n_conv_per_stage_decoder, deep_supervision)
|
| 172 |
+
|
| 173 |
+
def forward(self, x):
|
| 174 |
+
skips = self.encoder(x)
|
| 175 |
+
return self.decoder(skips)
|
| 176 |
+
|
| 177 |
+
def compute_conv_feature_map_size(self, input_size):
|
| 178 |
+
assert len(input_size) == convert_conv_op_to_dim(self.encoder.conv_op), "just give the image size without color/feature channels or " \
|
| 179 |
+
"batch channel. Do not give input_size=(b, c, x, y(, z)). " \
|
| 180 |
+
"Give input_size=(x, y(, z))!"
|
| 181 |
+
return self.encoder.compute_conv_feature_map_size(input_size) + self.decoder.compute_conv_feature_map_size(input_size)
|
| 182 |
+
|
| 183 |
+
@staticmethod
|
| 184 |
+
def initialize(module):
|
| 185 |
+
InitWeights_He(1e-2)(module)
|
| 186 |
+
init_last_bn_before_add_to_0(module)
|
| 187 |
+
|
| 188 |
+
|
| 189 |
+
if __name__ == '__main__':
|
| 190 |
+
data = torch.rand((1, 4, 128, 128, 128))
|
| 191 |
+
|
| 192 |
+
model = PlainConvUNet(4, 6, (32, 64, 125, 256, 320, 320), nn.Conv3d, 3, (1, 2, 2, 2, 2, 2), (2, 2, 2, 2, 2, 2), 4,
|
| 193 |
+
(2, 2, 2, 2, 2), False, nn.BatchNorm3d, None, None, None, nn.ReLU, deep_supervision=True)
|
| 194 |
+
|
| 195 |
+
if False:
|
| 196 |
+
import hiddenlayer as hl
|
| 197 |
+
|
| 198 |
+
g = hl.build_graph(model, data,
|
| 199 |
+
transforms=None)
|
| 200 |
+
g.save("network_architecture.pdf")
|
| 201 |
+
del g
|
| 202 |
+
|
| 203 |
+
print(model.compute_conv_feature_map_size(data.shape[2:]))
|
| 204 |
+
|
| 205 |
+
data = torch.rand((1, 4, 512, 512))
|
| 206 |
+
|
| 207 |
+
model = PlainConvUNet(4, 8, (32, 64, 125, 256, 512, 512, 512, 512), nn.Conv2d, 3, (1, 2, 2, 2, 2, 2, 2, 2), (2, 2, 2, 2, 2, 2, 2, 2), 4,
|
| 208 |
+
(2, 2, 2, 2, 2, 2, 2), False, nn.BatchNorm2d, None, None, None, nn.ReLU, deep_supervision=True)
|
| 209 |
+
|
| 210 |
+
if False:
|
| 211 |
+
import hiddenlayer as hl
|
| 212 |
+
|
| 213 |
+
g = hl.build_graph(model, data,
|
| 214 |
+
transforms=None)
|
| 215 |
+
g.save("network_architecture.pdf")
|
| 216 |
+
del g
|
| 217 |
+
|
| 218 |
+
print(model.compute_conv_feature_map_size(data.shape[2:]))
|
RADAR_inference/dynamic_network_architectures/architectures/unet_lightdecoder.py
ADDED
|
@@ -0,0 +1,218 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import Union, Type, List, Tuple
|
| 2 |
+
|
| 3 |
+
import torch
|
| 4 |
+
from dynamic_network_architectures.building_blocks.helper import convert_conv_op_to_dim
|
| 5 |
+
from dynamic_network_architectures.building_blocks.plain_conv_encoder import PlainConvEncoder
|
| 6 |
+
from dynamic_network_architectures.building_blocks.residual import BasicBlockD, BottleneckD
|
| 7 |
+
from dynamic_network_architectures.building_blocks.residual_encoders import ResidualEncoder
|
| 8 |
+
from dynamic_network_architectures.building_blocks.unet_decoder_light import UNetDecoder
|
| 9 |
+
from dynamic_network_architectures.building_blocks.unet_residual_decoder import UNetResDecoder
|
| 10 |
+
from dynamic_network_architectures.initialization.weight_init import InitWeights_He
|
| 11 |
+
from dynamic_network_architectures.initialization.weight_init import init_last_bn_before_add_to_0
|
| 12 |
+
from torch import nn
|
| 13 |
+
from torch.nn.modules.conv import _ConvNd
|
| 14 |
+
from torch.nn.modules.dropout import _DropoutNd
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
class PlainConvUNetLightD(nn.Module):
|
| 18 |
+
def __init__(self,
|
| 19 |
+
input_channels: int,
|
| 20 |
+
n_stages: int,
|
| 21 |
+
features_per_stage: Union[int, List[int], Tuple[int, ...]],
|
| 22 |
+
conv_op: Type[_ConvNd],
|
| 23 |
+
kernel_sizes: Union[int, List[int], Tuple[int, ...]],
|
| 24 |
+
strides: Union[int, List[int], Tuple[int, ...]],
|
| 25 |
+
n_conv_per_stage: Union[int, List[int], Tuple[int, ...]],
|
| 26 |
+
num_classes: int,
|
| 27 |
+
n_conv_per_stage_decoder: Union[int, Tuple[int, ...], List[int]],
|
| 28 |
+
conv_bias: bool = False,
|
| 29 |
+
norm_op: Union[None, Type[nn.Module]] = None,
|
| 30 |
+
norm_op_kwargs: dict = None,
|
| 31 |
+
dropout_op: Union[None, Type[_DropoutNd]] = None,
|
| 32 |
+
dropout_op_kwargs: dict = None,
|
| 33 |
+
nonlin: Union[None, Type[torch.nn.Module]] = None,
|
| 34 |
+
nonlin_kwargs: dict = None,
|
| 35 |
+
deep_supervision: bool = False,
|
| 36 |
+
nonlin_first: bool = False
|
| 37 |
+
):
|
| 38 |
+
"""
|
| 39 |
+
nonlin_first: if True you get conv -> nonlin -> norm. Else it's conv -> norm -> nonlin
|
| 40 |
+
"""
|
| 41 |
+
super().__init__()
|
| 42 |
+
if isinstance(n_conv_per_stage, int):
|
| 43 |
+
n_conv_per_stage = [n_conv_per_stage] * n_stages
|
| 44 |
+
if isinstance(n_conv_per_stage_decoder, int):
|
| 45 |
+
n_conv_per_stage_decoder = [n_conv_per_stage_decoder] * (n_stages - 1)
|
| 46 |
+
assert len(n_conv_per_stage) == n_stages, "n_conv_per_stage must have as many entries as we have " \
|
| 47 |
+
f"resolution stages. here: {n_stages}. " \
|
| 48 |
+
f"n_conv_per_stage: {n_conv_per_stage}"
|
| 49 |
+
assert len(n_conv_per_stage_decoder) == (n_stages - 1), "n_conv_per_stage_decoder must have one less entries " \
|
| 50 |
+
f"as we have resolution stages. here: {n_stages} " \
|
| 51 |
+
f"stages, so it should have {n_stages - 1} entries. " \
|
| 52 |
+
f"n_conv_per_stage_decoder: {n_conv_per_stage_decoder}"
|
| 53 |
+
self.encoder = PlainConvEncoder(input_channels, n_stages, features_per_stage, conv_op, kernel_sizes, strides,
|
| 54 |
+
n_conv_per_stage, conv_bias, norm_op, norm_op_kwargs, dropout_op,
|
| 55 |
+
dropout_op_kwargs, nonlin, nonlin_kwargs, return_skips=True,
|
| 56 |
+
nonlin_first=nonlin_first)
|
| 57 |
+
self.decoder = UNetDecoder(self.encoder, num_classes, n_conv_per_stage_decoder, deep_supervision,
|
| 58 |
+
nonlin_first=nonlin_first)
|
| 59 |
+
|
| 60 |
+
def forward(self, x):
|
| 61 |
+
skips = self.encoder(x)
|
| 62 |
+
return skips, self.decoder(skips)
|
| 63 |
+
|
| 64 |
+
def compute_conv_feature_map_size(self, input_size):
|
| 65 |
+
assert len(input_size) == convert_conv_op_to_dim(self.encoder.conv_op), "just give the image size without color/feature channels or " \
|
| 66 |
+
"batch channel. Do not give input_size=(b, c, x, y(, z)). " \
|
| 67 |
+
"Give input_size=(x, y(, z))!"
|
| 68 |
+
return self.encoder.compute_conv_feature_map_size(input_size) + self.decoder.compute_conv_feature_map_size(input_size)
|
| 69 |
+
|
| 70 |
+
@staticmethod
|
| 71 |
+
def initialize(module):
|
| 72 |
+
InitWeights_He(1e-2)(module)
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
class ResidualEncoderUNet(nn.Module):
|
| 76 |
+
def __init__(self,
|
| 77 |
+
input_channels: int,
|
| 78 |
+
n_stages: int,
|
| 79 |
+
features_per_stage: Union[int, List[int], Tuple[int, ...]],
|
| 80 |
+
conv_op: Type[_ConvNd],
|
| 81 |
+
kernel_sizes: Union[int, List[int], Tuple[int, ...]],
|
| 82 |
+
strides: Union[int, List[int], Tuple[int, ...]],
|
| 83 |
+
n_blocks_per_stage: Union[int, List[int], Tuple[int, ...]],
|
| 84 |
+
num_classes: int,
|
| 85 |
+
n_conv_per_stage_decoder: Union[int, Tuple[int, ...], List[int]],
|
| 86 |
+
conv_bias: bool = False,
|
| 87 |
+
norm_op: Union[None, Type[nn.Module]] = None,
|
| 88 |
+
norm_op_kwargs: dict = None,
|
| 89 |
+
dropout_op: Union[None, Type[_DropoutNd]] = None,
|
| 90 |
+
dropout_op_kwargs: dict = None,
|
| 91 |
+
nonlin: Union[None, Type[torch.nn.Module]] = None,
|
| 92 |
+
nonlin_kwargs: dict = None,
|
| 93 |
+
deep_supervision: bool = False,
|
| 94 |
+
block: Union[Type[BasicBlockD], Type[BottleneckD]] = BasicBlockD,
|
| 95 |
+
bottleneck_channels: Union[int, List[int], Tuple[int, ...]] = None,
|
| 96 |
+
stem_channels: int = None
|
| 97 |
+
):
|
| 98 |
+
super().__init__()
|
| 99 |
+
if isinstance(n_blocks_per_stage, int):
|
| 100 |
+
n_blocks_per_stage = [n_blocks_per_stage] * n_stages
|
| 101 |
+
if isinstance(n_conv_per_stage_decoder, int):
|
| 102 |
+
n_conv_per_stage_decoder = [n_conv_per_stage_decoder] * (n_stages - 1)
|
| 103 |
+
assert len(n_blocks_per_stage) == n_stages, "n_blocks_per_stage must have as many entries as we have " \
|
| 104 |
+
f"resolution stages. here: {n_stages}. " \
|
| 105 |
+
f"n_blocks_per_stage: {n_blocks_per_stage}"
|
| 106 |
+
assert len(n_conv_per_stage_decoder) == (n_stages - 1), "n_conv_per_stage_decoder must have one less entries " \
|
| 107 |
+
f"as we have resolution stages. here: {n_stages} " \
|
| 108 |
+
f"stages, so it should have {n_stages - 1} entries. " \
|
| 109 |
+
f"n_conv_per_stage_decoder: {n_conv_per_stage_decoder}"
|
| 110 |
+
self.encoder = ResidualEncoder(input_channels, n_stages, features_per_stage, conv_op, kernel_sizes, strides,
|
| 111 |
+
n_blocks_per_stage, conv_bias, norm_op, norm_op_kwargs, dropout_op,
|
| 112 |
+
dropout_op_kwargs, nonlin, nonlin_kwargs, block, bottleneck_channels,
|
| 113 |
+
return_skips=True, disable_default_stem=False, stem_channels=stem_channels)
|
| 114 |
+
self.decoder = UNetDecoder(self.encoder, num_classes, n_conv_per_stage_decoder, deep_supervision)
|
| 115 |
+
|
| 116 |
+
def forward(self, x):
|
| 117 |
+
skips = self.encoder(x)
|
| 118 |
+
return self.decoder(skips)
|
| 119 |
+
|
| 120 |
+
def compute_conv_feature_map_size(self, input_size):
|
| 121 |
+
assert len(input_size) == convert_conv_op_to_dim(self.encoder.conv_op), "just give the image size without color/feature channels or " \
|
| 122 |
+
"batch channel. Do not give input_size=(b, c, x, y(, z)). " \
|
| 123 |
+
"Give input_size=(x, y(, z))!"
|
| 124 |
+
return self.encoder.compute_conv_feature_map_size(input_size) + self.decoder.compute_conv_feature_map_size(input_size)
|
| 125 |
+
|
| 126 |
+
@staticmethod
|
| 127 |
+
def initialize(module):
|
| 128 |
+
InitWeights_He(1e-2)(module)
|
| 129 |
+
init_last_bn_before_add_to_0(module)
|
| 130 |
+
|
| 131 |
+
|
| 132 |
+
class ResidualUNet(nn.Module):
|
| 133 |
+
def __init__(self,
|
| 134 |
+
input_channels: int,
|
| 135 |
+
n_stages: int,
|
| 136 |
+
features_per_stage: Union[int, List[int], Tuple[int, ...]],
|
| 137 |
+
conv_op: Type[_ConvNd],
|
| 138 |
+
kernel_sizes: Union[int, List[int], Tuple[int, ...]],
|
| 139 |
+
strides: Union[int, List[int], Tuple[int, ...]],
|
| 140 |
+
n_blocks_per_stage: Union[int, List[int], Tuple[int, ...]],
|
| 141 |
+
num_classes: int,
|
| 142 |
+
n_conv_per_stage_decoder: Union[int, Tuple[int, ...], List[int]],
|
| 143 |
+
conv_bias: bool = False,
|
| 144 |
+
norm_op: Union[None, Type[nn.Module]] = None,
|
| 145 |
+
norm_op_kwargs: dict = None,
|
| 146 |
+
dropout_op: Union[None, Type[_DropoutNd]] = None,
|
| 147 |
+
dropout_op_kwargs: dict = None,
|
| 148 |
+
nonlin: Union[None, Type[torch.nn.Module]] = None,
|
| 149 |
+
nonlin_kwargs: dict = None,
|
| 150 |
+
deep_supervision: bool = False,
|
| 151 |
+
block: Union[Type[BasicBlockD], Type[BottleneckD]] = BasicBlockD,
|
| 152 |
+
bottleneck_channels: Union[int, List[int], Tuple[int, ...]] = None,
|
| 153 |
+
stem_channels: int = None
|
| 154 |
+
):
|
| 155 |
+
super().__init__()
|
| 156 |
+
if isinstance(n_blocks_per_stage, int):
|
| 157 |
+
n_blocks_per_stage = [n_blocks_per_stage] * n_stages
|
| 158 |
+
if isinstance(n_conv_per_stage_decoder, int):
|
| 159 |
+
n_conv_per_stage_decoder = [n_conv_per_stage_decoder] * (n_stages - 1)
|
| 160 |
+
assert len(n_blocks_per_stage) == n_stages, "n_blocks_per_stage must have as many entries as we have " \
|
| 161 |
+
f"resolution stages. here: {n_stages}. " \
|
| 162 |
+
f"n_blocks_per_stage: {n_blocks_per_stage}"
|
| 163 |
+
assert len(n_conv_per_stage_decoder) == (n_stages - 1), "n_conv_per_stage_decoder must have one less entries " \
|
| 164 |
+
f"as we have resolution stages. here: {n_stages} " \
|
| 165 |
+
f"stages, so it should have {n_stages - 1} entries. " \
|
| 166 |
+
f"n_conv_per_stage_decoder: {n_conv_per_stage_decoder}"
|
| 167 |
+
self.encoder = ResidualEncoder(input_channels, n_stages, features_per_stage, conv_op, kernel_sizes, strides,
|
| 168 |
+
n_blocks_per_stage, conv_bias, norm_op, norm_op_kwargs, dropout_op,
|
| 169 |
+
dropout_op_kwargs, nonlin, nonlin_kwargs, block, bottleneck_channels,
|
| 170 |
+
return_skips=True, disable_default_stem=False, stem_channels=stem_channels)
|
| 171 |
+
self.decoder = UNetResDecoder(self.encoder, num_classes, n_conv_per_stage_decoder, deep_supervision)
|
| 172 |
+
|
| 173 |
+
def forward(self, x):
|
| 174 |
+
skips = self.encoder(x)
|
| 175 |
+
return self.decoder(skips)
|
| 176 |
+
|
| 177 |
+
def compute_conv_feature_map_size(self, input_size):
|
| 178 |
+
assert len(input_size) == convert_conv_op_to_dim(self.encoder.conv_op), "just give the image size without color/feature channels or " \
|
| 179 |
+
"batch channel. Do not give input_size=(b, c, x, y(, z)). " \
|
| 180 |
+
"Give input_size=(x, y(, z))!"
|
| 181 |
+
return self.encoder.compute_conv_feature_map_size(input_size) + self.decoder.compute_conv_feature_map_size(input_size)
|
| 182 |
+
|
| 183 |
+
@staticmethod
|
| 184 |
+
def initialize(module):
|
| 185 |
+
InitWeights_He(1e-2)(module)
|
| 186 |
+
init_last_bn_before_add_to_0(module)
|
| 187 |
+
|
| 188 |
+
|
| 189 |
+
if __name__ == '__main__':
|
| 190 |
+
data = torch.rand((1, 4, 128, 128, 128))
|
| 191 |
+
|
| 192 |
+
model = PlainConvUNet(4, 6, (32, 64, 125, 256, 320, 320), nn.Conv3d, 3, (1, 2, 2, 2, 2, 2), (2, 2, 2, 2, 2, 2), 4,
|
| 193 |
+
(2, 2, 2, 2, 2), False, nn.BatchNorm3d, None, None, None, nn.ReLU, deep_supervision=True)
|
| 194 |
+
|
| 195 |
+
if False:
|
| 196 |
+
import hiddenlayer as hl
|
| 197 |
+
|
| 198 |
+
g = hl.build_graph(model, data,
|
| 199 |
+
transforms=None)
|
| 200 |
+
g.save("network_architecture.pdf")
|
| 201 |
+
del g
|
| 202 |
+
|
| 203 |
+
print(model.compute_conv_feature_map_size(data.shape[2:]))
|
| 204 |
+
|
| 205 |
+
data = torch.rand((1, 4, 512, 512))
|
| 206 |
+
|
| 207 |
+
model = PlainConvUNet(4, 8, (32, 64, 125, 256, 512, 512, 512, 512), nn.Conv2d, 3, (1, 2, 2, 2, 2, 2, 2, 2), (2, 2, 2, 2, 2, 2, 2, 2), 4,
|
| 208 |
+
(2, 2, 2, 2, 2, 2, 2), False, nn.BatchNorm2d, None, None, None, nn.ReLU, deep_supervision=True)
|
| 209 |
+
|
| 210 |
+
if False:
|
| 211 |
+
import hiddenlayer as hl
|
| 212 |
+
|
| 213 |
+
g = hl.build_graph(model, data,
|
| 214 |
+
transforms=None)
|
| 215 |
+
g.save("network_architecture.pdf")
|
| 216 |
+
del g
|
| 217 |
+
|
| 218 |
+
print(model.compute_conv_feature_map_size(data.shape[2:]))
|
RADAR_inference/dynamic_network_architectures/architectures/vgg.py
ADDED
|
@@ -0,0 +1,85 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from torch import nn
|
| 3 |
+
|
| 4 |
+
from dynamic_network_architectures.building_blocks.plain_conv_encoder import PlainConvEncoder
|
| 5 |
+
from dynamic_network_architectures.building_blocks.helper import get_matching_pool_op, get_default_network_config
|
| 6 |
+
|
| 7 |
+
_VGG_CONFIGS = {
|
| 8 |
+
'16': {'features_per_stage': (64, 128, 256, 512, 512, 512), 'n_conv_per_stage': (2, 2, 2, 3, 3, 3),
|
| 9 |
+
'strides': (1, 2, 2, 2, 2, 2)},
|
| 10 |
+
'19': {'features_per_stage': (64, 128, 256, 512, 512, 512), 'n_conv_per_stage': (2, 2, 3, 3, 4, 4),
|
| 11 |
+
'strides': (1, 2, 2, 2, 2, 2)},
|
| 12 |
+
'16_cifar': {'features_per_stage': (64, 128, 256, 512), 'n_conv_per_stage': (2, 3, 5, 5), 'strides': (1, 2, 2, 2)},
|
| 13 |
+
'19_cifar': {'features_per_stage': (64, 128, 256, 512), 'n_conv_per_stage': (3, 4, 5, 6), 'strides': (1, 2, 2, 2)},
|
| 14 |
+
}
|
| 15 |
+
|
| 16 |
+
_VGG_OPS = {
|
| 17 |
+
1: {'conv_op': nn.Conv1d, 'norm_op': nn.BatchNorm1d},
|
| 18 |
+
2: {'conv_op': nn.Conv2d, 'norm_op': nn.BatchNorm2d},
|
| 19 |
+
3: {'conv_op': nn.Conv3d, 'norm_op': nn.BatchNorm3d},
|
| 20 |
+
}
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
class VGG(nn.Module):
|
| 24 |
+
def __init__(self, n_classes: int, n_input_channel: int = 3, config='16', input_dimension=2):
|
| 25 |
+
"""
|
| 26 |
+
This is not 1:1 VGG because it does not have the bloated fully connected layers at the end. Since these were
|
| 27 |
+
counted towards the XX layers as well, we increase the number of convolutional layers so that we have the
|
| 28 |
+
desired number of conv layers in total
|
| 29 |
+
|
| 30 |
+
We also use batchnorm
|
| 31 |
+
"""
|
| 32 |
+
super().__init__()
|
| 33 |
+
cfg = _VGG_CONFIGS[config]
|
| 34 |
+
ops = get_default_network_config(dimension=input_dimension)
|
| 35 |
+
self.encoder = PlainConvEncoder(
|
| 36 |
+
n_input_channel, n_stages=len(cfg['features_per_stage']), features_per_stage=cfg['features_per_stage'],
|
| 37 |
+
conv_op=ops['conv_op'],
|
| 38 |
+
kernel_sizes=3, strides=cfg['strides'], n_conv_per_stage=cfg['n_conv_per_stage'], conv_bias=False,
|
| 39 |
+
norm_op=ops['norm_op'], norm_op_kwargs=None, dropout_op=None, dropout_op_kwargs=None, nonlin=nn.ReLU,
|
| 40 |
+
nonlin_kwargs={'inplace': True}, return_skips=False
|
| 41 |
+
)
|
| 42 |
+
self.gap = get_matching_pool_op(conv_op=ops['conv_op'], adaptive=True, pool_type='avg')(1)
|
| 43 |
+
self.classifier = nn.Linear(cfg['features_per_stage'][-1], n_classes, True)
|
| 44 |
+
|
| 45 |
+
def forward(self, x):
|
| 46 |
+
x = self.encoder(x)
|
| 47 |
+
x = self.gap(x).squeeze()
|
| 48 |
+
return self.classifier(x)
|
| 49 |
+
|
| 50 |
+
def compute_conv_feature_map_size(self, input_size):
|
| 51 |
+
return self.encoder.compute_conv_feature_map_size(input_size)
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
class VGG16(VGG):
|
| 55 |
+
def __init__(self, n_classes: int, n_input_channel: int = 3, input_dimension: int = 2):
|
| 56 |
+
super().__init__(n_classes, n_input_channel, config='16', input_dimension=input_dimension)
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
class VGG19(VGG):
|
| 60 |
+
def __init__(self, n_classes: int, n_input_channel: int = 3, input_dimension: int = 2):
|
| 61 |
+
super().__init__(n_classes, n_input_channel, config='19', input_dimension=input_dimension)
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
class VGG16_cifar(VGG):
|
| 65 |
+
def __init__(self, n_classes: int, n_input_channel: int = 3, input_dimension: int = 2):
|
| 66 |
+
super().__init__(n_classes, n_input_channel, config='16_cifar', input_dimension=input_dimension)
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
class VGG19_cifar(VGG):
|
| 70 |
+
def __init__(self, n_classes: int, n_input_channel: int = 3, input_dimension: int = 2):
|
| 71 |
+
super().__init__(n_classes, n_input_channel, config='19_cifar', input_dimension=input_dimension)
|
| 72 |
+
|
| 73 |
+
|
| 74 |
+
if __name__ == '__main__':
|
| 75 |
+
data = torch.rand((1, 3, 32, 32))
|
| 76 |
+
|
| 77 |
+
model = VGG19_cifar(10, 3)
|
| 78 |
+
import hiddenlayer as hl
|
| 79 |
+
|
| 80 |
+
g = hl.build_graph(model, data,
|
| 81 |
+
transforms=None)
|
| 82 |
+
g.save("network_architecture.pdf")
|
| 83 |
+
del g
|
| 84 |
+
|
| 85 |
+
print(model.compute_conv_feature_map_size((32, 32)))
|
RADAR_inference/dynamic_network_architectures/building_blocks/__init__.py
ADDED
|
File without changes
|
RADAR_inference/dynamic_network_architectures/building_blocks/helper.py
ADDED
|
@@ -0,0 +1,242 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import Type
|
| 2 |
+
import numpy as np
|
| 3 |
+
import torch.nn
|
| 4 |
+
from torch import nn
|
| 5 |
+
from torch.nn.modules.batchnorm import _BatchNorm
|
| 6 |
+
from torch.nn.modules.conv import _ConvNd, _ConvTransposeNd
|
| 7 |
+
from torch.nn.modules.dropout import _DropoutNd
|
| 8 |
+
from torch.nn.modules.instancenorm import _InstanceNorm
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
def convert_dim_to_conv_op(dimension: int) -> Type[_ConvNd]:
|
| 12 |
+
"""
|
| 13 |
+
:param dimension: 1, 2 or 3
|
| 14 |
+
:return: conv Class of corresponding dimension
|
| 15 |
+
"""
|
| 16 |
+
if dimension == 1:
|
| 17 |
+
return nn.Conv1d
|
| 18 |
+
elif dimension == 2:
|
| 19 |
+
return nn.Conv2d
|
| 20 |
+
elif dimension == 3:
|
| 21 |
+
return nn.Conv3d
|
| 22 |
+
else:
|
| 23 |
+
raise ValueError("Unknown dimension. Only 1, 2 and 3 are supported")
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
def convert_conv_op_to_dim(conv_op: Type[_ConvNd]) -> int:
|
| 27 |
+
"""
|
| 28 |
+
:param conv_op: conv class
|
| 29 |
+
:return: dimension: 1, 2 or 3
|
| 30 |
+
"""
|
| 31 |
+
if conv_op == nn.Conv1d:
|
| 32 |
+
return 1
|
| 33 |
+
elif conv_op == nn.Conv2d:
|
| 34 |
+
return 2
|
| 35 |
+
elif conv_op == nn.Conv3d:
|
| 36 |
+
return 3
|
| 37 |
+
else:
|
| 38 |
+
raise ValueError("Unknown dimension. Only 1d 2d and 3d conv are supported. got %s" % str(conv_op))
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
def get_matching_pool_op(conv_op: Type[_ConvNd] = None,
|
| 42 |
+
dimension: int = None,
|
| 43 |
+
adaptive=False,
|
| 44 |
+
pool_type: str = 'avg') -> Type[torch.nn.Module]:
|
| 45 |
+
"""
|
| 46 |
+
You MUST set EITHER conv_op OR dimension. Do not set both!
|
| 47 |
+
:param conv_op:
|
| 48 |
+
:param dimension:
|
| 49 |
+
:param adaptive:
|
| 50 |
+
:param pool_type: either 'avg' or 'max'
|
| 51 |
+
:return:
|
| 52 |
+
"""
|
| 53 |
+
assert not ((conv_op is not None) and (dimension is not None)), \
|
| 54 |
+
"You MUST set EITHER conv_op OR dimension. Do not set both!"
|
| 55 |
+
assert pool_type in ['avg', 'max'], 'pool_type must be either avg or max'
|
| 56 |
+
if conv_op is not None:
|
| 57 |
+
dimension = convert_conv_op_to_dim(conv_op)
|
| 58 |
+
assert dimension in [1, 2, 3], 'Dimension must be 1, 2 or 3'
|
| 59 |
+
|
| 60 |
+
if conv_op is not None:
|
| 61 |
+
dimension = convert_conv_op_to_dim(conv_op)
|
| 62 |
+
|
| 63 |
+
if dimension == 1:
|
| 64 |
+
if pool_type == 'avg':
|
| 65 |
+
if adaptive:
|
| 66 |
+
return nn.AdaptiveAvgPool1d
|
| 67 |
+
else:
|
| 68 |
+
return nn.AvgPool1d
|
| 69 |
+
elif pool_type == 'max':
|
| 70 |
+
if adaptive:
|
| 71 |
+
return nn.AdaptiveMaxPool1d
|
| 72 |
+
else:
|
| 73 |
+
return nn.MaxPool1d
|
| 74 |
+
elif dimension == 2:
|
| 75 |
+
if pool_type == 'avg':
|
| 76 |
+
if adaptive:
|
| 77 |
+
return nn.AdaptiveAvgPool2d
|
| 78 |
+
else:
|
| 79 |
+
return nn.AvgPool2d
|
| 80 |
+
elif pool_type == 'max':
|
| 81 |
+
if adaptive:
|
| 82 |
+
return nn.AdaptiveMaxPool2d
|
| 83 |
+
else:
|
| 84 |
+
return nn.MaxPool2d
|
| 85 |
+
elif dimension == 3:
|
| 86 |
+
if pool_type == 'avg':
|
| 87 |
+
if adaptive:
|
| 88 |
+
return nn.AdaptiveAvgPool3d
|
| 89 |
+
else:
|
| 90 |
+
return nn.AvgPool3d
|
| 91 |
+
elif pool_type == 'max':
|
| 92 |
+
if adaptive:
|
| 93 |
+
return nn.AdaptiveMaxPool3d
|
| 94 |
+
else:
|
| 95 |
+
return nn.MaxPool3d
|
| 96 |
+
|
| 97 |
+
|
| 98 |
+
def get_matching_instancenorm(conv_op: Type[_ConvNd] = None, dimension: int = None) -> Type[_InstanceNorm]:
|
| 99 |
+
"""
|
| 100 |
+
You MUST set EITHER conv_op OR dimension. Do not set both!
|
| 101 |
+
|
| 102 |
+
:param conv_op:
|
| 103 |
+
:param dimension:
|
| 104 |
+
:return:
|
| 105 |
+
"""
|
| 106 |
+
assert not ((conv_op is not None) and (dimension is not None)), \
|
| 107 |
+
"You MUST set EITHER conv_op OR dimension. Do not set both!"
|
| 108 |
+
if conv_op is not None:
|
| 109 |
+
dimension = convert_conv_op_to_dim(conv_op)
|
| 110 |
+
if dimension is not None:
|
| 111 |
+
assert dimension in [1, 2, 3], 'Dimension must be 1, 2 or 3'
|
| 112 |
+
if dimension == 1:
|
| 113 |
+
return nn.InstanceNorm1d
|
| 114 |
+
elif dimension == 2:
|
| 115 |
+
return nn.InstanceNorm2d
|
| 116 |
+
elif dimension == 3:
|
| 117 |
+
return nn.InstanceNorm3d
|
| 118 |
+
|
| 119 |
+
|
| 120 |
+
def get_matching_convtransp(conv_op: Type[_ConvNd] = None, dimension: int = None) -> Type[_ConvTransposeNd]:
|
| 121 |
+
"""
|
| 122 |
+
You MUST set EITHER conv_op OR dimension. Do not set both!
|
| 123 |
+
|
| 124 |
+
:param conv_op:
|
| 125 |
+
:param dimension:
|
| 126 |
+
:return:
|
| 127 |
+
"""
|
| 128 |
+
assert not ((conv_op is not None) and (dimension is not None)), \
|
| 129 |
+
"You MUST set EITHER conv_op OR dimension. Do not set both!"
|
| 130 |
+
if conv_op is not None:
|
| 131 |
+
dimension = convert_conv_op_to_dim(conv_op)
|
| 132 |
+
assert dimension in [1, 2, 3], 'Dimension must be 1, 2 or 3'
|
| 133 |
+
if dimension == 1:
|
| 134 |
+
return nn.ConvTranspose1d
|
| 135 |
+
elif dimension == 2:
|
| 136 |
+
return nn.ConvTranspose2d
|
| 137 |
+
elif dimension == 3:
|
| 138 |
+
return nn.ConvTranspose3d
|
| 139 |
+
|
| 140 |
+
|
| 141 |
+
def get_matching_batchnorm(conv_op: Type[_ConvNd] = None, dimension: int = None) -> Type[_BatchNorm]:
|
| 142 |
+
"""
|
| 143 |
+
You MUST set EITHER conv_op OR dimension. Do not set both!
|
| 144 |
+
|
| 145 |
+
:param conv_op:
|
| 146 |
+
:param dimension:
|
| 147 |
+
:return:
|
| 148 |
+
"""
|
| 149 |
+
assert not ((conv_op is not None) and (dimension is not None)), \
|
| 150 |
+
"You MUST set EITHER conv_op OR dimension. Do not set both!"
|
| 151 |
+
if conv_op is not None:
|
| 152 |
+
dimension = convert_conv_op_to_dim(conv_op)
|
| 153 |
+
assert dimension in [1, 2, 3], 'Dimension must be 1, 2 or 3'
|
| 154 |
+
if dimension == 1:
|
| 155 |
+
return nn.BatchNorm1d
|
| 156 |
+
elif dimension == 2:
|
| 157 |
+
return nn.BatchNorm2d
|
| 158 |
+
elif dimension == 3:
|
| 159 |
+
return nn.BatchNorm3d
|
| 160 |
+
|
| 161 |
+
|
| 162 |
+
def get_matching_dropout(conv_op: Type[_ConvNd] = None, dimension: int = None) -> Type[_DropoutNd]:
|
| 163 |
+
"""
|
| 164 |
+
You MUST set EITHER conv_op OR dimension. Do not set both!
|
| 165 |
+
|
| 166 |
+
:param conv_op:
|
| 167 |
+
:param dimension:
|
| 168 |
+
:return:
|
| 169 |
+
"""
|
| 170 |
+
assert not ((conv_op is not None) and (dimension is not None)), \
|
| 171 |
+
"You MUST set EITHER conv_op OR dimension. Do not set both!"
|
| 172 |
+
assert dimension in [1, 2, 3], 'Dimension must be 1, 2 or 3'
|
| 173 |
+
if dimension == 1:
|
| 174 |
+
return nn.Dropout
|
| 175 |
+
elif dimension == 2:
|
| 176 |
+
return nn.Dropout2d
|
| 177 |
+
elif dimension == 3:
|
| 178 |
+
return nn.Dropout3d
|
| 179 |
+
|
| 180 |
+
|
| 181 |
+
def maybe_convert_scalar_to_list(conv_op, scalar):
|
| 182 |
+
"""
|
| 183 |
+
useful for converting, for example, kernel_size=3 to [3, 3, 3] in case of nn.Conv3d
|
| 184 |
+
:param conv_op:
|
| 185 |
+
:param scalar:
|
| 186 |
+
:return:
|
| 187 |
+
"""
|
| 188 |
+
if not isinstance(scalar, (tuple, list, np.ndarray)):
|
| 189 |
+
if conv_op == nn.Conv2d:
|
| 190 |
+
return [scalar] * 2
|
| 191 |
+
elif conv_op == nn.Conv3d:
|
| 192 |
+
return [scalar] * 3
|
| 193 |
+
elif conv_op == nn.Conv1d:
|
| 194 |
+
return [scalar] * 1
|
| 195 |
+
else:
|
| 196 |
+
raise RuntimeError("Invalid conv op: %s" % str(conv_op))
|
| 197 |
+
else:
|
| 198 |
+
return scalar
|
| 199 |
+
|
| 200 |
+
|
| 201 |
+
def get_default_network_config(dimension: int = 2,
|
| 202 |
+
nonlin: str = "ReLU",
|
| 203 |
+
norm_type: str = "bn") -> dict:
|
| 204 |
+
"""
|
| 205 |
+
Use this to get a standard configuration. A network configuration looks like this:
|
| 206 |
+
|
| 207 |
+
config = {'conv_op': torch.nn.modules.conv.Conv2d,
|
| 208 |
+
'dropout_op': torch.nn.modules.dropout.Dropout2d,
|
| 209 |
+
'norm_op': torch.nn.modules.batchnorm.BatchNorm2d,
|
| 210 |
+
'norm_op_kwargs': {'eps': 1e-05, 'affine': True},
|
| 211 |
+
'nonlin': torch.nn.modules.activation.ReLU,
|
| 212 |
+
'nonlin_kwargs': {'inplace': True}}
|
| 213 |
+
|
| 214 |
+
There is no need to use get_default_network_config. You can create your own. Network configs are a convenient way of
|
| 215 |
+
setting dimensionality, normalization and nonlinearity.
|
| 216 |
+
|
| 217 |
+
:param dimension: integer denoting the dimension of the data. 1, 2 and 3 are accepted
|
| 218 |
+
:param nonlin: string (ReLU or LeakyReLU)
|
| 219 |
+
:param norm_type: string (bn=batch norm, in=instance norm)
|
| 220 |
+
torch.nn.Module
|
| 221 |
+
:return: dict
|
| 222 |
+
"""
|
| 223 |
+
config = {}
|
| 224 |
+
config['conv_op'] = convert_dim_to_conv_op(dimension)
|
| 225 |
+
config['dropout_op'] = get_matching_dropout(dimension=dimension)
|
| 226 |
+
if norm_type == "bn":
|
| 227 |
+
config['norm_op'] = get_matching_batchnorm(dimension=dimension)
|
| 228 |
+
elif norm_type == "in":
|
| 229 |
+
config['norm_op'] = get_matching_instancenorm(dimension=dimension)
|
| 230 |
+
|
| 231 |
+
config['norm_op_kwargs'] = None # this will use defaults
|
| 232 |
+
|
| 233 |
+
if nonlin == "LeakyReLU":
|
| 234 |
+
config['nonlin'] = nn.LeakyReLU
|
| 235 |
+
config['nonlin_kwargs'] = {'negative_slope': 1e-2, 'inplace': True}
|
| 236 |
+
elif nonlin == "ReLU":
|
| 237 |
+
config['nonlin'] = nn.ReLU
|
| 238 |
+
config['nonlin_kwargs'] = {'inplace': True}
|
| 239 |
+
else:
|
| 240 |
+
raise NotImplementedError('Unknown nonlin %s. Only "LeakyReLU" and "ReLU" are supported for now' % nonlin)
|
| 241 |
+
|
| 242 |
+
return config
|
RADAR_inference/dynamic_network_architectures/building_blocks/plain_conv_encoder.py
ADDED
|
@@ -0,0 +1,105 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from torch import nn
|
| 3 |
+
import numpy as np
|
| 4 |
+
from typing import Union, Type, List, Tuple
|
| 5 |
+
|
| 6 |
+
from torch.nn.modules.conv import _ConvNd
|
| 7 |
+
from torch.nn.modules.dropout import _DropoutNd
|
| 8 |
+
from dynamic_network_architectures.building_blocks.simple_conv_blocks import StackedConvBlocks
|
| 9 |
+
from dynamic_network_architectures.building_blocks.helper import maybe_convert_scalar_to_list, get_matching_pool_op
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
class PlainConvEncoder(nn.Module):
|
| 13 |
+
def __init__(self,
|
| 14 |
+
input_channels: int,
|
| 15 |
+
n_stages: int,
|
| 16 |
+
features_per_stage: Union[int, List[int], Tuple[int, ...]],
|
| 17 |
+
conv_op: Type[_ConvNd],
|
| 18 |
+
kernel_sizes: Union[int, List[int], Tuple[int, ...]],
|
| 19 |
+
strides: Union[int, List[int], Tuple[int, ...]],
|
| 20 |
+
n_conv_per_stage: Union[int, List[int], Tuple[int, ...]],
|
| 21 |
+
conv_bias: bool = False,
|
| 22 |
+
norm_op: Union[None, Type[nn.Module]] = None,
|
| 23 |
+
norm_op_kwargs: dict = None,
|
| 24 |
+
dropout_op: Union[None, Type[_DropoutNd]] = None,
|
| 25 |
+
dropout_op_kwargs: dict = None,
|
| 26 |
+
nonlin: Union[None, Type[torch.nn.Module]] = None,
|
| 27 |
+
nonlin_kwargs: dict = None,
|
| 28 |
+
return_skips: bool = False,
|
| 29 |
+
nonlin_first: bool = False,
|
| 30 |
+
pool: str = 'conv'
|
| 31 |
+
):
|
| 32 |
+
|
| 33 |
+
super().__init__()
|
| 34 |
+
if isinstance(kernel_sizes, int):
|
| 35 |
+
kernel_sizes = [kernel_sizes] * n_stages
|
| 36 |
+
if isinstance(features_per_stage, int):
|
| 37 |
+
features_per_stage = [features_per_stage] * n_stages
|
| 38 |
+
if isinstance(n_conv_per_stage, int):
|
| 39 |
+
n_conv_per_stage = [n_conv_per_stage] * n_stages
|
| 40 |
+
if isinstance(strides, int):
|
| 41 |
+
strides = [strides] * n_stages
|
| 42 |
+
assert len(kernel_sizes) == n_stages, "kernel_sizes must have as many entries as we have resolution stages (n_stages)"
|
| 43 |
+
assert len(n_conv_per_stage) == n_stages, "n_conv_per_stage must have as many entries as we have resolution stages (n_stages)"
|
| 44 |
+
assert len(features_per_stage) == n_stages, "features_per_stage must have as many entries as we have resolution stages (n_stages)"
|
| 45 |
+
assert len(strides) == n_stages, "strides must have as many entries as we have resolution stages (n_stages). " \
|
| 46 |
+
"Important: first entry is recommended to be 1, else we run strided conv drectly on the input"
|
| 47 |
+
|
| 48 |
+
stages = []
|
| 49 |
+
for s in range(n_stages):
|
| 50 |
+
stage_modules = []
|
| 51 |
+
if pool == 'max' or pool == 'avg':
|
| 52 |
+
if (isinstance(strides[s], int) and strides[s] != 1) or \
|
| 53 |
+
isinstance(strides[s], (tuple, list)) and any([i != 1 for i in strides[s]]):
|
| 54 |
+
stage_modules.append(get_matching_pool_op(conv_op, pool_type=pool)(kernel_size=strides[s], stride=strides[s]))
|
| 55 |
+
conv_stride = 1
|
| 56 |
+
elif pool == 'conv':
|
| 57 |
+
conv_stride = strides[s]
|
| 58 |
+
else:
|
| 59 |
+
raise RuntimeError()
|
| 60 |
+
stage_modules.append(StackedConvBlocks(
|
| 61 |
+
n_conv_per_stage[s], conv_op, input_channels, features_per_stage[s], kernel_sizes[s], conv_stride,
|
| 62 |
+
conv_bias, norm_op, norm_op_kwargs, dropout_op, dropout_op_kwargs, nonlin, nonlin_kwargs, nonlin_first
|
| 63 |
+
))
|
| 64 |
+
stages.append(nn.Sequential(*stage_modules))
|
| 65 |
+
input_channels = features_per_stage[s]
|
| 66 |
+
|
| 67 |
+
self.stages = nn.Sequential(*stages)
|
| 68 |
+
self.output_channels = features_per_stage
|
| 69 |
+
self.strides = [maybe_convert_scalar_to_list(conv_op, i) for i in strides]
|
| 70 |
+
self.return_skips = return_skips
|
| 71 |
+
|
| 72 |
+
# we store some things that a potential decoder needs
|
| 73 |
+
self.conv_op = conv_op
|
| 74 |
+
self.norm_op = norm_op
|
| 75 |
+
self.norm_op_kwargs = norm_op_kwargs
|
| 76 |
+
self.nonlin = nonlin
|
| 77 |
+
self.nonlin_kwargs = nonlin_kwargs
|
| 78 |
+
self.dropout_op = dropout_op
|
| 79 |
+
self.dropout_op_kwargs = dropout_op_kwargs
|
| 80 |
+
self.conv_bias = conv_bias
|
| 81 |
+
self.kernel_sizes = kernel_sizes
|
| 82 |
+
|
| 83 |
+
def forward(self, x):
|
| 84 |
+
ret = []
|
| 85 |
+
for s in self.stages:
|
| 86 |
+
x = s(x)
|
| 87 |
+
ret.append(x)
|
| 88 |
+
if self.return_skips:
|
| 89 |
+
return ret
|
| 90 |
+
else:
|
| 91 |
+
return ret[-1]
|
| 92 |
+
|
| 93 |
+
def compute_conv_feature_map_size(self, input_size):
|
| 94 |
+
output = np.int64(0)
|
| 95 |
+
for s in range(len(self.stages)):
|
| 96 |
+
if isinstance(self.stages[s], nn.Sequential):
|
| 97 |
+
for sq in self.stages[s]:
|
| 98 |
+
if hasattr(sq, 'compute_conv_feature_map_size'):
|
| 99 |
+
output += self.stages[s][-1].compute_conv_feature_map_size(input_size)
|
| 100 |
+
else:
|
| 101 |
+
output += self.stages[s].compute_conv_feature_map_size(input_size)
|
| 102 |
+
input_size = [i // j for i, j in zip(input_size, self.strides[s])]
|
| 103 |
+
return output
|
| 104 |
+
|
| 105 |
+
|
RADAR_inference/dynamic_network_architectures/building_blocks/regularization.py
ADDED
|
@@ -0,0 +1,86 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from torch import nn
|
| 2 |
+
|
| 3 |
+
|
| 4 |
+
def drop_path(x, drop_prob: float = 0., training: bool = False, scale_by_keep: bool = True):
|
| 5 |
+
"""
|
| 6 |
+
This function is taken from the timm package (https://github.com/rwightman/pytorch-image-models/blob/master/timm/models/layers/drop.py).
|
| 7 |
+
|
| 8 |
+
Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
|
| 9 |
+
This is the same as the DropConnect impl I created for EfficientNet, etc networks, however,
|
| 10 |
+
the original name is misleading as 'Drop Connect' is a different form of dropout in a separate paper...
|
| 11 |
+
See discussion: https://github.com/tensorflow/tpu/issues/494#issuecomment-532968956 ... I've opted for
|
| 12 |
+
changing the layer and argument names to 'drop path' rather than mix DropConnect as a layer name and use
|
| 13 |
+
'survival rate' as the argument.
|
| 14 |
+
"""
|
| 15 |
+
if drop_prob == 0. or not training:
|
| 16 |
+
return x
|
| 17 |
+
keep_prob = 1 - drop_prob
|
| 18 |
+
shape = (x.shape[0],) + (1,) * (x.ndim - 1) # work with diff dim tensors, not just 2D ConvNets
|
| 19 |
+
random_tensor = x.new_empty(shape).bernoulli_(keep_prob)
|
| 20 |
+
if keep_prob > 0.0 and scale_by_keep:
|
| 21 |
+
random_tensor.div_(keep_prob)
|
| 22 |
+
return x * random_tensor
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
class DropPath(nn.Module):
|
| 26 |
+
"""
|
| 27 |
+
This class is taken from the timm package (https://github.com/rwightman/pytorch-image-models/blob/master/timm/models/layers/drop.py).
|
| 28 |
+
|
| 29 |
+
Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
|
| 30 |
+
"""
|
| 31 |
+
def __init__(self, drop_prob: float = 0., scale_by_keep: bool = True):
|
| 32 |
+
super(DropPath, self).__init__()
|
| 33 |
+
self.drop_prob = drop_prob
|
| 34 |
+
self.scale_by_keep = scale_by_keep
|
| 35 |
+
|
| 36 |
+
def forward(self, x):
|
| 37 |
+
return drop_path(x, self.drop_prob, self.training, self.scale_by_keep)
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
class SqueezeExcite(nn.Module):
|
| 41 |
+
"""
|
| 42 |
+
This class is taken from the timm package (https://github.com/rwightman/pytorch-image-models/blob/master/timm/models/layers/squeeze_excite.py)
|
| 43 |
+
and slightly modified so that the convolution type can be adapted.
|
| 44 |
+
|
| 45 |
+
SE Module as defined in original SE-Nets with a few additions
|
| 46 |
+
Additions include:
|
| 47 |
+
* divisor can be specified to keep channels % div == 0 (default: 8)
|
| 48 |
+
* reduction channels can be specified directly by arg (if rd_channels is set)
|
| 49 |
+
* reduction channels can be specified by float rd_ratio (default: 1/16)
|
| 50 |
+
* global max pooling can be added to the squeeze aggregation
|
| 51 |
+
* customizable activation, normalization, and gate layer
|
| 52 |
+
"""
|
| 53 |
+
def __init__(
|
| 54 |
+
self, channels, conv_op, rd_ratio=1. / 16, rd_channels=None, rd_divisor=8, add_maxpool=False,
|
| 55 |
+
act_layer=nn.ReLU, norm_layer=None, gate_layer=nn.Sigmoid):
|
| 56 |
+
super(SqueezeExcite, self).__init__()
|
| 57 |
+
self.add_maxpool = add_maxpool
|
| 58 |
+
if not rd_channels:
|
| 59 |
+
rd_channels = make_divisible(channels * rd_ratio, rd_divisor, round_limit=0.)
|
| 60 |
+
self.fc1 = conv_op(channels, rd_channels, kernel_size=1, bias=True)
|
| 61 |
+
self.bn = norm_layer(rd_channels) if norm_layer else nn.Identity()
|
| 62 |
+
self.act = act_layer(inplace=True)
|
| 63 |
+
self.fc2 = conv_op(rd_channels, channels, kernel_size=1, bias=True)
|
| 64 |
+
self.gate = gate_layer()
|
| 65 |
+
|
| 66 |
+
def forward(self, x):
|
| 67 |
+
x_se = x.mean((2, 3), keepdim=True)
|
| 68 |
+
if self.add_maxpool:
|
| 69 |
+
# experimental codepath, may remove or change
|
| 70 |
+
x_se = 0.5 * x_se + 0.5 * x.amax((2, 3), keepdim=True)
|
| 71 |
+
x_se = self.fc1(x_se)
|
| 72 |
+
x_se = self.act(self.bn(x_se))
|
| 73 |
+
x_se = self.fc2(x_se)
|
| 74 |
+
return x * self.gate(x_se)
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
def make_divisible(v, divisor=8, min_value=None, round_limit=.9):
|
| 78 |
+
"""
|
| 79 |
+
This function is taken from the timm package (https://github.com/rwightman/pytorch-image-models/blob/b7cb8d0337b3e7b50516849805ddb9be5fc11644/timm/models/layers/helpers.py#L25)
|
| 80 |
+
"""
|
| 81 |
+
min_value = min_value or divisor
|
| 82 |
+
new_v = max(min_value, int(v + divisor / 2) // divisor * divisor)
|
| 83 |
+
# Make sure that round down does not go down by more than 10%.
|
| 84 |
+
if new_v < round_limit * v:
|
| 85 |
+
new_v += divisor
|
| 86 |
+
return new_v
|
RADAR_inference/dynamic_network_architectures/building_blocks/residual.py
ADDED
|
@@ -0,0 +1,371 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import Tuple, List, Union, Type
|
| 2 |
+
import torch.nn
|
| 3 |
+
from torch import nn
|
| 4 |
+
from torch.nn.modules.conv import _ConvNd
|
| 5 |
+
from torch.nn.modules.dropout import _DropoutNd
|
| 6 |
+
|
| 7 |
+
from dynamic_network_architectures.building_blocks.helper import maybe_convert_scalar_to_list, get_matching_pool_op
|
| 8 |
+
from dynamic_network_architectures.building_blocks.simple_conv_blocks import ConvDropoutNormReLU
|
| 9 |
+
from dynamic_network_architectures.building_blocks.regularization import DropPath, SqueezeExcite
|
| 10 |
+
import numpy as np
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
class BasicBlockD(nn.Module):
|
| 14 |
+
def __init__(self,
|
| 15 |
+
conv_op: Type[_ConvNd],
|
| 16 |
+
input_channels: int,
|
| 17 |
+
output_channels: int,
|
| 18 |
+
kernel_size: Union[int, List[int], Tuple[int, ...]],
|
| 19 |
+
stride: Union[int, List[int], Tuple[int, ...]],
|
| 20 |
+
conv_bias: bool = False,
|
| 21 |
+
norm_op: Union[None, Type[nn.Module]] = None,
|
| 22 |
+
norm_op_kwargs: dict = None,
|
| 23 |
+
dropout_op: Union[None, Type[_DropoutNd]] = None,
|
| 24 |
+
dropout_op_kwargs: dict = None,
|
| 25 |
+
nonlin: Union[None, Type[torch.nn.Module]] = None,
|
| 26 |
+
nonlin_kwargs: dict = None,
|
| 27 |
+
stochastic_depth_p: float = 0.0,
|
| 28 |
+
squeeze_excitation: bool = False,
|
| 29 |
+
squeeze_excitation_reduction_ratio: float = 1. / 16,
|
| 30 |
+
# todo wideresnet?
|
| 31 |
+
):
|
| 32 |
+
"""
|
| 33 |
+
This implementation follows ResNet-D:
|
| 34 |
+
|
| 35 |
+
He, Tong, et al. "Bag of tricks for image classification with convolutional neural networks."
|
| 36 |
+
Proceedings of the IEEE/CVF Conference on Computer Vision and Pattern Recognition. 2019.
|
| 37 |
+
|
| 38 |
+
The skip has an avgpool (if needed) followed by 1x1 conv instead of just a strided 1x1 conv
|
| 39 |
+
|
| 40 |
+
:param conv_op:
|
| 41 |
+
:param input_channels:
|
| 42 |
+
:param output_channels:
|
| 43 |
+
:param kernel_size: refers only to convs in feature extraction path, not to 1x1x1 conv in skip
|
| 44 |
+
:param stride: only applies to first conv (and skip). Second conv always has stride 1
|
| 45 |
+
:param conv_bias:
|
| 46 |
+
:param norm_op:
|
| 47 |
+
:param norm_op_kwargs:
|
| 48 |
+
:param dropout_op: only the first conv can have dropout. The second never has
|
| 49 |
+
:param dropout_op_kwargs:
|
| 50 |
+
:param nonlin:
|
| 51 |
+
:param nonlin_kwargs:
|
| 52 |
+
:param stochastic_depth_p:
|
| 53 |
+
:param squeeze_excitation:
|
| 54 |
+
:param squeeze_excitation_reduction_ratio:
|
| 55 |
+
"""
|
| 56 |
+
super().__init__()
|
| 57 |
+
self.input_channels = input_channels
|
| 58 |
+
self.output_channels = output_channels
|
| 59 |
+
stride = maybe_convert_scalar_to_list(conv_op, stride)
|
| 60 |
+
self.stride = stride
|
| 61 |
+
|
| 62 |
+
kernel_size = maybe_convert_scalar_to_list(conv_op, kernel_size)
|
| 63 |
+
|
| 64 |
+
if norm_op_kwargs is None:
|
| 65 |
+
norm_op_kwargs = {}
|
| 66 |
+
if nonlin_kwargs is None:
|
| 67 |
+
nonlin_kwargs = {}
|
| 68 |
+
|
| 69 |
+
self.conv1 = ConvDropoutNormReLU(conv_op, input_channels, output_channels, kernel_size, stride, conv_bias,
|
| 70 |
+
norm_op, norm_op_kwargs, dropout_op, dropout_op_kwargs, nonlin, nonlin_kwargs)
|
| 71 |
+
self.conv2 = ConvDropoutNormReLU(conv_op, output_channels, output_channels, kernel_size, 1, conv_bias, norm_op,
|
| 72 |
+
norm_op_kwargs, None, None, None, None)
|
| 73 |
+
|
| 74 |
+
self.nonlin2 = nonlin(**nonlin_kwargs) if nonlin is not None else lambda x: x
|
| 75 |
+
|
| 76 |
+
# Stochastic Depth
|
| 77 |
+
self.apply_stochastic_depth = False if stochastic_depth_p == 0.0 else True
|
| 78 |
+
if self.apply_stochastic_depth:
|
| 79 |
+
self.drop_path = DropPath(drop_prob=stochastic_depth_p)
|
| 80 |
+
|
| 81 |
+
# Squeeze Excitation
|
| 82 |
+
self.apply_se = squeeze_excitation
|
| 83 |
+
if self.apply_se:
|
| 84 |
+
self.squeeze_excitation = SqueezeExcite(self.output_channels, conv_op,
|
| 85 |
+
rd_ratio=squeeze_excitation_reduction_ratio, rd_divisor=8)
|
| 86 |
+
|
| 87 |
+
has_stride = (isinstance(stride, int) and stride != 1) or any([i != 1 for i in stride])
|
| 88 |
+
requires_projection = (input_channels != output_channels)
|
| 89 |
+
|
| 90 |
+
if has_stride or requires_projection:
|
| 91 |
+
ops = []
|
| 92 |
+
if has_stride:
|
| 93 |
+
ops.append(get_matching_pool_op(conv_op=conv_op, adaptive=False, pool_type='avg')(stride, stride))
|
| 94 |
+
if requires_projection:
|
| 95 |
+
ops.append(
|
| 96 |
+
ConvDropoutNormReLU(conv_op, input_channels, output_channels, 1, 1, False, norm_op,
|
| 97 |
+
norm_op_kwargs, None, None, None, None
|
| 98 |
+
)
|
| 99 |
+
)
|
| 100 |
+
self.skip = nn.Sequential(*ops)
|
| 101 |
+
else:
|
| 102 |
+
self.skip = lambda x: x
|
| 103 |
+
|
| 104 |
+
def forward(self, x):
|
| 105 |
+
residual = self.skip(x)
|
| 106 |
+
out = self.conv2(self.conv1(x))
|
| 107 |
+
if self.apply_stochastic_depth:
|
| 108 |
+
out = self.drop_path(out)
|
| 109 |
+
if self.apply_se:
|
| 110 |
+
out = self.squeeze_excitation(out)
|
| 111 |
+
out += residual
|
| 112 |
+
return self.nonlin2(out)
|
| 113 |
+
|
| 114 |
+
def compute_conv_feature_map_size(self, input_size):
|
| 115 |
+
assert len(input_size) == len(self.stride), "just give the image size without color/feature channels or " \
|
| 116 |
+
"batch channel. Do not give input_size=(b, c, x, y(, z)). " \
|
| 117 |
+
"Give input_size=(x, y(, z))!"
|
| 118 |
+
size_after_stride = [i // j for i, j in zip(input_size, self.stride)]
|
| 119 |
+
# conv1
|
| 120 |
+
output_size_conv1 = np.prod([self.output_channels, *size_after_stride], dtype=np.int64)
|
| 121 |
+
# conv2
|
| 122 |
+
output_size_conv2 = np.prod([self.output_channels, *size_after_stride], dtype=np.int64)
|
| 123 |
+
# skip conv (if applicable)
|
| 124 |
+
if (self.input_channels != self.output_channels) or any([i != j for i, j in zip(input_size, size_after_stride)]):
|
| 125 |
+
assert isinstance(self.skip, nn.Sequential)
|
| 126 |
+
output_size_skip = np.prod([self.output_channels, *size_after_stride], dtype=np.int64)
|
| 127 |
+
else:
|
| 128 |
+
assert not isinstance(self.skip, nn.Sequential)
|
| 129 |
+
output_size_skip = 0
|
| 130 |
+
return output_size_conv1 + output_size_conv2 + output_size_skip
|
| 131 |
+
|
| 132 |
+
|
| 133 |
+
class BottleneckD(nn.Module):
|
| 134 |
+
def __init__(self,
|
| 135 |
+
conv_op: Type[_ConvNd],
|
| 136 |
+
input_channels: int,
|
| 137 |
+
bottleneck_channels: int,
|
| 138 |
+
output_channels: int,
|
| 139 |
+
kernel_size: Union[int, List[int], Tuple[int, ...]],
|
| 140 |
+
stride: Union[int, List[int], Tuple[int, ...]],
|
| 141 |
+
conv_bias: bool = False,
|
| 142 |
+
norm_op: Union[None, Type[nn.Module]] = None,
|
| 143 |
+
norm_op_kwargs: dict = None,
|
| 144 |
+
dropout_op: Union[None, Type[_DropoutNd]] = None,
|
| 145 |
+
dropout_op_kwargs: dict = None,
|
| 146 |
+
nonlin: Union[None, Type[torch.nn.Module]] = None,
|
| 147 |
+
nonlin_kwargs: dict = None,
|
| 148 |
+
stochastic_depth_p: float = 0.0,
|
| 149 |
+
squeeze_excitation: bool = False,
|
| 150 |
+
squeeze_excitation_reduction_ratio: float = 1. / 16
|
| 151 |
+
):
|
| 152 |
+
"""
|
| 153 |
+
This implementation follows ResNet-D:
|
| 154 |
+
|
| 155 |
+
He, Tong, et al. "Bag of tricks for image classification with convolutional neural networks."
|
| 156 |
+
Proceedings of the IEEE/CVF Conference on Computer Vision and Pattern Recognition. 2019.
|
| 157 |
+
|
| 158 |
+
The stride sits in the 3x3 conv instead of the 1x1 conv!
|
| 159 |
+
The skip has an avgpool (if needed) followed by 1x1 conv instead of just a strided 1x1 conv
|
| 160 |
+
|
| 161 |
+
:param conv_op:
|
| 162 |
+
:param input_channels:
|
| 163 |
+
:param output_channels:
|
| 164 |
+
:param kernel_size: only affects the conv in the middle (typically 3x3). The other convs remain 1x1
|
| 165 |
+
:param stride: only applies to the conv in the middle (and skip). Note that this deviates from the canonical
|
| 166 |
+
ResNet implementation where the stride is applied to the first 1x1 conv. (This implementation follows ResNet-D)
|
| 167 |
+
:param conv_bias:
|
| 168 |
+
:param norm_op:
|
| 169 |
+
:param norm_op_kwargs:
|
| 170 |
+
:param dropout_op: only the second (kernel_size) conv can have dropout. The first and last conv (1x1(x1)) never have it
|
| 171 |
+
:param dropout_op_kwargs:
|
| 172 |
+
:param nonlin:
|
| 173 |
+
:param nonlin_kwargs:
|
| 174 |
+
:param stochastic_depth_p:
|
| 175 |
+
:param squeeze_excitation:
|
| 176 |
+
:param squeeze_excitation_reduction_ratio:
|
| 177 |
+
"""
|
| 178 |
+
super().__init__()
|
| 179 |
+
self.input_channels = input_channels
|
| 180 |
+
self.output_channels = output_channels
|
| 181 |
+
self.bottleneck_channels = bottleneck_channels
|
| 182 |
+
stride = maybe_convert_scalar_to_list(conv_op, stride)
|
| 183 |
+
self.stride = stride
|
| 184 |
+
|
| 185 |
+
kernel_size = maybe_convert_scalar_to_list(conv_op, kernel_size)
|
| 186 |
+
if norm_op_kwargs is None:
|
| 187 |
+
norm_op_kwargs = {}
|
| 188 |
+
if nonlin_kwargs is None:
|
| 189 |
+
nonlin_kwargs = {}
|
| 190 |
+
|
| 191 |
+
self.conv1 = ConvDropoutNormReLU(conv_op, input_channels, bottleneck_channels, 1, 1, conv_bias,
|
| 192 |
+
norm_op, norm_op_kwargs, None, None, nonlin, nonlin_kwargs)
|
| 193 |
+
self.conv2 = ConvDropoutNormReLU(conv_op, bottleneck_channels, bottleneck_channels, kernel_size, stride,
|
| 194 |
+
conv_bias,
|
| 195 |
+
norm_op, norm_op_kwargs, dropout_op, dropout_op_kwargs, nonlin, nonlin_kwargs)
|
| 196 |
+
self.conv3 = ConvDropoutNormReLU(conv_op, bottleneck_channels, output_channels, 1, 1, conv_bias, norm_op,
|
| 197 |
+
norm_op_kwargs, None, None, None, None)
|
| 198 |
+
|
| 199 |
+
self.nonlin3 = nonlin(**nonlin_kwargs) if nonlin is not None else lambda x: x
|
| 200 |
+
|
| 201 |
+
# Stochastic Depth
|
| 202 |
+
self.apply_stochastic_depth = False if stochastic_depth_p == 0.0 else True
|
| 203 |
+
if self.apply_stochastic_depth:
|
| 204 |
+
self.drop_path = DropPath(drop_prob=stochastic_depth_p)
|
| 205 |
+
|
| 206 |
+
# Squeeze Excitation
|
| 207 |
+
self.apply_se = squeeze_excitation
|
| 208 |
+
if self.apply_se:
|
| 209 |
+
self.squeeze_excitation = SqueezeExcite(self.output_channels, conv_op,
|
| 210 |
+
rd_ratio=squeeze_excitation_reduction_ratio, rd_divisor=8)
|
| 211 |
+
|
| 212 |
+
has_stride = (isinstance(stride, int) and stride != 1) or any([i != 1 for i in stride])
|
| 213 |
+
requires_projection = (input_channels != output_channels)
|
| 214 |
+
|
| 215 |
+
if has_stride or requires_projection:
|
| 216 |
+
ops = []
|
| 217 |
+
if has_stride:
|
| 218 |
+
ops.append(get_matching_pool_op(conv_op=conv_op, adaptive=False, pool_type='avg')(stride, stride))
|
| 219 |
+
if requires_projection:
|
| 220 |
+
ops.append(
|
| 221 |
+
ConvDropoutNormReLU(conv_op, input_channels, output_channels, 1, 1, False,
|
| 222 |
+
norm_op, norm_op_kwargs, None, None, None, None
|
| 223 |
+
)
|
| 224 |
+
)
|
| 225 |
+
self.skip = nn.Sequential(*ops)
|
| 226 |
+
else:
|
| 227 |
+
self.skip = lambda x: x
|
| 228 |
+
|
| 229 |
+
def forward(self, x):
|
| 230 |
+
residual = self.skip(x)
|
| 231 |
+
out = self.conv3(self.conv2(self.conv1(x)))
|
| 232 |
+
if self.apply_stochastic_depth:
|
| 233 |
+
out = self.drop_path(out)
|
| 234 |
+
if self.apply_se:
|
| 235 |
+
out = self.squeeze_excitation(out)
|
| 236 |
+
out += residual
|
| 237 |
+
return self.nonlin3(out)
|
| 238 |
+
|
| 239 |
+
def compute_conv_feature_map_size(self, input_size):
|
| 240 |
+
assert len(input_size) == len(self.stride), "just give the image size without color/feature channels or " \
|
| 241 |
+
"batch channel. Do not give input_size=(b, c, x, y(, z)). " \
|
| 242 |
+
"Give input_size=(x, y(, z))!"
|
| 243 |
+
size_after_stride = [i // j for i, j in zip(input_size, self.stride)]
|
| 244 |
+
# conv1
|
| 245 |
+
output_size_conv1 = np.prod([self.bottleneck_channels, *input_size], dtype=np.int64)
|
| 246 |
+
# conv2
|
| 247 |
+
output_size_conv2 = np.prod([self.bottleneck_channels, *size_after_stride], dtype=np.int64)
|
| 248 |
+
# conv3
|
| 249 |
+
output_size_conv3 = np.prod([self.output_channels, *size_after_stride], dtype=np.int64)
|
| 250 |
+
# skip conv (if applicable)
|
| 251 |
+
if (self.input_channels != self.output_channels) or any([i != j for i, j in zip(input_size, size_after_stride)]):
|
| 252 |
+
assert isinstance(self.skip, nn.Sequential)
|
| 253 |
+
output_size_skip = np.prod([self.output_channels, *size_after_stride], dtype=np.int64)
|
| 254 |
+
else:
|
| 255 |
+
assert not isinstance(self.skip, nn.Sequential)
|
| 256 |
+
output_size_skip = 0
|
| 257 |
+
return output_size_conv1 + output_size_conv2 + output_size_conv3 + output_size_skip
|
| 258 |
+
|
| 259 |
+
|
| 260 |
+
class StackedResidualBlocks(nn.Module):
|
| 261 |
+
def __init__(self,
|
| 262 |
+
n_blocks: int,
|
| 263 |
+
conv_op: Type[_ConvNd],
|
| 264 |
+
input_channels: int,
|
| 265 |
+
output_channels: Union[int, List[int], Tuple[int, ...]],
|
| 266 |
+
kernel_size: Union[int, List[int], Tuple[int, ...]],
|
| 267 |
+
initial_stride: Union[int, List[int], Tuple[int, ...]],
|
| 268 |
+
conv_bias: bool = False,
|
| 269 |
+
norm_op: Union[None, Type[nn.Module]] = None,
|
| 270 |
+
norm_op_kwargs: dict = None,
|
| 271 |
+
dropout_op: Union[None, Type[_DropoutNd]] = None,
|
| 272 |
+
dropout_op_kwargs: dict = None,
|
| 273 |
+
nonlin: Union[None, Type[torch.nn.Module]] = None,
|
| 274 |
+
nonlin_kwargs: dict = None,
|
| 275 |
+
block: Union[Type[BasicBlockD], Type[BottleneckD]] = BasicBlockD,
|
| 276 |
+
bottleneck_channels: Union[int, List[int], Tuple[int, ...]] = None,
|
| 277 |
+
stochastic_depth_p: float = 0.0,
|
| 278 |
+
squeeze_excitation: bool = False,
|
| 279 |
+
squeeze_excitation_reduction_ratio: float = 1. / 16
|
| 280 |
+
):
|
| 281 |
+
"""
|
| 282 |
+
Stack multiple instances of block.
|
| 283 |
+
|
| 284 |
+
:param n_blocks: number of residual blocks
|
| 285 |
+
:param conv_op: nn.ConvNd class
|
| 286 |
+
:param input_channels: only relevant for forst block in the sequence. This is the input number of features.
|
| 287 |
+
After the first block, the number of features in the main path to which the residuals are added is output_channels
|
| 288 |
+
:param output_channels: number of features in the main path to which the residuals are added (and also the
|
| 289 |
+
number of features of the output)
|
| 290 |
+
:param kernel_size: kernel size for all nxn (n!=1) convolutions. Default: 3x3
|
| 291 |
+
:param initial_stride: only affects the first block. All subsequent blocks have stride 1
|
| 292 |
+
:param conv_bias: usually False
|
| 293 |
+
:param norm_op: nn.BatchNormNd, InstanceNormNd etc
|
| 294 |
+
:param norm_op_kwargs: dictionary of kwargs. Leave empty ({}) for defaults
|
| 295 |
+
:param dropout_op: nn.DropoutNd, can be None for no dropout
|
| 296 |
+
:param dropout_op_kwargs:
|
| 297 |
+
:param nonlin:
|
| 298 |
+
:param nonlin_kwargs:
|
| 299 |
+
:param block: BasicBlockD or BottleneckD
|
| 300 |
+
:param bottleneck_channels: if block is BottleneckD then we need to know the number of bottleneck features.
|
| 301 |
+
Bottleneck will use first 1x1 conv to reduce input to bottleneck features, then run the nxn (see kernel_size)
|
| 302 |
+
conv on that (bottleneck -> bottleneck). Finally the output will be projected back to output_channels
|
| 303 |
+
(bottleneck -> output_channels) with the final 1x1 conv
|
| 304 |
+
:param stochastic_depth_p: probability of applying stochastic depth in residual blocks
|
| 305 |
+
:param squeeze_excitation: whether to apply squeeze and excitation or not
|
| 306 |
+
:param squeeze_excitation_reduction_ratio: ratio by how much squeeze and excitation should reduce channels
|
| 307 |
+
respective to number of out channels of respective block
|
| 308 |
+
"""
|
| 309 |
+
super().__init__()
|
| 310 |
+
assert n_blocks > 0, 'n_blocks must be > 0'
|
| 311 |
+
assert block in [BasicBlockD, BottleneckD], 'block must be BasicBlockD or BottleneckD'
|
| 312 |
+
if not isinstance(output_channels, (tuple, list)):
|
| 313 |
+
output_channels = [output_channels] * n_blocks
|
| 314 |
+
if not isinstance(bottleneck_channels, (tuple, list)):
|
| 315 |
+
bottleneck_channels = [bottleneck_channels] * n_blocks
|
| 316 |
+
|
| 317 |
+
if block == BasicBlockD:
|
| 318 |
+
blocks = nn.Sequential(
|
| 319 |
+
block(conv_op, input_channels, output_channels[0], kernel_size, initial_stride, conv_bias,
|
| 320 |
+
norm_op, norm_op_kwargs, dropout_op, dropout_op_kwargs, nonlin, nonlin_kwargs, stochastic_depth_p,
|
| 321 |
+
squeeze_excitation, squeeze_excitation_reduction_ratio),
|
| 322 |
+
*[block(conv_op, output_channels[n - 1], output_channels[n], kernel_size, 1, conv_bias, norm_op,
|
| 323 |
+
norm_op_kwargs, dropout_op, dropout_op_kwargs, nonlin, nonlin_kwargs, stochastic_depth_p,
|
| 324 |
+
squeeze_excitation, squeeze_excitation_reduction_ratio) for n in range(1, n_blocks)]
|
| 325 |
+
)
|
| 326 |
+
else:
|
| 327 |
+
blocks = nn.Sequential(
|
| 328 |
+
block(conv_op, input_channels, bottleneck_channels[0], output_channels[0], kernel_size,
|
| 329 |
+
initial_stride, conv_bias, norm_op, norm_op_kwargs, dropout_op, dropout_op_kwargs,
|
| 330 |
+
nonlin, nonlin_kwargs, stochastic_depth_p, squeeze_excitation, squeeze_excitation_reduction_ratio),
|
| 331 |
+
*[block(conv_op, output_channels[n - 1], bottleneck_channels[n], output_channels[n], kernel_size,
|
| 332 |
+
1, conv_bias, norm_op, norm_op_kwargs, dropout_op, dropout_op_kwargs,
|
| 333 |
+
nonlin, nonlin_kwargs, stochastic_depth_p, squeeze_excitation,
|
| 334 |
+
squeeze_excitation_reduction_ratio) for n in range(1, n_blocks)]
|
| 335 |
+
)
|
| 336 |
+
self.blocks = blocks
|
| 337 |
+
self.initial_stride = maybe_convert_scalar_to_list(conv_op, initial_stride)
|
| 338 |
+
self.output_channels = output_channels[-1]
|
| 339 |
+
|
| 340 |
+
def forward(self, x):
|
| 341 |
+
return self.blocks(x)
|
| 342 |
+
|
| 343 |
+
def compute_conv_feature_map_size(self, input_size):
|
| 344 |
+
assert len(input_size) == len(self.initial_stride), "just give the image size without color/feature channels or " \
|
| 345 |
+
"batch channel. Do not give input_size=(b, c, x, y(, z)). " \
|
| 346 |
+
"Give input_size=(x, y(, z))!"
|
| 347 |
+
output = self.blocks[0].compute_conv_feature_map_size(input_size)
|
| 348 |
+
size_after_stride = [i // j for i, j in zip(input_size, self.initial_stride)]
|
| 349 |
+
for b in self.blocks[1:]:
|
| 350 |
+
output += b.compute_conv_feature_map_size(size_after_stride)
|
| 351 |
+
return output
|
| 352 |
+
|
| 353 |
+
|
| 354 |
+
if __name__ == '__main__':
|
| 355 |
+
data = torch.rand((1, 3, 40, 32))
|
| 356 |
+
|
| 357 |
+
stx = StackedResidualBlocks(2, nn.Conv2d, 24, (16, 16), (3, 3), (1, 2),
|
| 358 |
+
norm_op=nn.BatchNorm2d, nonlin=nn.ReLU, nonlin_kwargs={'inplace': True},
|
| 359 |
+
block=BottleneckD, bottleneck_channels=3)
|
| 360 |
+
model = nn.Sequential(ConvDropoutNormReLU(nn.Conv2d,
|
| 361 |
+
3, 24, 3, 1, True, nn.BatchNorm2d, {}, None, None, nn.LeakyReLU,
|
| 362 |
+
{'inplace': True}),
|
| 363 |
+
stx)
|
| 364 |
+
import hiddenlayer as hl
|
| 365 |
+
|
| 366 |
+
g = hl.build_graph(model, data,
|
| 367 |
+
transforms=None)
|
| 368 |
+
g.save("network_architecture.pdf")
|
| 369 |
+
del g
|
| 370 |
+
|
| 371 |
+
print(stx.compute_conv_feature_map_size((40, 32)))
|
RADAR_inference/dynamic_network_architectures/building_blocks/residual_encoders.py
ADDED
|
@@ -0,0 +1,172 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from torch import nn
|
| 3 |
+
import numpy as np
|
| 4 |
+
from typing import Union, Type, List, Tuple
|
| 5 |
+
|
| 6 |
+
from torch.nn.modules.conv import _ConvNd
|
| 7 |
+
from torch.nn.modules.dropout import _DropoutNd
|
| 8 |
+
from dynamic_network_architectures.building_blocks.residual import StackedResidualBlocks, BottleneckD, BasicBlockD
|
| 9 |
+
from dynamic_network_architectures.building_blocks.helper import maybe_convert_scalar_to_list, get_matching_pool_op
|
| 10 |
+
from dynamic_network_architectures.building_blocks.simple_conv_blocks import StackedConvBlocks
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
class ResidualEncoder(nn.Module):
|
| 14 |
+
def __init__(self,
|
| 15 |
+
input_channels: int,
|
| 16 |
+
n_stages: int,
|
| 17 |
+
features_per_stage: Union[int, List[int], Tuple[int, ...]],
|
| 18 |
+
conv_op: Type[_ConvNd],
|
| 19 |
+
kernel_sizes: Union[int, List[int], Tuple[int, ...]],
|
| 20 |
+
strides: Union[int, List[int], Tuple[int, ...], Tuple[Tuple[int, ...], ...]],
|
| 21 |
+
n_blocks_per_stage: Union[int, List[int], Tuple[int, ...]],
|
| 22 |
+
conv_bias: bool = False,
|
| 23 |
+
norm_op: Union[None, Type[nn.Module]] = None,
|
| 24 |
+
norm_op_kwargs: dict = None,
|
| 25 |
+
dropout_op: Union[None, Type[_DropoutNd]] = None,
|
| 26 |
+
dropout_op_kwargs: dict = None,
|
| 27 |
+
nonlin: Union[None, Type[torch.nn.Module]] = None,
|
| 28 |
+
nonlin_kwargs: dict = None,
|
| 29 |
+
block: Union[Type[BasicBlockD], Type[BottleneckD]] = BasicBlockD,
|
| 30 |
+
bottleneck_channels: Union[int, List[int], Tuple[int, ...]] = None,
|
| 31 |
+
return_skips: bool = False,
|
| 32 |
+
disable_default_stem: bool = False,
|
| 33 |
+
stem_channels: int = None,
|
| 34 |
+
pool_type: str = 'conv',
|
| 35 |
+
stochastic_depth_p: float = 0.0,
|
| 36 |
+
squeeze_excitation: bool = False,
|
| 37 |
+
squeeze_excitation_reduction_ratio: float = 1. / 16
|
| 38 |
+
):
|
| 39 |
+
"""
|
| 40 |
+
|
| 41 |
+
:param input_channels:
|
| 42 |
+
:param n_stages:
|
| 43 |
+
:param features_per_stage: Note: If the block is BottleneckD, then this number is supposed to be the number of
|
| 44 |
+
features AFTER the expansion (which is not coded implicitly in this repository)! See todo!
|
| 45 |
+
:param conv_op:
|
| 46 |
+
:param kernel_sizes:
|
| 47 |
+
:param strides:
|
| 48 |
+
:param n_blocks_per_stage:
|
| 49 |
+
:param conv_bias:
|
| 50 |
+
:param norm_op:
|
| 51 |
+
:param norm_op_kwargs:
|
| 52 |
+
:param dropout_op:
|
| 53 |
+
:param dropout_op_kwargs:
|
| 54 |
+
:param nonlin:
|
| 55 |
+
:param nonlin_kwargs:
|
| 56 |
+
:param block:
|
| 57 |
+
:param bottleneck_channels: only needed if block is BottleneckD
|
| 58 |
+
:param return_skips: set this to True if used as encoder in a U-Net like network
|
| 59 |
+
:param disable_default_stem: If True then no stem will be created. You need to build your own and ensure it is executed first, see todo.
|
| 60 |
+
The stem in this implementation does not so stride/pooling so building your own stem is a necessity if you need this.
|
| 61 |
+
:param stem_channels: if None, features_per_stage[0] will be used for the default stem. Not recommended for BottleneckD
|
| 62 |
+
:param pool_type: if conv, strided conv will be used. avg = average pooling, max = max pooling
|
| 63 |
+
"""
|
| 64 |
+
super().__init__()
|
| 65 |
+
if isinstance(kernel_sizes, int):
|
| 66 |
+
kernel_sizes = [kernel_sizes] * n_stages
|
| 67 |
+
if isinstance(features_per_stage, int):
|
| 68 |
+
features_per_stage = [features_per_stage] * n_stages
|
| 69 |
+
if isinstance(n_blocks_per_stage, int):
|
| 70 |
+
n_blocks_per_stage = [n_blocks_per_stage] * n_stages
|
| 71 |
+
if isinstance(strides, int):
|
| 72 |
+
strides = [strides] * n_stages
|
| 73 |
+
if bottleneck_channels is None or isinstance(bottleneck_channels, int):
|
| 74 |
+
bottleneck_channels = [bottleneck_channels] * n_stages
|
| 75 |
+
assert len(
|
| 76 |
+
bottleneck_channels) == n_stages, "bottleneck_channels must be None or have as many entries as we have resolution stages (n_stages)"
|
| 77 |
+
assert len(
|
| 78 |
+
kernel_sizes) == n_stages, "kernel_sizes must have as many entries as we have resolution stages (n_stages)"
|
| 79 |
+
assert len(
|
| 80 |
+
n_blocks_per_stage) == n_stages, "n_conv_per_stage must have as many entries as we have resolution stages (n_stages)"
|
| 81 |
+
assert len(
|
| 82 |
+
features_per_stage) == n_stages, "features_per_stage must have as many entries as we have resolution stages (n_stages)"
|
| 83 |
+
assert len(strides) == n_stages, "strides must have as many entries as we have resolution stages (n_stages). " \
|
| 84 |
+
"Important: first entry is recommended to be 1, else we run strided conv drectly on the input"
|
| 85 |
+
|
| 86 |
+
pool_op = get_matching_pool_op(conv_op, pool_type=pool_type) if pool_type != 'conv' else None
|
| 87 |
+
|
| 88 |
+
# build a stem, Todo maybe we need more flexibility for this in the future. For now, if you need a custom
|
| 89 |
+
# stem you can just disable the stem and build your own.
|
| 90 |
+
# THE STEM DOES NOT DO STRIDE/POOLING IN THIS IMPLEMENTATION
|
| 91 |
+
if not disable_default_stem:
|
| 92 |
+
if stem_channels is None:
|
| 93 |
+
stem_channels = features_per_stage[0]
|
| 94 |
+
self.stem = StackedConvBlocks(1, conv_op, input_channels, stem_channels, kernel_sizes[0], 1, conv_bias,
|
| 95 |
+
norm_op, norm_op_kwargs, dropout_op, dropout_op_kwargs, nonlin, nonlin_kwargs)
|
| 96 |
+
input_channels = stem_channels
|
| 97 |
+
else:
|
| 98 |
+
self.stem = None
|
| 99 |
+
|
| 100 |
+
# now build the network
|
| 101 |
+
stages = []
|
| 102 |
+
for s in range(n_stages):
|
| 103 |
+
stride_for_conv = strides[s] if pool_op is None else 1
|
| 104 |
+
|
| 105 |
+
stage = StackedResidualBlocks(
|
| 106 |
+
n_blocks_per_stage[s], conv_op, input_channels, features_per_stage[s], kernel_sizes[s], stride_for_conv,
|
| 107 |
+
conv_bias, norm_op, norm_op_kwargs, dropout_op, dropout_op_kwargs, nonlin, nonlin_kwargs,
|
| 108 |
+
block=block, bottleneck_channels=bottleneck_channels[s], stochastic_depth_p=stochastic_depth_p,
|
| 109 |
+
squeeze_excitation=squeeze_excitation,
|
| 110 |
+
squeeze_excitation_reduction_ratio=squeeze_excitation_reduction_ratio
|
| 111 |
+
)
|
| 112 |
+
|
| 113 |
+
if pool_op is not None:
|
| 114 |
+
stage = nn.Sequential(pool_op(strides[s]), stage)
|
| 115 |
+
|
| 116 |
+
stages.append(stage)
|
| 117 |
+
input_channels = features_per_stage[s]
|
| 118 |
+
|
| 119 |
+
self.stages = nn.Sequential(*stages)
|
| 120 |
+
self.output_channels = features_per_stage
|
| 121 |
+
self.strides = [maybe_convert_scalar_to_list(conv_op, i) for i in strides]
|
| 122 |
+
self.return_skips = return_skips
|
| 123 |
+
|
| 124 |
+
# we store some things that a potential decoder needs
|
| 125 |
+
self.conv_op = conv_op
|
| 126 |
+
self.norm_op = norm_op
|
| 127 |
+
self.norm_op_kwargs = norm_op_kwargs
|
| 128 |
+
self.nonlin = nonlin
|
| 129 |
+
self.nonlin_kwargs = nonlin_kwargs
|
| 130 |
+
self.dropout_op = dropout_op
|
| 131 |
+
self.dropout_op_kwargs = dropout_op_kwargs
|
| 132 |
+
self.conv_bias = conv_bias
|
| 133 |
+
self.kernel_sizes = kernel_sizes
|
| 134 |
+
|
| 135 |
+
def forward(self, x):
|
| 136 |
+
if self.stem is not None:
|
| 137 |
+
x = self.stem(x)
|
| 138 |
+
ret = []
|
| 139 |
+
for s in self.stages:
|
| 140 |
+
x = s(x)
|
| 141 |
+
ret.append(x)
|
| 142 |
+
if self.return_skips:
|
| 143 |
+
return ret
|
| 144 |
+
else:
|
| 145 |
+
return ret[-1]
|
| 146 |
+
|
| 147 |
+
def compute_conv_feature_map_size(self, input_size):
|
| 148 |
+
if self.stem is not None:
|
| 149 |
+
output = self.stem.compute_conv_feature_map_size(input_size)
|
| 150 |
+
else:
|
| 151 |
+
output = np.int64(0)
|
| 152 |
+
|
| 153 |
+
for s in range(len(self.stages)):
|
| 154 |
+
output += self.stages[s].compute_conv_feature_map_size(input_size)
|
| 155 |
+
input_size = [i // j for i, j in zip(input_size, self.strides[s])]
|
| 156 |
+
|
| 157 |
+
return output
|
| 158 |
+
|
| 159 |
+
|
| 160 |
+
if __name__ == '__main__':
|
| 161 |
+
data = torch.rand((1, 3, 128, 160))
|
| 162 |
+
|
| 163 |
+
model = ResidualEncoder(3, 5, (2, 4, 6, 8, 10), nn.Conv2d, 3, ((1, 1), 2, (2, 2), (2, 2), (2, 2)), 2, False,
|
| 164 |
+
nn.BatchNorm2d, None, None, None, nn.ReLU, None, stem_channels=7)
|
| 165 |
+
import hiddenlayer as hl
|
| 166 |
+
|
| 167 |
+
g = hl.build_graph(model, data,
|
| 168 |
+
transforms=None)
|
| 169 |
+
g.save("network_architecture.pdf")
|
| 170 |
+
del g
|
| 171 |
+
|
| 172 |
+
print(model.compute_conv_feature_map_size((128, 160)))
|
RADAR_inference/dynamic_network_architectures/building_blocks/simple_conv_blocks.py
ADDED
|
@@ -0,0 +1,167 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import Tuple, List, Union, Type
|
| 2 |
+
|
| 3 |
+
import numpy as np
|
| 4 |
+
import torch.nn
|
| 5 |
+
from torch import nn
|
| 6 |
+
from torch.nn.modules.conv import _ConvNd
|
| 7 |
+
from torch.nn.modules.dropout import _DropoutNd
|
| 8 |
+
|
| 9 |
+
from dynamic_network_architectures.building_blocks.helper import maybe_convert_scalar_to_list
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
class ConvDropoutNormReLU(nn.Module):
|
| 13 |
+
def __init__(self,
|
| 14 |
+
conv_op: Type[_ConvNd],
|
| 15 |
+
input_channels: int,
|
| 16 |
+
output_channels: int,
|
| 17 |
+
kernel_size: Union[int, List[int], Tuple[int, ...]],
|
| 18 |
+
stride: Union[int, List[int], Tuple[int, ...]],
|
| 19 |
+
conv_bias: bool = False,
|
| 20 |
+
norm_op: Union[None, Type[nn.Module]] = None,
|
| 21 |
+
norm_op_kwargs: dict = None,
|
| 22 |
+
dropout_op: Union[None, Type[_DropoutNd]] = None,
|
| 23 |
+
dropout_op_kwargs: dict = None,
|
| 24 |
+
nonlin: Union[None, Type[torch.nn.Module]] = None,
|
| 25 |
+
nonlin_kwargs: dict = None,
|
| 26 |
+
nonlin_first: bool = False
|
| 27 |
+
):
|
| 28 |
+
super(ConvDropoutNormReLU, self).__init__()
|
| 29 |
+
self.input_channels = input_channels
|
| 30 |
+
self.output_channels = output_channels
|
| 31 |
+
stride = maybe_convert_scalar_to_list(conv_op, stride)
|
| 32 |
+
self.stride = stride
|
| 33 |
+
|
| 34 |
+
kernel_size = maybe_convert_scalar_to_list(conv_op, kernel_size)
|
| 35 |
+
if norm_op_kwargs is None:
|
| 36 |
+
norm_op_kwargs = {}
|
| 37 |
+
if nonlin_kwargs is None:
|
| 38 |
+
nonlin_kwargs = {}
|
| 39 |
+
|
| 40 |
+
ops = []
|
| 41 |
+
|
| 42 |
+
self.conv = conv_op(
|
| 43 |
+
input_channels,
|
| 44 |
+
output_channels,
|
| 45 |
+
kernel_size,
|
| 46 |
+
stride,
|
| 47 |
+
padding=[(i - 1) // 2 for i in kernel_size],
|
| 48 |
+
dilation=1,
|
| 49 |
+
bias=conv_bias,
|
| 50 |
+
)
|
| 51 |
+
ops.append(self.conv)
|
| 52 |
+
|
| 53 |
+
if dropout_op is not None:
|
| 54 |
+
self.dropout = dropout_op(**dropout_op_kwargs)
|
| 55 |
+
ops.append(self.dropout)
|
| 56 |
+
|
| 57 |
+
if norm_op is not None:
|
| 58 |
+
self.norm = norm_op(output_channels, **norm_op_kwargs)
|
| 59 |
+
ops.append(self.norm)
|
| 60 |
+
|
| 61 |
+
if nonlin is not None:
|
| 62 |
+
self.nonlin = nonlin(**nonlin_kwargs)
|
| 63 |
+
ops.append(self.nonlin)
|
| 64 |
+
|
| 65 |
+
if nonlin_first and (norm_op is not None and nonlin is not None):
|
| 66 |
+
ops[-1], ops[-2] = ops[-2], ops[-1]
|
| 67 |
+
|
| 68 |
+
self.all_modules = nn.Sequential(*ops)
|
| 69 |
+
|
| 70 |
+
def forward(self, x):
|
| 71 |
+
return self.all_modules(x)
|
| 72 |
+
|
| 73 |
+
def compute_conv_feature_map_size(self, input_size):
|
| 74 |
+
assert len(input_size) == len(self.stride), "just give the image size without color/feature channels or " \
|
| 75 |
+
"batch channel. Do not give input_size=(b, c, x, y(, z)). " \
|
| 76 |
+
"Give input_size=(x, y(, z))!"
|
| 77 |
+
output_size = [i // j for i, j in zip(input_size, self.stride)] # we always do same padding
|
| 78 |
+
return np.prod([self.output_channels, *output_size], dtype=np.int64)
|
| 79 |
+
|
| 80 |
+
|
| 81 |
+
class StackedConvBlocks(nn.Module):
|
| 82 |
+
def __init__(self,
|
| 83 |
+
num_convs: int,
|
| 84 |
+
conv_op: Type[_ConvNd],
|
| 85 |
+
input_channels: int,
|
| 86 |
+
output_channels: Union[int, List[int], Tuple[int, ...]],
|
| 87 |
+
kernel_size: Union[int, List[int], Tuple[int, ...]],
|
| 88 |
+
initial_stride: Union[int, List[int], Tuple[int, ...]],
|
| 89 |
+
conv_bias: bool = False,
|
| 90 |
+
norm_op: Union[None, Type[nn.Module]] = None,
|
| 91 |
+
norm_op_kwargs: dict = None,
|
| 92 |
+
dropout_op: Union[None, Type[_DropoutNd]] = None,
|
| 93 |
+
dropout_op_kwargs: dict = None,
|
| 94 |
+
nonlin: Union[None, Type[torch.nn.Module]] = None,
|
| 95 |
+
nonlin_kwargs: dict = None,
|
| 96 |
+
nonlin_first: bool = False
|
| 97 |
+
):
|
| 98 |
+
"""
|
| 99 |
+
|
| 100 |
+
:param conv_op:
|
| 101 |
+
:param num_convs:
|
| 102 |
+
:param input_channels:
|
| 103 |
+
:param output_channels: can be int or a list/tuple of int. If list/tuple are provided, each entry is for
|
| 104 |
+
one conv. The length of the list/tuple must then naturally be num_convs
|
| 105 |
+
:param kernel_size:
|
| 106 |
+
:param initial_stride:
|
| 107 |
+
:param conv_bias:
|
| 108 |
+
:param norm_op:
|
| 109 |
+
:param norm_op_kwargs:
|
| 110 |
+
:param dropout_op:
|
| 111 |
+
:param dropout_op_kwargs:
|
| 112 |
+
:param nonlin:
|
| 113 |
+
:param nonlin_kwargs:
|
| 114 |
+
"""
|
| 115 |
+
super().__init__()
|
| 116 |
+
if not isinstance(output_channels, (tuple, list)):
|
| 117 |
+
output_channels = [output_channels] * num_convs
|
| 118 |
+
|
| 119 |
+
self.convs = nn.Sequential(
|
| 120 |
+
ConvDropoutNormReLU(
|
| 121 |
+
conv_op, input_channels, output_channels[0], kernel_size, initial_stride, conv_bias, norm_op,
|
| 122 |
+
norm_op_kwargs, dropout_op, dropout_op_kwargs, nonlin, nonlin_kwargs, nonlin_first
|
| 123 |
+
),
|
| 124 |
+
*[
|
| 125 |
+
ConvDropoutNormReLU(
|
| 126 |
+
conv_op, output_channels[i - 1], output_channels[i], kernel_size, 1, conv_bias, norm_op,
|
| 127 |
+
norm_op_kwargs, dropout_op, dropout_op_kwargs, nonlin, nonlin_kwargs, nonlin_first
|
| 128 |
+
)
|
| 129 |
+
for i in range(1, num_convs)
|
| 130 |
+
]
|
| 131 |
+
)
|
| 132 |
+
|
| 133 |
+
self.output_channels = output_channels[-1]
|
| 134 |
+
self.initial_stride = maybe_convert_scalar_to_list(conv_op, initial_stride)
|
| 135 |
+
|
| 136 |
+
def forward(self, x):
|
| 137 |
+
return self.convs(x)
|
| 138 |
+
|
| 139 |
+
def compute_conv_feature_map_size(self, input_size):
|
| 140 |
+
assert len(input_size) == len(self.initial_stride), "just give the image size without color/feature channels or " \
|
| 141 |
+
"batch channel. Do not give input_size=(b, c, x, y(, z)). " \
|
| 142 |
+
"Give input_size=(x, y(, z))!"
|
| 143 |
+
output = self.convs[0].compute_conv_feature_map_size(input_size)
|
| 144 |
+
size_after_stride = [i // j for i, j in zip(input_size, self.initial_stride)]
|
| 145 |
+
for b in self.convs[1:]:
|
| 146 |
+
output += b.compute_conv_feature_map_size(size_after_stride)
|
| 147 |
+
return output
|
| 148 |
+
|
| 149 |
+
|
| 150 |
+
if __name__ == '__main__':
|
| 151 |
+
data = torch.rand((1, 3, 40, 32))
|
| 152 |
+
|
| 153 |
+
stx = StackedConvBlocks(2, nn.Conv2d, 24, 16, (3, 3), 2,
|
| 154 |
+
norm_op=nn.BatchNorm2d, nonlin=nn.ReLU, nonlin_kwargs={'inplace': True},
|
| 155 |
+
)
|
| 156 |
+
model = nn.Sequential(ConvDropoutNormReLU(nn.Conv2d,
|
| 157 |
+
3, 24, 3, 1, True, nn.BatchNorm2d, {}, None, None, nn.LeakyReLU,
|
| 158 |
+
{'inplace': True}),
|
| 159 |
+
stx)
|
| 160 |
+
import hiddenlayer as hl
|
| 161 |
+
|
| 162 |
+
g = hl.build_graph(model, data,
|
| 163 |
+
transforms=None)
|
| 164 |
+
g.save("network_architecture.pdf")
|
| 165 |
+
del g
|
| 166 |
+
|
| 167 |
+
stx.compute_conv_feature_map_size((40, 32))
|
RADAR_inference/dynamic_network_architectures/building_blocks/unet_decoder.py
ADDED
|
@@ -0,0 +1,154 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import numpy as np
|
| 2 |
+
import torch
|
| 3 |
+
from torch import nn
|
| 4 |
+
from typing import Union, List, Tuple, Type
|
| 5 |
+
|
| 6 |
+
from torch.nn.modules.dropout import _DropoutNd
|
| 7 |
+
|
| 8 |
+
from dynamic_network_architectures.building_blocks.simple_conv_blocks import StackedConvBlocks
|
| 9 |
+
from dynamic_network_architectures.building_blocks.helper import get_matching_convtransp
|
| 10 |
+
from dynamic_network_architectures.building_blocks.residual_encoders import ResidualEncoder
|
| 11 |
+
from dynamic_network_architectures.building_blocks.plain_conv_encoder import PlainConvEncoder
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
class UNetDecoder(nn.Module):
|
| 15 |
+
def __init__(self,
|
| 16 |
+
encoder: Union[PlainConvEncoder, ResidualEncoder],
|
| 17 |
+
num_classes: int,
|
| 18 |
+
n_conv_per_stage: Union[int, Tuple[int, ...], List[int]],
|
| 19 |
+
deep_supervision,
|
| 20 |
+
nonlin_first: bool = False,
|
| 21 |
+
norm_op: Union[None, Type[nn.Module]] = None,
|
| 22 |
+
norm_op_kwargs: dict = None,
|
| 23 |
+
dropout_op: Union[None, Type[_DropoutNd]] = None,
|
| 24 |
+
dropout_op_kwargs: dict = None,
|
| 25 |
+
nonlin: Union[None, Type[torch.nn.Module]] = None,
|
| 26 |
+
nonlin_kwargs: dict = None,
|
| 27 |
+
conv_bias: bool = None
|
| 28 |
+
):
|
| 29 |
+
"""
|
| 30 |
+
This class needs the skips of the encoder as input in its forward.
|
| 31 |
+
|
| 32 |
+
the encoder goes all the way to the bottleneck, so that's where the decoder picks up. stages in the decoder
|
| 33 |
+
are sorted by order of computation, so the first stage has the lowest resolution and takes the bottleneck
|
| 34 |
+
features and the lowest skip as inputs
|
| 35 |
+
the decoder has two (three) parts in each stage:
|
| 36 |
+
1) conv transpose to upsample the feature maps of the stage below it (or the bottleneck in case of the first stage)
|
| 37 |
+
2) n_conv_per_stage conv blocks to let the two inputs get to know each other and merge
|
| 38 |
+
3) (optional if deep_supervision=True) a segmentation output Todo: enable upsample logits?
|
| 39 |
+
:param encoder:
|
| 40 |
+
:param num_classes:
|
| 41 |
+
:param n_conv_per_stage:
|
| 42 |
+
:param deep_supervision:
|
| 43 |
+
"""
|
| 44 |
+
super().__init__()
|
| 45 |
+
self.deep_supervision = deep_supervision
|
| 46 |
+
self.encoder = encoder
|
| 47 |
+
self.num_classes = num_classes
|
| 48 |
+
n_stages_encoder = len(encoder.output_channels)
|
| 49 |
+
if isinstance(n_conv_per_stage, int):
|
| 50 |
+
n_conv_per_stage = [n_conv_per_stage] * (n_stages_encoder - 1)
|
| 51 |
+
assert len(n_conv_per_stage) == n_stages_encoder - 1, "n_conv_per_stage must have as many entries as we have " \
|
| 52 |
+
"resolution stages - 1 (n_stages in encoder - 1), " \
|
| 53 |
+
"here: %d" % n_stages_encoder
|
| 54 |
+
|
| 55 |
+
transpconv_op = get_matching_convtransp(conv_op=encoder.conv_op)
|
| 56 |
+
conv_bias = encoder.conv_bias if conv_bias is None else conv_bias
|
| 57 |
+
norm_op = encoder.norm_op if norm_op is None else norm_op
|
| 58 |
+
norm_op_kwargs = encoder.norm_op_kwargs if norm_op_kwargs is None else norm_op_kwargs
|
| 59 |
+
dropout_op = encoder.dropout_op if dropout_op is None else dropout_op
|
| 60 |
+
dropout_op_kwargs = encoder.dropout_op_kwargs if dropout_op_kwargs is None else dropout_op_kwargs
|
| 61 |
+
nonlin = encoder.nonlin if nonlin is None else nonlin
|
| 62 |
+
nonlin_kwargs = encoder.nonlin_kwargs if nonlin_kwargs is None else nonlin_kwargs
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
# we start with the bottleneck and work out way up
|
| 66 |
+
stages = []
|
| 67 |
+
transpconvs = []
|
| 68 |
+
seg_layers = []
|
| 69 |
+
for s in range(1, n_stages_encoder):
|
| 70 |
+
input_features_below = encoder.output_channels[-s]
|
| 71 |
+
input_features_skip = encoder.output_channels[-(s + 1)]
|
| 72 |
+
stride_for_transpconv = encoder.strides[-s]
|
| 73 |
+
transpconvs.append(transpconv_op(
|
| 74 |
+
input_features_below, input_features_skip, stride_for_transpconv, stride_for_transpconv,
|
| 75 |
+
bias=conv_bias
|
| 76 |
+
))
|
| 77 |
+
# input features to conv is 2x input_features_skip (concat input_features_skip with transpconv output)
|
| 78 |
+
stages.append(StackedConvBlocks(
|
| 79 |
+
n_conv_per_stage[s-1], encoder.conv_op, 2 * input_features_skip, input_features_skip,
|
| 80 |
+
encoder.kernel_sizes[-(s + 1)], 1,
|
| 81 |
+
conv_bias,
|
| 82 |
+
norm_op,
|
| 83 |
+
norm_op_kwargs,
|
| 84 |
+
dropout_op,
|
| 85 |
+
dropout_op_kwargs,
|
| 86 |
+
nonlin,
|
| 87 |
+
nonlin_kwargs,
|
| 88 |
+
nonlin_first
|
| 89 |
+
))
|
| 90 |
+
|
| 91 |
+
# we always build the deep supervision outputs so that we can always load parameters. If we don't do this
|
| 92 |
+
# then a model trained with deep_supervision=True could not easily be loaded at inference time where
|
| 93 |
+
# deep supervision is not needed. It's just a convenience thing
|
| 94 |
+
seg_layers.append(encoder.conv_op(input_features_skip, num_classes, 1, 1, 0, bias=True))
|
| 95 |
+
|
| 96 |
+
self.stages = nn.ModuleList(stages)
|
| 97 |
+
self.transpconvs = nn.ModuleList(transpconvs)
|
| 98 |
+
self.seg_layers = nn.ModuleList(seg_layers)
|
| 99 |
+
|
| 100 |
+
def forward(self, skips):
|
| 101 |
+
"""
|
| 102 |
+
we expect to get the skips in the order they were computed, so the bottleneck should be the last entry
|
| 103 |
+
:param skips:
|
| 104 |
+
:return:
|
| 105 |
+
"""
|
| 106 |
+
lres_input = skips[-1]
|
| 107 |
+
seg_outputs = []
|
| 108 |
+
for s in range(len(self.stages)):
|
| 109 |
+
x = self.transpconvs[s](lres_input)
|
| 110 |
+
x = torch.cat((x, skips[-(s+2)]), 1)
|
| 111 |
+
x = self.stages[s](x)
|
| 112 |
+
if self.deep_supervision:
|
| 113 |
+
seg_outputs.append(self.seg_layers[s](x))
|
| 114 |
+
elif s == (len(self.stages) - 1):
|
| 115 |
+
seg_outputs.append(self.seg_layers[-1](x))
|
| 116 |
+
lres_input = x
|
| 117 |
+
|
| 118 |
+
# invert seg outputs so that the largest segmentation prediction is returned first
|
| 119 |
+
seg_outputs = seg_outputs[::-1]
|
| 120 |
+
|
| 121 |
+
if not self.deep_supervision:
|
| 122 |
+
r = seg_outputs[0]
|
| 123 |
+
else:
|
| 124 |
+
r = seg_outputs
|
| 125 |
+
return r
|
| 126 |
+
|
| 127 |
+
def compute_conv_feature_map_size(self, input_size):
|
| 128 |
+
"""
|
| 129 |
+
IMPORTANT: input_size is the input_size of the encoder!
|
| 130 |
+
:param input_size:
|
| 131 |
+
:return:
|
| 132 |
+
"""
|
| 133 |
+
# first we need to compute the skip sizes. Skip bottleneck because all output feature maps of our ops will at
|
| 134 |
+
# least have the size of the skip above that (therefore -1)
|
| 135 |
+
skip_sizes = []
|
| 136 |
+
for s in range(len(self.encoder.strides) - 1):
|
| 137 |
+
skip_sizes.append([i // j for i, j in zip(input_size, self.encoder.strides[s])])
|
| 138 |
+
input_size = skip_sizes[-1]
|
| 139 |
+
# print(skip_sizes)
|
| 140 |
+
|
| 141 |
+
assert len(skip_sizes) == len(self.stages)
|
| 142 |
+
|
| 143 |
+
# our ops are the other way around, so let's match things up
|
| 144 |
+
output = np.int64(0)
|
| 145 |
+
for s in range(len(self.stages)):
|
| 146 |
+
# print(skip_sizes[-(s+1)], self.encoder.output_channels[-(s+2)])
|
| 147 |
+
# conv blocks
|
| 148 |
+
output += self.stages[s].compute_conv_feature_map_size(skip_sizes[-(s+1)])
|
| 149 |
+
# trans conv
|
| 150 |
+
output += np.prod([self.encoder.output_channels[-(s+2)], *skip_sizes[-(s+1)]], dtype=np.int64)
|
| 151 |
+
# segmentation
|
| 152 |
+
if self.deep_supervision or (s == (len(self.stages) - 1)):
|
| 153 |
+
output += np.prod([self.num_classes, *skip_sizes[-(s+1)]], dtype=np.int64)
|
| 154 |
+
return output
|
RADAR_inference/dynamic_network_architectures/building_blocks/unet_decoder_light.py
ADDED
|
@@ -0,0 +1,154 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import numpy as np
|
| 2 |
+
import torch
|
| 3 |
+
from torch import nn
|
| 4 |
+
from typing import Union, List, Tuple, Type
|
| 5 |
+
|
| 6 |
+
from torch.nn.modules.dropout import _DropoutNd
|
| 7 |
+
|
| 8 |
+
from dynamic_network_architectures.building_blocks.simple_conv_blocks import StackedConvBlocks
|
| 9 |
+
from dynamic_network_architectures.building_blocks.helper import get_matching_convtransp
|
| 10 |
+
from dynamic_network_architectures.building_blocks.residual_encoders import ResidualEncoder
|
| 11 |
+
from dynamic_network_architectures.building_blocks.plain_conv_encoder import PlainConvEncoder
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
class UNetDecoder(nn.Module):
|
| 15 |
+
def __init__(self,
|
| 16 |
+
encoder: Union[PlainConvEncoder, ResidualEncoder],
|
| 17 |
+
num_classes: int,
|
| 18 |
+
n_conv_per_stage: Union[int, Tuple[int, ...], List[int]],
|
| 19 |
+
deep_supervision,
|
| 20 |
+
nonlin_first: bool = False,
|
| 21 |
+
norm_op: Union[None, Type[nn.Module]] = None,
|
| 22 |
+
norm_op_kwargs: dict = None,
|
| 23 |
+
dropout_op: Union[None, Type[_DropoutNd]] = None,
|
| 24 |
+
dropout_op_kwargs: dict = None,
|
| 25 |
+
nonlin: Union[None, Type[torch.nn.Module]] = None,
|
| 26 |
+
nonlin_kwargs: dict = None,
|
| 27 |
+
conv_bias: bool = None
|
| 28 |
+
):
|
| 29 |
+
"""
|
| 30 |
+
This class needs the skips of the encoder as input in its forward.
|
| 31 |
+
|
| 32 |
+
the encoder goes all the way to the bottleneck, so that's where the decoder picks up. stages in the decoder
|
| 33 |
+
are sorted by order of computation, so the first stage has the lowest resolution and takes the bottleneck
|
| 34 |
+
features and the lowest skip as inputs
|
| 35 |
+
the decoder has two (three) parts in each stage:
|
| 36 |
+
1) conv transpose to upsample the feature maps of the stage below it (or the bottleneck in case of the first stage)
|
| 37 |
+
2) n_conv_per_stage conv blocks to let the two inputs get to know each other and merge
|
| 38 |
+
3) (optional if deep_supervision=True) a segmentation output Todo: enable upsample logits?
|
| 39 |
+
:param encoder:
|
| 40 |
+
:param num_classes:
|
| 41 |
+
:param n_conv_per_stage:
|
| 42 |
+
:param deep_supervision:
|
| 43 |
+
"""
|
| 44 |
+
super().__init__()
|
| 45 |
+
self.deep_supervision = deep_supervision
|
| 46 |
+
self.encoder = encoder
|
| 47 |
+
self.num_classes = num_classes
|
| 48 |
+
n_stages_encoder = len(encoder.output_channels)
|
| 49 |
+
if isinstance(n_conv_per_stage, int):
|
| 50 |
+
n_conv_per_stage = [n_conv_per_stage] * (n_stages_encoder - 1)
|
| 51 |
+
# assert len(n_conv_per_stage) == n_stages_encoder - 1, "n_conv_per_stage must have as many entries as we have " \
|
| 52 |
+
# "resolution stages - 1 (n_stages in encoder - 1), " \
|
| 53 |
+
# "here: %d" % n_stages_encoder
|
| 54 |
+
|
| 55 |
+
transpconv_op = get_matching_convtransp(conv_op=encoder.conv_op)
|
| 56 |
+
conv_bias = encoder.conv_bias if conv_bias is None else conv_bias
|
| 57 |
+
norm_op = encoder.norm_op if norm_op is None else norm_op
|
| 58 |
+
norm_op_kwargs = encoder.norm_op_kwargs if norm_op_kwargs is None else norm_op_kwargs
|
| 59 |
+
dropout_op = encoder.dropout_op if dropout_op is None else dropout_op
|
| 60 |
+
dropout_op_kwargs = encoder.dropout_op_kwargs if dropout_op_kwargs is None else dropout_op_kwargs
|
| 61 |
+
nonlin = encoder.nonlin if nonlin is None else nonlin
|
| 62 |
+
nonlin_kwargs = encoder.nonlin_kwargs if nonlin_kwargs is None else nonlin_kwargs
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
stages = []
|
| 66 |
+
transpconvs = []
|
| 67 |
+
seg_layers = []
|
| 68 |
+
for s in range(1, n_stages_encoder-1):
|
| 69 |
+
input_features_below = encoder.output_channels[-s]
|
| 70 |
+
input_features_skip = encoder.output_channels[-(s + 1)]
|
| 71 |
+
stride_for_transpconv = encoder.strides[-s]
|
| 72 |
+
transpconvs.append(transpconv_op(
|
| 73 |
+
input_features_below, input_features_skip, stride_for_transpconv, stride_for_transpconv,
|
| 74 |
+
bias=conv_bias
|
| 75 |
+
))
|
| 76 |
+
# input features to conv is 2x input_features_skip (concat input_features_skip with transpconv output)
|
| 77 |
+
stages.append(StackedConvBlocks(
|
| 78 |
+
n_conv_per_stage[s-1], encoder.conv_op, 2 * input_features_skip, input_features_skip,
|
| 79 |
+
encoder.kernel_sizes[-(s + 1)], 1,
|
| 80 |
+
conv_bias,
|
| 81 |
+
norm_op,
|
| 82 |
+
norm_op_kwargs,
|
| 83 |
+
dropout_op,
|
| 84 |
+
dropout_op_kwargs,
|
| 85 |
+
nonlin,
|
| 86 |
+
nonlin_kwargs,
|
| 87 |
+
nonlin_first
|
| 88 |
+
))
|
| 89 |
+
|
| 90 |
+
# we always build the deep supervision outputs so that we can always load parameters. If we don't do this
|
| 91 |
+
# then a model trained with deep_supervision=True could not easily be loaded at inference time where
|
| 92 |
+
# deep supervision is not needed. It's just a convenience thing
|
| 93 |
+
seg_layers.append(encoder.conv_op(input_features_skip, num_classes, 1, 1, 0, bias=True))
|
| 94 |
+
|
| 95 |
+
self.stages = nn.ModuleList(stages)
|
| 96 |
+
self.transpconvs = nn.ModuleList(transpconvs)
|
| 97 |
+
self.seg_layers = nn.ModuleList(seg_layers)
|
| 98 |
+
|
| 99 |
+
def forward(self, skips):
|
| 100 |
+
"""
|
| 101 |
+
we expect to get the skips in the order they were computed, so the bottleneck should be the last entry
|
| 102 |
+
:param skips:
|
| 103 |
+
:return:
|
| 104 |
+
"""
|
| 105 |
+
# print('debug light nnunvetv2 decoder')
|
| 106 |
+
lres_input = skips[-1]
|
| 107 |
+
seg_outputs = []
|
| 108 |
+
for s in range(len(self.stages)):
|
| 109 |
+
x = self.transpconvs[s](lres_input)
|
| 110 |
+
x = torch.cat((x, skips[-(s+2)]), 1)
|
| 111 |
+
x = self.stages[s](x)
|
| 112 |
+
if self.deep_supervision:
|
| 113 |
+
seg_outputs.append(self.seg_layers[s](x))
|
| 114 |
+
elif s == (len(self.stages) - 1):
|
| 115 |
+
seg_outputs.append(self.seg_layers[-1](x))
|
| 116 |
+
lres_input = x
|
| 117 |
+
|
| 118 |
+
# invert seg outputs so that the largest segmentation prediction is returned first
|
| 119 |
+
seg_outputs = seg_outputs[::-1]
|
| 120 |
+
|
| 121 |
+
if not self.deep_supervision:
|
| 122 |
+
r = seg_outputs[0]
|
| 123 |
+
else:
|
| 124 |
+
r = seg_outputs
|
| 125 |
+
return r
|
| 126 |
+
|
| 127 |
+
def compute_conv_feature_map_size(self, input_size):
|
| 128 |
+
"""
|
| 129 |
+
IMPORTANT: input_size is the input_size of the encoder!
|
| 130 |
+
:param input_size:
|
| 131 |
+
:return:
|
| 132 |
+
"""
|
| 133 |
+
# first we need to compute the skip sizes. Skip bottleneck because all output feature maps of our ops will at
|
| 134 |
+
# least have the size of the skip above that (therefore -1)
|
| 135 |
+
skip_sizes = []
|
| 136 |
+
for s in range(len(self.encoder.strides) - 1):
|
| 137 |
+
skip_sizes.append([i // j for i, j in zip(input_size, self.encoder.strides[s])])
|
| 138 |
+
input_size = skip_sizes[-1]
|
| 139 |
+
# print(skip_sizes)
|
| 140 |
+
|
| 141 |
+
assert len(skip_sizes) == len(self.stages)
|
| 142 |
+
|
| 143 |
+
# our ops are the other way around, so let's match things up
|
| 144 |
+
output = np.int64(0)
|
| 145 |
+
for s in range(len(self.stages)):
|
| 146 |
+
# print(skip_sizes[-(s+1)], self.encoder.output_channels[-(s+2)])
|
| 147 |
+
# conv blocks
|
| 148 |
+
output += self.stages[s].compute_conv_feature_map_size(skip_sizes[-(s+1)])
|
| 149 |
+
# trans conv
|
| 150 |
+
output += np.prod([self.encoder.output_channels[-(s+2)], *skip_sizes[-(s+1)]], dtype=np.int64)
|
| 151 |
+
# segmentation
|
| 152 |
+
if self.deep_supervision or (s == (len(self.stages) - 1)):
|
| 153 |
+
output += np.prod([self.num_classes, *skip_sizes[-(s+1)]], dtype=np.int64)
|
| 154 |
+
return output
|
RADAR_inference/dynamic_network_architectures/building_blocks/unet_residual_decoder.py
ADDED
|
@@ -0,0 +1,155 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import numpy as np
|
| 2 |
+
from typing import Union, Tuple, List, Type
|
| 3 |
+
|
| 4 |
+
import torch
|
| 5 |
+
from dynamic_network_architectures.building_blocks.helper import get_matching_convtransp
|
| 6 |
+
from dynamic_network_architectures.building_blocks.plain_conv_encoder import PlainConvEncoder
|
| 7 |
+
from dynamic_network_architectures.building_blocks.residual import StackedResidualBlocks
|
| 8 |
+
from dynamic_network_architectures.building_blocks.residual_encoders import ResidualEncoder
|
| 9 |
+
from torch import nn
|
| 10 |
+
from torch.nn.modules.dropout import _DropoutNd
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
class UNetResDecoder(nn.Module):
|
| 14 |
+
def __init__(self,
|
| 15 |
+
encoder: Union[PlainConvEncoder, ResidualEncoder],
|
| 16 |
+
num_classes: int,
|
| 17 |
+
n_conv_per_stage: Union[int, Tuple[int, ...], List[int]],
|
| 18 |
+
deep_supervision,
|
| 19 |
+
nonlin_first: bool = False,
|
| 20 |
+
norm_op: Union[None, Type[nn.Module]] = None,
|
| 21 |
+
norm_op_kwargs: dict = None,
|
| 22 |
+
dropout_op: Union[None, Type[_DropoutNd]] = None,
|
| 23 |
+
dropout_op_kwargs: dict = None,
|
| 24 |
+
nonlin: Union[None, Type[torch.nn.Module]] = None,
|
| 25 |
+
nonlin_kwargs: dict = None,
|
| 26 |
+
conv_bias: bool = None
|
| 27 |
+
):
|
| 28 |
+
"""
|
| 29 |
+
This class needs the skips of the encoder as input in its forward.
|
| 30 |
+
|
| 31 |
+
the encoder goes all the way to the bottleneck, so that's where the decoder picks up. stages in the decoder
|
| 32 |
+
are sorted by order of computation, so the first stage has the lowest resolution and takes the bottleneck
|
| 33 |
+
features and the lowest skip as inputs
|
| 34 |
+
the decoder has two (three) parts in each stage:
|
| 35 |
+
1) conv transpose to upsample the feature maps of the stage below it (or the bottleneck in case of the first stage)
|
| 36 |
+
2) n_conv_per_stage conv blocks to let the two inputs get to know each other and merge
|
| 37 |
+
3) (optional if deep_supervision=True) a segmentation output Todo: enable upsample logits?
|
| 38 |
+
:param encoder:
|
| 39 |
+
:param num_classes:
|
| 40 |
+
:param n_conv_per_stage:
|
| 41 |
+
:param deep_supervision:
|
| 42 |
+
"""
|
| 43 |
+
super().__init__()
|
| 44 |
+
self.deep_supervision = deep_supervision
|
| 45 |
+
self.encoder = encoder
|
| 46 |
+
self.num_classes = num_classes
|
| 47 |
+
n_stages_encoder = len(encoder.output_channels)
|
| 48 |
+
if isinstance(n_conv_per_stage, int):
|
| 49 |
+
n_conv_per_stage = [n_conv_per_stage] * (n_stages_encoder - 1)
|
| 50 |
+
assert len(n_conv_per_stage) == n_stages_encoder - 1, "n_conv_per_stage must have as many entries as we have " \
|
| 51 |
+
"resolution stages - 1 (n_stages in encoder - 1), " \
|
| 52 |
+
"here: %d" % n_stages_encoder
|
| 53 |
+
|
| 54 |
+
transpconv_op = get_matching_convtransp(conv_op=encoder.conv_op)
|
| 55 |
+
conv_bias = encoder.conv_bias if conv_bias is None else conv_bias
|
| 56 |
+
norm_op = encoder.norm_op if norm_op is None else norm_op
|
| 57 |
+
norm_op_kwargs = encoder.norm_op_kwargs if norm_op_kwargs is None else norm_op_kwargs
|
| 58 |
+
dropout_op = encoder.dropout_op if dropout_op is None else dropout_op
|
| 59 |
+
dropout_op_kwargs = encoder.dropout_op_kwargs if dropout_op_kwargs is None else dropout_op_kwargs
|
| 60 |
+
nonlin = encoder.nonlin if nonlin is None else nonlin
|
| 61 |
+
nonlin_kwargs = encoder.nonlin_kwargs if nonlin_kwargs is None else nonlin_kwargs
|
| 62 |
+
|
| 63 |
+
# we start with the bottleneck and work out way up
|
| 64 |
+
stages = []
|
| 65 |
+
transpconvs = []
|
| 66 |
+
seg_layers = []
|
| 67 |
+
for s in range(1, n_stages_encoder):
|
| 68 |
+
input_features_below = encoder.output_channels[-s]
|
| 69 |
+
input_features_skip = encoder.output_channels[-(s + 1)]
|
| 70 |
+
stride_for_transpconv = encoder.strides[-s]
|
| 71 |
+
transpconvs.append(transpconv_op(
|
| 72 |
+
input_features_below, input_features_skip, stride_for_transpconv, stride_for_transpconv,
|
| 73 |
+
bias=encoder.conv_bias
|
| 74 |
+
))
|
| 75 |
+
# input features to conv is 2x input_features_skip (concat input_features_skip with transpconv output)
|
| 76 |
+
stages.append(StackedResidualBlocks(
|
| 77 |
+
n_blocks=n_conv_per_stage[s - 1],
|
| 78 |
+
conv_op=encoder.conv_op,
|
| 79 |
+
input_channels=2 * input_features_skip,
|
| 80 |
+
output_channels=input_features_skip,
|
| 81 |
+
kernel_size=encoder.kernel_sizes[-(s + 1)],
|
| 82 |
+
initial_stride=1,
|
| 83 |
+
conv_bias=conv_bias,
|
| 84 |
+
norm_op=norm_op,
|
| 85 |
+
norm_op_kwargs=norm_op_kwargs,
|
| 86 |
+
dropout_op=dropout_op,
|
| 87 |
+
dropout_op_kwargs=dropout_op_kwargs,
|
| 88 |
+
nonlin=nonlin,
|
| 89 |
+
nonlin_kwargs=nonlin_kwargs,
|
| 90 |
+
))
|
| 91 |
+
|
| 92 |
+
# we always build the deep supervision outputs so that we can always load parameters. If we don't do this
|
| 93 |
+
# then a model trained with deep_supervision=True could not easily be loaded at inference time where
|
| 94 |
+
# deep supervision is not needed. It's just a convenience thing
|
| 95 |
+
seg_layers.append(encoder.conv_op(input_features_skip, num_classes, 1, 1, 0, bias=True))
|
| 96 |
+
|
| 97 |
+
self.stages = nn.ModuleList(stages)
|
| 98 |
+
self.transpconvs = nn.ModuleList(transpconvs)
|
| 99 |
+
self.seg_layers = nn.ModuleList(seg_layers)
|
| 100 |
+
|
| 101 |
+
def forward(self, skips):
|
| 102 |
+
"""
|
| 103 |
+
we expect to get the skips in the order they were computed, so the bottleneck should be the last entry
|
| 104 |
+
:param skips:
|
| 105 |
+
:return:
|
| 106 |
+
"""
|
| 107 |
+
lres_input = skips[-1]
|
| 108 |
+
seg_outputs = []
|
| 109 |
+
for s in range(len(self.stages)):
|
| 110 |
+
x = self.transpconvs[s](lres_input)
|
| 111 |
+
x = torch.cat((x, skips[-(s + 2)]), 1)
|
| 112 |
+
x = self.stages[s](x)
|
| 113 |
+
if self.deep_supervision:
|
| 114 |
+
seg_outputs.append(self.seg_layers[s](x))
|
| 115 |
+
elif s == (len(self.stages) - 1):
|
| 116 |
+
seg_outputs.append(self.seg_layers[-1](x))
|
| 117 |
+
lres_input = x
|
| 118 |
+
|
| 119 |
+
# invert seg outputs so that the largest segmentation prediction is returned first
|
| 120 |
+
seg_outputs = seg_outputs[::-1]
|
| 121 |
+
|
| 122 |
+
if not self.deep_supervision:
|
| 123 |
+
r = seg_outputs[0]
|
| 124 |
+
else:
|
| 125 |
+
r = seg_outputs
|
| 126 |
+
return r
|
| 127 |
+
|
| 128 |
+
def compute_conv_feature_map_size(self, input_size):
|
| 129 |
+
"""
|
| 130 |
+
IMPORTANT: input_size is the input_size of the encoder!
|
| 131 |
+
:param input_size:
|
| 132 |
+
:return:
|
| 133 |
+
"""
|
| 134 |
+
# first we need to compute the skip sizes. Skip bottleneck because all output feature maps of our ops will at
|
| 135 |
+
# least have the size of the skip above that (therefore -1)
|
| 136 |
+
skip_sizes = []
|
| 137 |
+
for s in range(len(self.encoder.strides) - 1):
|
| 138 |
+
skip_sizes.append([i // j for i, j in zip(input_size, self.encoder.strides[s])])
|
| 139 |
+
input_size = skip_sizes[-1]
|
| 140 |
+
# print(skip_sizes)
|
| 141 |
+
|
| 142 |
+
assert len(skip_sizes) == len(self.stages)
|
| 143 |
+
|
| 144 |
+
# our ops are the other way around, so let's match things up
|
| 145 |
+
output = np.int64(0)
|
| 146 |
+
for s in range(len(self.stages)):
|
| 147 |
+
# print(skip_sizes[-(s+1)], self.encoder.output_channels[-(s+2)])
|
| 148 |
+
# conv blocks
|
| 149 |
+
output += self.stages[s].compute_conv_feature_map_size(skip_sizes[-(s + 1)])
|
| 150 |
+
# trans conv
|
| 151 |
+
output += np.prod([self.encoder.output_channels[-(s + 2)], *skip_sizes[-(s + 1)]], dtype=np.int64)
|
| 152 |
+
# segmentation
|
| 153 |
+
if self.deep_supervision or (s == (len(self.stages) - 1)):
|
| 154 |
+
output += np.prod([self.num_classes, *skip_sizes[-(s + 1)]], dtype=np.int64)
|
| 155 |
+
return output
|
RADAR_inference/dynamic_network_architectures/initialization/__init__.py
ADDED
|
File without changes
|
RADAR_inference/dynamic_network_architectures/initialization/weight_init.py
ADDED
|
@@ -0,0 +1,34 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from torch import nn
|
| 2 |
+
|
| 3 |
+
from dynamic_network_architectures.building_blocks.residual import BasicBlockD, BottleneckD
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
class InitWeights_He(object):
|
| 7 |
+
def __init__(self, neg_slope: float = 1e-2):
|
| 8 |
+
self.neg_slope = neg_slope
|
| 9 |
+
|
| 10 |
+
def __call__(self, module):
|
| 11 |
+
if isinstance(module, nn.Conv3d) or isinstance(module, nn.Conv2d) or isinstance(module, nn.ConvTranspose2d) or isinstance(module, nn.ConvTranspose3d):
|
| 12 |
+
module.weight = nn.init.kaiming_normal_(module.weight, a=self.neg_slope)
|
| 13 |
+
if module.bias is not None:
|
| 14 |
+
module.bias = nn.init.constant_(module.bias, 0)
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
class InitWeights_XavierUniform(object):
|
| 18 |
+
def __init__(self, gain: int = 1):
|
| 19 |
+
self.gain = gain
|
| 20 |
+
|
| 21 |
+
def __call__(self, module):
|
| 22 |
+
if isinstance(module, nn.Conv3d) or isinstance(module, nn.Conv2d) or isinstance(module, nn.ConvTranspose2d) or isinstance(module, nn.ConvTranspose3d):
|
| 23 |
+
module.weight = nn.init.xavier_uniform_(module.weight, self.gain)
|
| 24 |
+
if module.bias is not None:
|
| 25 |
+
module.bias = nn.init.constant_(module.bias, 0)
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def init_last_bn_before_add_to_0(module):
|
| 29 |
+
if isinstance(module, BasicBlockD):
|
| 30 |
+
module.conv2.norm.weight = nn.init.constant_(module.conv2.norm.weight, 0)
|
| 31 |
+
module.conv2.norm.bias = nn.init.constant_(module.conv2.norm.bias, 0)
|
| 32 |
+
if isinstance(module, BottleneckD):
|
| 33 |
+
module.conv3.norm.weight = nn.init.constant_(module.conv3.norm.weight, 0)
|
| 34 |
+
module.conv3.norm.bias = nn.init.constant_(module.conv3.norm.bias, 0)
|
RADAR_inference/dynamic_network_architectures/med.py
ADDED
|
@@ -0,0 +1,1502 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Copyright (c) 2022, salesforce.com, inc.
|
| 3 |
+
All rights reserved.
|
| 4 |
+
SPDX-License-Identifier: BSD-3-Clause
|
| 5 |
+
For full license text, see the LICENSE file in the repo root or https://opensource.org/licenses/BSD-3-Clause
|
| 6 |
+
|
| 7 |
+
Based on huggingface code base
|
| 8 |
+
https://github.com/huggingface/transformers/blob/v4.15.0/src/transformers/models/bert
|
| 9 |
+
"""
|
| 10 |
+
|
| 11 |
+
import math
|
| 12 |
+
import os
|
| 13 |
+
import warnings
|
| 14 |
+
from dataclasses import dataclass
|
| 15 |
+
from typing import Optional, Tuple
|
| 16 |
+
|
| 17 |
+
import torch
|
| 18 |
+
from torch import Tensor, device
|
| 19 |
+
import torch.utils.checkpoint
|
| 20 |
+
from torch import nn
|
| 21 |
+
from torch.nn import CrossEntropyLoss
|
| 22 |
+
import torch.nn.functional as F
|
| 23 |
+
from transformers import BatchEncoding, PreTrainedTokenizer
|
| 24 |
+
|
| 25 |
+
from transformers.activations import ACT2FN
|
| 26 |
+
from transformers.file_utils import (
|
| 27 |
+
ModelOutput,
|
| 28 |
+
)
|
| 29 |
+
from transformers.modeling_outputs import (
|
| 30 |
+
BaseModelOutputWithPastAndCrossAttentions,
|
| 31 |
+
BaseModelOutputWithPoolingAndCrossAttentions,
|
| 32 |
+
CausalLMOutputWithCrossAttentions,
|
| 33 |
+
MaskedLMOutput,
|
| 34 |
+
MultipleChoiceModelOutput,
|
| 35 |
+
NextSentencePredictorOutput,
|
| 36 |
+
QuestionAnsweringModelOutput,
|
| 37 |
+
SequenceClassifierOutput,
|
| 38 |
+
TokenClassifierOutput,
|
| 39 |
+
)
|
| 40 |
+
from transformers.modeling_utils import (
|
| 41 |
+
PreTrainedModel,
|
| 42 |
+
apply_chunking_to_forward,
|
| 43 |
+
find_pruneable_heads_and_indices,
|
| 44 |
+
prune_linear_layer,
|
| 45 |
+
)
|
| 46 |
+
from transformers.utils import logging
|
| 47 |
+
from transformers.models.bert.configuration_bert import BertConfig
|
| 48 |
+
|
| 49 |
+
# from lavis.models.base_model import BaseEncoder
|
| 50 |
+
|
| 51 |
+
logging.set_verbosity_error()
|
| 52 |
+
logger = logging.get_logger(__name__)
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
class BaseEncoder(nn.Module):
|
| 56 |
+
"""
|
| 57 |
+
Base class for primitive encoders, such as ViT, TimeSformer, etc.
|
| 58 |
+
"""
|
| 59 |
+
|
| 60 |
+
def __init__(self):
|
| 61 |
+
super().__init__()
|
| 62 |
+
|
| 63 |
+
def forward_features(self, samples, **kwargs):
|
| 64 |
+
raise NotImplementedError
|
| 65 |
+
|
| 66 |
+
@property
|
| 67 |
+
def device(self):
|
| 68 |
+
return list(self.parameters())[0].device
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
class BertEmbeddings(nn.Module):
|
| 72 |
+
"""Construct the embeddings from word and position embeddings."""
|
| 73 |
+
|
| 74 |
+
def __init__(self, config):
|
| 75 |
+
super().__init__()
|
| 76 |
+
self.word_embeddings = nn.Embedding(
|
| 77 |
+
config.vocab_size, config.hidden_size, padding_idx=config.pad_token_id
|
| 78 |
+
)
|
| 79 |
+
self.position_embeddings = nn.Embedding(
|
| 80 |
+
config.max_position_embeddings, config.hidden_size
|
| 81 |
+
)
|
| 82 |
+
|
| 83 |
+
if config.add_type_embeddings:
|
| 84 |
+
self.token_type_embeddings = nn.Embedding(
|
| 85 |
+
config.type_vocab_size, config.hidden_size
|
| 86 |
+
)
|
| 87 |
+
|
| 88 |
+
# self.LayerNorm is not snake-cased to stick with TensorFlow model variable name and be able to load
|
| 89 |
+
# any TensorFlow checkpoint file
|
| 90 |
+
self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
|
| 91 |
+
self.dropout = nn.Dropout(config.hidden_dropout_prob)
|
| 92 |
+
|
| 93 |
+
# position_ids (1, len position emb) is contiguous in memory and exported when serialized
|
| 94 |
+
self.register_buffer(
|
| 95 |
+
"position_ids", torch.arange(config.max_position_embeddings).expand((1, -1))
|
| 96 |
+
)
|
| 97 |
+
self.position_embedding_type = getattr(
|
| 98 |
+
config, "position_embedding_type", "absolute"
|
| 99 |
+
)
|
| 100 |
+
|
| 101 |
+
self.config = config
|
| 102 |
+
|
| 103 |
+
def forward(
|
| 104 |
+
self,
|
| 105 |
+
input_ids=None,
|
| 106 |
+
token_type_ids=None,
|
| 107 |
+
position_ids=None,
|
| 108 |
+
inputs_embeds=None,
|
| 109 |
+
past_key_values_length=0,
|
| 110 |
+
):
|
| 111 |
+
if input_ids is not None:
|
| 112 |
+
input_shape = input_ids.size()
|
| 113 |
+
else:
|
| 114 |
+
input_shape = inputs_embeds.size()[:-1]
|
| 115 |
+
|
| 116 |
+
seq_length = input_shape[1]
|
| 117 |
+
|
| 118 |
+
if position_ids is None:
|
| 119 |
+
position_ids = self.position_ids[
|
| 120 |
+
:, past_key_values_length : seq_length + past_key_values_length
|
| 121 |
+
]
|
| 122 |
+
|
| 123 |
+
if inputs_embeds is None:
|
| 124 |
+
inputs_embeds = self.word_embeddings(input_ids)
|
| 125 |
+
|
| 126 |
+
if token_type_ids is not None:
|
| 127 |
+
token_type_embeddings = self.token_type_embeddings(token_type_ids)
|
| 128 |
+
|
| 129 |
+
embeddings = inputs_embeds + token_type_embeddings
|
| 130 |
+
else:
|
| 131 |
+
embeddings = inputs_embeds
|
| 132 |
+
|
| 133 |
+
if self.position_embedding_type == "absolute":
|
| 134 |
+
position_embeddings = self.position_embeddings(position_ids)
|
| 135 |
+
embeddings += position_embeddings
|
| 136 |
+
embeddings = self.LayerNorm(embeddings)
|
| 137 |
+
embeddings = self.dropout(embeddings)
|
| 138 |
+
return embeddings
|
| 139 |
+
|
| 140 |
+
|
| 141 |
+
class BertSelfAttention(nn.Module):
|
| 142 |
+
def __init__(self, config, is_cross_attention):
|
| 143 |
+
super().__init__()
|
| 144 |
+
self.config = config
|
| 145 |
+
if config.hidden_size % config.num_attention_heads != 0 and not hasattr(
|
| 146 |
+
config, "embedding_size"
|
| 147 |
+
):
|
| 148 |
+
raise ValueError(
|
| 149 |
+
"The hidden size (%d) is not a multiple of the number of attention "
|
| 150 |
+
"heads (%d)" % (config.hidden_size, config.num_attention_heads)
|
| 151 |
+
)
|
| 152 |
+
|
| 153 |
+
self.num_attention_heads = config.num_attention_heads
|
| 154 |
+
self.attention_head_size = int(config.hidden_size / config.num_attention_heads)
|
| 155 |
+
self.all_head_size = self.num_attention_heads * self.attention_head_size
|
| 156 |
+
|
| 157 |
+
self.query = nn.Linear(config.hidden_size, self.all_head_size)
|
| 158 |
+
if is_cross_attention:
|
| 159 |
+
self.key = nn.Linear(config.encoder_width, self.all_head_size)
|
| 160 |
+
self.value = nn.Linear(config.encoder_width, self.all_head_size)
|
| 161 |
+
else:
|
| 162 |
+
self.key = nn.Linear(config.hidden_size, self.all_head_size)
|
| 163 |
+
self.value = nn.Linear(config.hidden_size, self.all_head_size)
|
| 164 |
+
|
| 165 |
+
self.dropout = nn.Dropout(config.attention_probs_dropout_prob)
|
| 166 |
+
self.position_embedding_type = getattr(
|
| 167 |
+
config, "position_embedding_type", "absolute"
|
| 168 |
+
)
|
| 169 |
+
if (
|
| 170 |
+
self.position_embedding_type == "relative_key"
|
| 171 |
+
or self.position_embedding_type == "relative_key_query"
|
| 172 |
+
):
|
| 173 |
+
self.max_position_embeddings = config.max_position_embeddings
|
| 174 |
+
self.distance_embedding = nn.Embedding(
|
| 175 |
+
2 * config.max_position_embeddings - 1, self.attention_head_size
|
| 176 |
+
)
|
| 177 |
+
self.save_attention = False
|
| 178 |
+
|
| 179 |
+
def save_attn_gradients(self, attn_gradients):
|
| 180 |
+
self.attn_gradients = attn_gradients
|
| 181 |
+
|
| 182 |
+
def get_attn_gradients(self):
|
| 183 |
+
return self.attn_gradients
|
| 184 |
+
|
| 185 |
+
def save_attention_map(self, attention_map):
|
| 186 |
+
self.attention_map = attention_map
|
| 187 |
+
|
| 188 |
+
def get_attention_map(self):
|
| 189 |
+
return self.attention_map
|
| 190 |
+
|
| 191 |
+
def transpose_for_scores(self, x):
|
| 192 |
+
new_x_shape = x.size()[:-1] + (
|
| 193 |
+
self.num_attention_heads,
|
| 194 |
+
self.attention_head_size,
|
| 195 |
+
)
|
| 196 |
+
x = x.view(*new_x_shape)
|
| 197 |
+
return x.permute(0, 2, 1, 3)
|
| 198 |
+
|
| 199 |
+
def forward(
|
| 200 |
+
self,
|
| 201 |
+
hidden_states,
|
| 202 |
+
attention_mask=None,
|
| 203 |
+
head_mask=None,
|
| 204 |
+
encoder_hidden_states=None,
|
| 205 |
+
encoder_attention_mask=None,
|
| 206 |
+
past_key_value=None,
|
| 207 |
+
output_attentions=False,
|
| 208 |
+
):
|
| 209 |
+
mixed_query_layer = self.query(hidden_states)
|
| 210 |
+
|
| 211 |
+
# If this is instantiated as a cross-attention module, the keys
|
| 212 |
+
# and values come from an encoder; the attention mask needs to be
|
| 213 |
+
# such that the encoder's padding tokens are not attended to.
|
| 214 |
+
is_cross_attention = encoder_hidden_states is not None
|
| 215 |
+
|
| 216 |
+
if is_cross_attention:
|
| 217 |
+
key_layer = self.transpose_for_scores(self.key(encoder_hidden_states))
|
| 218 |
+
value_layer = self.transpose_for_scores(self.value(encoder_hidden_states))
|
| 219 |
+
attention_mask = encoder_attention_mask
|
| 220 |
+
elif past_key_value is not None:
|
| 221 |
+
key_layer = self.transpose_for_scores(self.key(hidden_states))
|
| 222 |
+
value_layer = self.transpose_for_scores(self.value(hidden_states))
|
| 223 |
+
key_layer = torch.cat([past_key_value[0], key_layer], dim=2)
|
| 224 |
+
value_layer = torch.cat([past_key_value[1], value_layer], dim=2)
|
| 225 |
+
else:
|
| 226 |
+
key_layer = self.transpose_for_scores(self.key(hidden_states))
|
| 227 |
+
value_layer = self.transpose_for_scores(self.value(hidden_states))
|
| 228 |
+
|
| 229 |
+
query_layer = self.transpose_for_scores(mixed_query_layer)
|
| 230 |
+
|
| 231 |
+
past_key_value = (key_layer, value_layer)
|
| 232 |
+
|
| 233 |
+
# Take the dot product between "query" and "key" to get the raw attention scores.
|
| 234 |
+
attention_scores = torch.matmul(query_layer, key_layer.transpose(-1, -2))
|
| 235 |
+
|
| 236 |
+
if (
|
| 237 |
+
self.position_embedding_type == "relative_key"
|
| 238 |
+
or self.position_embedding_type == "relative_key_query"
|
| 239 |
+
):
|
| 240 |
+
seq_length = hidden_states.size()[1]
|
| 241 |
+
position_ids_l = torch.arange(
|
| 242 |
+
seq_length, dtype=torch.long, device=hidden_states.device
|
| 243 |
+
).view(-1, 1)
|
| 244 |
+
position_ids_r = torch.arange(
|
| 245 |
+
seq_length, dtype=torch.long, device=hidden_states.device
|
| 246 |
+
).view(1, -1)
|
| 247 |
+
distance = position_ids_l - position_ids_r
|
| 248 |
+
positional_embedding = self.distance_embedding(
|
| 249 |
+
distance + self.max_position_embeddings - 1
|
| 250 |
+
)
|
| 251 |
+
positional_embedding = positional_embedding.to(
|
| 252 |
+
dtype=query_layer.dtype
|
| 253 |
+
) # fp16 compatibility
|
| 254 |
+
|
| 255 |
+
if self.position_embedding_type == "relative_key":
|
| 256 |
+
relative_position_scores = torch.einsum(
|
| 257 |
+
"bhld,lrd->bhlr", query_layer, positional_embedding
|
| 258 |
+
)
|
| 259 |
+
attention_scores = attention_scores + relative_position_scores
|
| 260 |
+
elif self.position_embedding_type == "relative_key_query":
|
| 261 |
+
relative_position_scores_query = torch.einsum(
|
| 262 |
+
"bhld,lrd->bhlr", query_layer, positional_embedding
|
| 263 |
+
)
|
| 264 |
+
relative_position_scores_key = torch.einsum(
|
| 265 |
+
"bhrd,lrd->bhlr", key_layer, positional_embedding
|
| 266 |
+
)
|
| 267 |
+
attention_scores = (
|
| 268 |
+
attention_scores
|
| 269 |
+
+ relative_position_scores_query
|
| 270 |
+
+ relative_position_scores_key
|
| 271 |
+
)
|
| 272 |
+
|
| 273 |
+
attention_scores = attention_scores / math.sqrt(self.attention_head_size)
|
| 274 |
+
if attention_mask is not None:
|
| 275 |
+
# Apply the attention mask is (precomputed for all layers in BertModel forward() function)
|
| 276 |
+
attention_scores = attention_scores + attention_mask
|
| 277 |
+
|
| 278 |
+
# Normalize the attention scores to probabilities.
|
| 279 |
+
attention_probs = nn.Softmax(dim=-1)(attention_scores)
|
| 280 |
+
|
| 281 |
+
if is_cross_attention and self.save_attention:
|
| 282 |
+
self.save_attention_map(attention_probs)
|
| 283 |
+
attention_probs.register_hook(self.save_attn_gradients)
|
| 284 |
+
|
| 285 |
+
# This is actually dropping out entire tokens to attend to, which might
|
| 286 |
+
# seem a bit unusual, but is taken from the original Transformer paper.
|
| 287 |
+
attention_probs_dropped = self.dropout(attention_probs)
|
| 288 |
+
|
| 289 |
+
# Mask heads if we want to
|
| 290 |
+
if head_mask is not None:
|
| 291 |
+
attention_probs_dropped = attention_probs_dropped * head_mask
|
| 292 |
+
|
| 293 |
+
context_layer = torch.matmul(attention_probs_dropped, value_layer)
|
| 294 |
+
|
| 295 |
+
context_layer = context_layer.permute(0, 2, 1, 3).contiguous()
|
| 296 |
+
new_context_layer_shape = context_layer.size()[:-2] + (self.all_head_size,)
|
| 297 |
+
context_layer = context_layer.view(*new_context_layer_shape)
|
| 298 |
+
|
| 299 |
+
outputs = (
|
| 300 |
+
(context_layer, attention_probs) if output_attentions else (context_layer,)
|
| 301 |
+
)
|
| 302 |
+
|
| 303 |
+
outputs = outputs + (past_key_value,)
|
| 304 |
+
return outputs
|
| 305 |
+
|
| 306 |
+
|
| 307 |
+
class BertSelfOutput(nn.Module):
|
| 308 |
+
def __init__(self, config):
|
| 309 |
+
super().__init__()
|
| 310 |
+
self.dense = nn.Linear(config.hidden_size, config.hidden_size)
|
| 311 |
+
self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
|
| 312 |
+
self.dropout = nn.Dropout(config.hidden_dropout_prob)
|
| 313 |
+
|
| 314 |
+
def forward(self, hidden_states, input_tensor):
|
| 315 |
+
hidden_states = self.dense(hidden_states)
|
| 316 |
+
hidden_states = self.dropout(hidden_states)
|
| 317 |
+
hidden_states = self.LayerNorm(hidden_states + input_tensor)
|
| 318 |
+
return hidden_states
|
| 319 |
+
|
| 320 |
+
|
| 321 |
+
class BertAttention(nn.Module):
|
| 322 |
+
def __init__(self, config, is_cross_attention=False):
|
| 323 |
+
super().__init__()
|
| 324 |
+
self.self = BertSelfAttention(config, is_cross_attention)
|
| 325 |
+
self.output = BertSelfOutput(config)
|
| 326 |
+
self.pruned_heads = set()
|
| 327 |
+
|
| 328 |
+
def prune_heads(self, heads):
|
| 329 |
+
if len(heads) == 0:
|
| 330 |
+
return
|
| 331 |
+
heads, index = find_pruneable_heads_and_indices(
|
| 332 |
+
heads,
|
| 333 |
+
self.self.num_attention_heads,
|
| 334 |
+
self.self.attention_head_size,
|
| 335 |
+
self.pruned_heads,
|
| 336 |
+
)
|
| 337 |
+
|
| 338 |
+
# Prune linear layers
|
| 339 |
+
self.self.query = prune_linear_layer(self.self.query, index)
|
| 340 |
+
self.self.key = prune_linear_layer(self.self.key, index)
|
| 341 |
+
self.self.value = prune_linear_layer(self.self.value, index)
|
| 342 |
+
self.output.dense = prune_linear_layer(self.output.dense, index, dim=1)
|
| 343 |
+
|
| 344 |
+
# Update hyper params and store pruned heads
|
| 345 |
+
self.self.num_attention_heads = self.self.num_attention_heads - len(heads)
|
| 346 |
+
self.self.all_head_size = (
|
| 347 |
+
self.self.attention_head_size * self.self.num_attention_heads
|
| 348 |
+
)
|
| 349 |
+
self.pruned_heads = self.pruned_heads.union(heads)
|
| 350 |
+
|
| 351 |
+
def forward(
|
| 352 |
+
self,
|
| 353 |
+
hidden_states,
|
| 354 |
+
attention_mask=None,
|
| 355 |
+
head_mask=None,
|
| 356 |
+
encoder_hidden_states=None,
|
| 357 |
+
encoder_attention_mask=None,
|
| 358 |
+
past_key_value=None,
|
| 359 |
+
output_attentions=False,
|
| 360 |
+
):
|
| 361 |
+
self_outputs = self.self(
|
| 362 |
+
hidden_states,
|
| 363 |
+
attention_mask,
|
| 364 |
+
head_mask,
|
| 365 |
+
encoder_hidden_states,
|
| 366 |
+
encoder_attention_mask,
|
| 367 |
+
past_key_value,
|
| 368 |
+
output_attentions,
|
| 369 |
+
)
|
| 370 |
+
attention_output = self.output(self_outputs[0], hidden_states)
|
| 371 |
+
outputs = (attention_output,) + self_outputs[
|
| 372 |
+
1:
|
| 373 |
+
] # add attentions if we output them
|
| 374 |
+
return outputs
|
| 375 |
+
|
| 376 |
+
|
| 377 |
+
class BertIntermediate(nn.Module):
|
| 378 |
+
def __init__(self, config):
|
| 379 |
+
super().__init__()
|
| 380 |
+
self.dense = nn.Linear(config.hidden_size, config.intermediate_size)
|
| 381 |
+
if isinstance(config.hidden_act, str):
|
| 382 |
+
self.intermediate_act_fn = ACT2FN[config.hidden_act]
|
| 383 |
+
else:
|
| 384 |
+
self.intermediate_act_fn = config.hidden_act
|
| 385 |
+
|
| 386 |
+
def forward(self, hidden_states):
|
| 387 |
+
hidden_states = self.dense(hidden_states)
|
| 388 |
+
hidden_states = self.intermediate_act_fn(hidden_states)
|
| 389 |
+
return hidden_states
|
| 390 |
+
|
| 391 |
+
|
| 392 |
+
class BertOutput(nn.Module):
|
| 393 |
+
def __init__(self, config):
|
| 394 |
+
super().__init__()
|
| 395 |
+
self.dense = nn.Linear(config.intermediate_size, config.hidden_size)
|
| 396 |
+
self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
|
| 397 |
+
self.dropout = nn.Dropout(config.hidden_dropout_prob)
|
| 398 |
+
|
| 399 |
+
def forward(self, hidden_states, input_tensor):
|
| 400 |
+
hidden_states = self.dense(hidden_states)
|
| 401 |
+
hidden_states = self.dropout(hidden_states)
|
| 402 |
+
hidden_states = self.LayerNorm(hidden_states + input_tensor)
|
| 403 |
+
return hidden_states
|
| 404 |
+
|
| 405 |
+
|
| 406 |
+
class BertLayer(nn.Module):
|
| 407 |
+
def __init__(self, config, layer_num):
|
| 408 |
+
super().__init__()
|
| 409 |
+
self.config = config
|
| 410 |
+
self.chunk_size_feed_forward = config.chunk_size_feed_forward
|
| 411 |
+
self.seq_len_dim = 1
|
| 412 |
+
self.attention = BertAttention(config)
|
| 413 |
+
self.layer_num = layer_num
|
| 414 |
+
|
| 415 |
+
# compatibility for ALBEF and BLIP
|
| 416 |
+
try:
|
| 417 |
+
# ALBEF & ALPRO
|
| 418 |
+
fusion_layer = self.config.fusion_layer
|
| 419 |
+
add_cross_attention = (
|
| 420 |
+
fusion_layer <= layer_num and self.config.add_cross_attention
|
| 421 |
+
)
|
| 422 |
+
|
| 423 |
+
self.fusion_layer = fusion_layer
|
| 424 |
+
except AttributeError:
|
| 425 |
+
# BLIP
|
| 426 |
+
self.fusion_layer = self.config.num_hidden_layers
|
| 427 |
+
add_cross_attention = self.config.add_cross_attention
|
| 428 |
+
|
| 429 |
+
# if self.config.add_cross_attention:
|
| 430 |
+
if add_cross_attention:
|
| 431 |
+
self.crossattention = BertAttention(
|
| 432 |
+
config, is_cross_attention=self.config.add_cross_attention
|
| 433 |
+
)
|
| 434 |
+
self.intermediate = BertIntermediate(config)
|
| 435 |
+
self.output = BertOutput(config)
|
| 436 |
+
|
| 437 |
+
def forward(
|
| 438 |
+
self,
|
| 439 |
+
hidden_states,
|
| 440 |
+
attention_mask=None,
|
| 441 |
+
head_mask=None,
|
| 442 |
+
encoder_hidden_states=None,
|
| 443 |
+
encoder_attention_mask=None,
|
| 444 |
+
past_key_value=None,
|
| 445 |
+
output_attentions=False,
|
| 446 |
+
mode=None,
|
| 447 |
+
):
|
| 448 |
+
# decoder uni-directional self-attention cached key/values tuple is at positions 1,2
|
| 449 |
+
self_attn_past_key_value = (
|
| 450 |
+
past_key_value[:2] if past_key_value is not None else None
|
| 451 |
+
)
|
| 452 |
+
self_attention_outputs = self.attention(
|
| 453 |
+
hidden_states,
|
| 454 |
+
attention_mask,
|
| 455 |
+
head_mask,
|
| 456 |
+
output_attentions=output_attentions,
|
| 457 |
+
past_key_value=self_attn_past_key_value,
|
| 458 |
+
)
|
| 459 |
+
attention_output = self_attention_outputs[0]
|
| 460 |
+
|
| 461 |
+
outputs = self_attention_outputs[1:-1]
|
| 462 |
+
present_key_value = self_attention_outputs[-1]
|
| 463 |
+
|
| 464 |
+
# TODO line 482 in albef/models/xbert.py
|
| 465 |
+
# compatibility for ALBEF and BLIP
|
| 466 |
+
if mode in ["multimodal", "fusion"] and hasattr(self, "crossattention"):
|
| 467 |
+
assert (
|
| 468 |
+
encoder_hidden_states is not None
|
| 469 |
+
), "encoder_hidden_states must be given for cross-attention layers"
|
| 470 |
+
|
| 471 |
+
if isinstance(encoder_hidden_states, list):
|
| 472 |
+
cross_attention_outputs = self.crossattention(
|
| 473 |
+
attention_output,
|
| 474 |
+
attention_mask,
|
| 475 |
+
head_mask,
|
| 476 |
+
encoder_hidden_states[
|
| 477 |
+
(self.layer_num - self.fusion_layer)
|
| 478 |
+
% len(encoder_hidden_states)
|
| 479 |
+
],
|
| 480 |
+
encoder_attention_mask[
|
| 481 |
+
(self.layer_num - self.fusion_layer)
|
| 482 |
+
% len(encoder_hidden_states)
|
| 483 |
+
],
|
| 484 |
+
output_attentions=output_attentions,
|
| 485 |
+
)
|
| 486 |
+
attention_output = cross_attention_outputs[0]
|
| 487 |
+
outputs = outputs + cross_attention_outputs[1:-1]
|
| 488 |
+
|
| 489 |
+
else:
|
| 490 |
+
cross_attention_outputs = self.crossattention(
|
| 491 |
+
attention_output,
|
| 492 |
+
attention_mask,
|
| 493 |
+
head_mask,
|
| 494 |
+
encoder_hidden_states,
|
| 495 |
+
encoder_attention_mask,
|
| 496 |
+
output_attentions=output_attentions,
|
| 497 |
+
)
|
| 498 |
+
attention_output = cross_attention_outputs[0]
|
| 499 |
+
outputs = (
|
| 500 |
+
outputs + cross_attention_outputs[1:-1]
|
| 501 |
+
) # add cross attentions if we output attention weights
|
| 502 |
+
layer_output = apply_chunking_to_forward(
|
| 503 |
+
self.feed_forward_chunk,
|
| 504 |
+
self.chunk_size_feed_forward,
|
| 505 |
+
self.seq_len_dim,
|
| 506 |
+
attention_output,
|
| 507 |
+
)
|
| 508 |
+
outputs = (layer_output,) + outputs
|
| 509 |
+
|
| 510 |
+
outputs = outputs + (present_key_value,)
|
| 511 |
+
|
| 512 |
+
return outputs
|
| 513 |
+
|
| 514 |
+
def feed_forward_chunk(self, attention_output):
|
| 515 |
+
intermediate_output = self.intermediate(attention_output)
|
| 516 |
+
layer_output = self.output(intermediate_output, attention_output)
|
| 517 |
+
return layer_output
|
| 518 |
+
|
| 519 |
+
|
| 520 |
+
class BertEncoder(nn.Module):
|
| 521 |
+
def __init__(self, config):
|
| 522 |
+
super().__init__()
|
| 523 |
+
self.config = config
|
| 524 |
+
self.layer = nn.ModuleList(
|
| 525 |
+
[BertLayer(config, i) for i in range(config.num_hidden_layers)]
|
| 526 |
+
)
|
| 527 |
+
self.gradient_checkpointing = False
|
| 528 |
+
|
| 529 |
+
def forward(
|
| 530 |
+
self,
|
| 531 |
+
hidden_states,
|
| 532 |
+
attention_mask=None,
|
| 533 |
+
head_mask=None,
|
| 534 |
+
encoder_hidden_states=None,
|
| 535 |
+
encoder_attention_mask=None,
|
| 536 |
+
past_key_values=None,
|
| 537 |
+
use_cache=None,
|
| 538 |
+
output_attentions=False,
|
| 539 |
+
output_hidden_states=False,
|
| 540 |
+
return_dict=True,
|
| 541 |
+
mode="multimodal",
|
| 542 |
+
):
|
| 543 |
+
all_hidden_states = () if output_hidden_states else None
|
| 544 |
+
all_self_attentions = () if output_attentions else None
|
| 545 |
+
all_cross_attentions = (
|
| 546 |
+
() if output_attentions and self.config.add_cross_attention else None
|
| 547 |
+
)
|
| 548 |
+
|
| 549 |
+
next_decoder_cache = () if use_cache else None
|
| 550 |
+
|
| 551 |
+
try:
|
| 552 |
+
# ALBEF
|
| 553 |
+
fusion_layer = self.config.fusion_layer
|
| 554 |
+
except AttributeError:
|
| 555 |
+
# BLIP
|
| 556 |
+
fusion_layer = self.config.num_hidden_layers
|
| 557 |
+
|
| 558 |
+
if mode == "text":
|
| 559 |
+
start_layer = 0
|
| 560 |
+
# output_layer = self.config.fusion_layer
|
| 561 |
+
output_layer = fusion_layer
|
| 562 |
+
|
| 563 |
+
elif mode == "fusion":
|
| 564 |
+
# start_layer = self.config.fusion_layer
|
| 565 |
+
start_layer = fusion_layer
|
| 566 |
+
output_layer = self.config.num_hidden_layers
|
| 567 |
+
|
| 568 |
+
elif mode == "multimodal":
|
| 569 |
+
start_layer = 0
|
| 570 |
+
output_layer = self.config.num_hidden_layers
|
| 571 |
+
|
| 572 |
+
# compatibility for ALBEF and BLIP
|
| 573 |
+
# for i in range(self.config.num_hidden_layers):
|
| 574 |
+
for i in range(start_layer, output_layer):
|
| 575 |
+
layer_module = self.layer[i]
|
| 576 |
+
if output_hidden_states:
|
| 577 |
+
all_hidden_states = all_hidden_states + (hidden_states,)
|
| 578 |
+
|
| 579 |
+
layer_head_mask = head_mask[i] if head_mask is not None else None
|
| 580 |
+
past_key_value = past_key_values[i] if past_key_values is not None else None
|
| 581 |
+
|
| 582 |
+
# TODO pay attention to this.
|
| 583 |
+
if self.gradient_checkpointing and self.training:
|
| 584 |
+
|
| 585 |
+
if use_cache:
|
| 586 |
+
logger.warn(
|
| 587 |
+
"`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`..."
|
| 588 |
+
)
|
| 589 |
+
use_cache = False
|
| 590 |
+
|
| 591 |
+
def create_custom_forward(module):
|
| 592 |
+
def custom_forward(*inputs):
|
| 593 |
+
return module(*inputs, past_key_value, output_attentions)
|
| 594 |
+
|
| 595 |
+
return custom_forward
|
| 596 |
+
|
| 597 |
+
layer_outputs = torch.utils.checkpoint.checkpoint(
|
| 598 |
+
create_custom_forward(layer_module),
|
| 599 |
+
hidden_states,
|
| 600 |
+
attention_mask,
|
| 601 |
+
layer_head_mask,
|
| 602 |
+
encoder_hidden_states,
|
| 603 |
+
encoder_attention_mask,
|
| 604 |
+
mode=mode,
|
| 605 |
+
)
|
| 606 |
+
else:
|
| 607 |
+
layer_outputs = layer_module(
|
| 608 |
+
hidden_states,
|
| 609 |
+
attention_mask,
|
| 610 |
+
layer_head_mask,
|
| 611 |
+
encoder_hidden_states,
|
| 612 |
+
encoder_attention_mask,
|
| 613 |
+
past_key_value,
|
| 614 |
+
output_attentions,
|
| 615 |
+
mode=mode,
|
| 616 |
+
)
|
| 617 |
+
|
| 618 |
+
hidden_states = layer_outputs[0]
|
| 619 |
+
if use_cache:
|
| 620 |
+
next_decoder_cache += (layer_outputs[-1],)
|
| 621 |
+
if output_attentions:
|
| 622 |
+
all_self_attentions = all_self_attentions + (layer_outputs[1],)
|
| 623 |
+
|
| 624 |
+
if output_hidden_states:
|
| 625 |
+
all_hidden_states = all_hidden_states + (hidden_states,)
|
| 626 |
+
|
| 627 |
+
if not return_dict:
|
| 628 |
+
return tuple(
|
| 629 |
+
v
|
| 630 |
+
for v in [
|
| 631 |
+
hidden_states,
|
| 632 |
+
next_decoder_cache,
|
| 633 |
+
all_hidden_states,
|
| 634 |
+
all_self_attentions,
|
| 635 |
+
all_cross_attentions,
|
| 636 |
+
]
|
| 637 |
+
if v is not None
|
| 638 |
+
)
|
| 639 |
+
return BaseModelOutputWithPastAndCrossAttentions(
|
| 640 |
+
last_hidden_state=hidden_states,
|
| 641 |
+
past_key_values=next_decoder_cache,
|
| 642 |
+
hidden_states=all_hidden_states,
|
| 643 |
+
attentions=all_self_attentions,
|
| 644 |
+
cross_attentions=all_cross_attentions,
|
| 645 |
+
)
|
| 646 |
+
|
| 647 |
+
|
| 648 |
+
class BertPooler(nn.Module):
|
| 649 |
+
def __init__(self, config):
|
| 650 |
+
super().__init__()
|
| 651 |
+
self.dense = nn.Linear(config.hidden_size, config.hidden_size)
|
| 652 |
+
self.activation = nn.Tanh()
|
| 653 |
+
|
| 654 |
+
def forward(self, hidden_states):
|
| 655 |
+
# We "pool" the model by simply taking the hidden state corresponding
|
| 656 |
+
# to the first token.
|
| 657 |
+
first_token_tensor = hidden_states[:, 0]
|
| 658 |
+
pooled_output = self.dense(first_token_tensor)
|
| 659 |
+
pooled_output = self.activation(pooled_output)
|
| 660 |
+
return pooled_output
|
| 661 |
+
|
| 662 |
+
|
| 663 |
+
class BertPredictionHeadTransform(nn.Module):
|
| 664 |
+
def __init__(self, config):
|
| 665 |
+
super().__init__()
|
| 666 |
+
self.dense = nn.Linear(config.hidden_size, config.hidden_size)
|
| 667 |
+
if isinstance(config.hidden_act, str):
|
| 668 |
+
self.transform_act_fn = ACT2FN[config.hidden_act]
|
| 669 |
+
else:
|
| 670 |
+
self.transform_act_fn = config.hidden_act
|
| 671 |
+
self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
|
| 672 |
+
|
| 673 |
+
def forward(self, hidden_states):
|
| 674 |
+
hidden_states = self.dense(hidden_states)
|
| 675 |
+
hidden_states = self.transform_act_fn(hidden_states)
|
| 676 |
+
hidden_states = self.LayerNorm(hidden_states)
|
| 677 |
+
return hidden_states
|
| 678 |
+
|
| 679 |
+
|
| 680 |
+
class BertLMPredictionHead(nn.Module):
|
| 681 |
+
def __init__(self, config):
|
| 682 |
+
super().__init__()
|
| 683 |
+
self.transform = BertPredictionHeadTransform(config)
|
| 684 |
+
|
| 685 |
+
# The output weights are the same as the input embeddings, but there is
|
| 686 |
+
# an output-only bias for each token.
|
| 687 |
+
self.decoder = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
|
| 688 |
+
|
| 689 |
+
self.bias = nn.Parameter(torch.zeros(config.vocab_size))
|
| 690 |
+
|
| 691 |
+
# Need a link between the two variables so that the bias is correctly resized with `resize_token_embeddings`
|
| 692 |
+
self.decoder.bias = self.bias
|
| 693 |
+
|
| 694 |
+
def forward(self, hidden_states):
|
| 695 |
+
hidden_states = self.transform(hidden_states)
|
| 696 |
+
hidden_states = self.decoder(hidden_states)
|
| 697 |
+
return hidden_states
|
| 698 |
+
|
| 699 |
+
|
| 700 |
+
class BertOnlyMLMHead(nn.Module):
|
| 701 |
+
def __init__(self, config):
|
| 702 |
+
super().__init__()
|
| 703 |
+
self.predictions = BertLMPredictionHead(config)
|
| 704 |
+
|
| 705 |
+
def forward(self, sequence_output):
|
| 706 |
+
prediction_scores = self.predictions(sequence_output)
|
| 707 |
+
return prediction_scores
|
| 708 |
+
|
| 709 |
+
|
| 710 |
+
class BertPreTrainedModel(PreTrainedModel):
|
| 711 |
+
"""
|
| 712 |
+
An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
|
| 713 |
+
models.
|
| 714 |
+
"""
|
| 715 |
+
|
| 716 |
+
config_class = BertConfig
|
| 717 |
+
base_model_prefix = "bert"
|
| 718 |
+
_keys_to_ignore_on_load_missing = [r"position_ids"]
|
| 719 |
+
|
| 720 |
+
def _init_weights(self, module):
|
| 721 |
+
"""Initialize the weights"""
|
| 722 |
+
if isinstance(module, (nn.Linear, nn.Embedding)):
|
| 723 |
+
# Slightly different from the TF version which uses truncated_normal for initialization
|
| 724 |
+
# cf https://github.com/pytorch/pytorch/pull/5617
|
| 725 |
+
module.weight.data.normal_(mean=0.0, std=self.config.initializer_range)
|
| 726 |
+
elif isinstance(module, nn.LayerNorm):
|
| 727 |
+
module.bias.data.zero_()
|
| 728 |
+
module.weight.data.fill_(1.0)
|
| 729 |
+
if isinstance(module, nn.Linear) and module.bias is not None:
|
| 730 |
+
module.bias.data.zero_()
|
| 731 |
+
|
| 732 |
+
|
| 733 |
+
class BertModel(BertPreTrainedModel):
|
| 734 |
+
"""
|
| 735 |
+
The model can behave as an encoder (with only self-attention) as well as a decoder, in which case a layer of
|
| 736 |
+
cross-attention is added between the self-attention layers, following the architecture described in `Attention is
|
| 737 |
+
all you need <https://arxiv.org/abs/1706.03762>`__ by Ashish Vaswani, Noam Shazeer, Niki Parmar, Jakob Uszkoreit,
|
| 738 |
+
Llion Jones, Aidan N. Gomez, Lukasz Kaiser and Illia Polosukhin.
|
| 739 |
+
argument and :obj:`add_cross_attention` set to :obj:`True`; an :obj:`encoder_hidden_states` is then expected as an
|
| 740 |
+
input to the forward pass.
|
| 741 |
+
"""
|
| 742 |
+
|
| 743 |
+
def __init__(self, config, add_pooling_layer=True):
|
| 744 |
+
super().__init__(config)
|
| 745 |
+
self.config = config
|
| 746 |
+
|
| 747 |
+
self.embeddings = BertEmbeddings(config)
|
| 748 |
+
|
| 749 |
+
self.encoder = BertEncoder(config)
|
| 750 |
+
|
| 751 |
+
self.pooler = BertPooler(config) if add_pooling_layer else None
|
| 752 |
+
|
| 753 |
+
self.init_weights()
|
| 754 |
+
|
| 755 |
+
def get_input_embeddings(self):
|
| 756 |
+
return self.embeddings.word_embeddings
|
| 757 |
+
|
| 758 |
+
def set_input_embeddings(self, value):
|
| 759 |
+
self.embeddings.word_embeddings = value
|
| 760 |
+
|
| 761 |
+
def _prune_heads(self, heads_to_prune):
|
| 762 |
+
"""
|
| 763 |
+
Prunes heads of the model. heads_to_prune: dict of {layer_num: list of heads to prune in this layer} See base
|
| 764 |
+
class PreTrainedModel
|
| 765 |
+
"""
|
| 766 |
+
for layer, heads in heads_to_prune.items():
|
| 767 |
+
self.encoder.layer[layer].attention.prune_heads(heads)
|
| 768 |
+
|
| 769 |
+
def get_extended_attention_mask(
|
| 770 |
+
self,
|
| 771 |
+
attention_mask: Tensor,
|
| 772 |
+
input_shape: Tuple[int],
|
| 773 |
+
device: device,
|
| 774 |
+
is_decoder: bool,
|
| 775 |
+
) -> Tensor:
|
| 776 |
+
"""
|
| 777 |
+
Makes broadcastable attention and causal masks so that future and masked tokens are ignored.
|
| 778 |
+
|
| 779 |
+
Arguments:
|
| 780 |
+
attention_mask (:obj:`torch.Tensor`):
|
| 781 |
+
Mask with ones indicating tokens to attend to, zeros for tokens to ignore.
|
| 782 |
+
input_shape (:obj:`Tuple[int]`):
|
| 783 |
+
The shape of the input to the model.
|
| 784 |
+
device: (:obj:`torch.device`):
|
| 785 |
+
The device of the input to the model.
|
| 786 |
+
|
| 787 |
+
Returns:
|
| 788 |
+
:obj:`torch.Tensor` The extended attention mask, with a the same dtype as :obj:`attention_mask.dtype`.
|
| 789 |
+
"""
|
| 790 |
+
# We can provide a self-attention mask of dimensions [batch_size, from_seq_length, to_seq_length]
|
| 791 |
+
# ourselves in which case we just need to make it broadcastable to all heads.
|
| 792 |
+
if attention_mask.dim() == 3:
|
| 793 |
+
extended_attention_mask = attention_mask[:, None, :, :]
|
| 794 |
+
elif attention_mask.dim() == 2:
|
| 795 |
+
# Provided a padding mask of dimensions [batch_size, seq_length]
|
| 796 |
+
# - if the model is a decoder, apply a causal mask in addition to the padding mask
|
| 797 |
+
# - if the model is an encoder, make the mask broadcastable to [batch_size, num_heads, seq_length, seq_length]
|
| 798 |
+
if is_decoder:
|
| 799 |
+
batch_size, seq_length = input_shape
|
| 800 |
+
|
| 801 |
+
seq_ids = torch.arange(seq_length, device=device)
|
| 802 |
+
causal_mask = (
|
| 803 |
+
seq_ids[None, None, :].repeat(batch_size, seq_length, 1)
|
| 804 |
+
<= seq_ids[None, :, None]
|
| 805 |
+
)
|
| 806 |
+
# in case past_key_values are used we need to add a prefix ones mask to the causal mask
|
| 807 |
+
# causal and attention masks must have same type with pytorch version < 1.3
|
| 808 |
+
causal_mask = causal_mask.to(attention_mask.dtype)
|
| 809 |
+
|
| 810 |
+
if causal_mask.shape[1] < attention_mask.shape[1]:
|
| 811 |
+
prefix_seq_len = attention_mask.shape[1] - causal_mask.shape[1]
|
| 812 |
+
causal_mask = torch.cat(
|
| 813 |
+
[
|
| 814 |
+
torch.ones(
|
| 815 |
+
(batch_size, seq_length, prefix_seq_len),
|
| 816 |
+
device=device,
|
| 817 |
+
dtype=causal_mask.dtype,
|
| 818 |
+
),
|
| 819 |
+
causal_mask,
|
| 820 |
+
],
|
| 821 |
+
axis=-1,
|
| 822 |
+
)
|
| 823 |
+
|
| 824 |
+
extended_attention_mask = (
|
| 825 |
+
causal_mask[:, None, :, :] * attention_mask[:, None, None, :]
|
| 826 |
+
)
|
| 827 |
+
else:
|
| 828 |
+
extended_attention_mask = attention_mask[:, None, None, :]
|
| 829 |
+
else:
|
| 830 |
+
raise ValueError(
|
| 831 |
+
"Wrong shape for input_ids (shape {}) or attention_mask (shape {})".format(
|
| 832 |
+
input_shape, attention_mask.shape
|
| 833 |
+
)
|
| 834 |
+
)
|
| 835 |
+
|
| 836 |
+
# Since attention_mask is 1.0 for positions we want to attend and 0.0 for
|
| 837 |
+
# masked positions, this operation will create a tensor which is 0.0 for
|
| 838 |
+
# positions we want to attend and -10000.0 for masked positions.
|
| 839 |
+
# Since we are adding it to the raw scores before the softmax, this is
|
| 840 |
+
# effectively the same as removing these entirely.
|
| 841 |
+
extended_attention_mask = extended_attention_mask.to(
|
| 842 |
+
dtype=self.dtype
|
| 843 |
+
) # fp16 compatibility
|
| 844 |
+
extended_attention_mask = (1.0 - extended_attention_mask) * -10000.0
|
| 845 |
+
return extended_attention_mask
|
| 846 |
+
|
| 847 |
+
def forward(
|
| 848 |
+
self,
|
| 849 |
+
input_ids=None,
|
| 850 |
+
attention_mask=None,
|
| 851 |
+
token_type_ids=None,
|
| 852 |
+
position_ids=None,
|
| 853 |
+
head_mask=None,
|
| 854 |
+
inputs_embeds=None,
|
| 855 |
+
encoder_embeds=None,
|
| 856 |
+
encoder_hidden_states=None,
|
| 857 |
+
encoder_attention_mask=None,
|
| 858 |
+
past_key_values=None,
|
| 859 |
+
use_cache=None,
|
| 860 |
+
output_attentions=None,
|
| 861 |
+
output_hidden_states=None,
|
| 862 |
+
return_dict=None,
|
| 863 |
+
is_decoder=False,
|
| 864 |
+
mode="multimodal",
|
| 865 |
+
):
|
| 866 |
+
r"""
|
| 867 |
+
encoder_hidden_states (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, sequence_length, hidden_size)`, `optional`):
|
| 868 |
+
Sequence of hidden-states at the output of the last layer of the encoder. Used in the cross-attention if
|
| 869 |
+
the model is configured as a decoder.
|
| 870 |
+
encoder_attention_mask (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, sequence_length)`, `optional`):
|
| 871 |
+
Mask to avoid performing attention on the padding token indices of the encoder input. This mask is used in
|
| 872 |
+
the cross-attention if the model is configured as a decoder. Mask values selected in ``[0, 1]``:
|
| 873 |
+
- 1 for tokens that are **not masked**,
|
| 874 |
+
- 0 for tokens that are **masked**.
|
| 875 |
+
past_key_values (:obj:`tuple(tuple(torch.FloatTensor))` of length :obj:`config.n_layers` with each tuple having 4 tensors of shape :obj:`(batch_size, num_heads, sequence_length - 1, embed_size_per_head)`):
|
| 876 |
+
Contains precomputed key and value hidden states of the attention blocks. Can be used to speed up decoding.
|
| 877 |
+
If :obj:`past_key_values` are used, the user can optionally input only the last :obj:`decoder_input_ids`
|
| 878 |
+
(those that don't have their past key value states given to this model) of shape :obj:`(batch_size, 1)`
|
| 879 |
+
instead of all :obj:`decoder_input_ids` of shape :obj:`(batch_size, sequence_length)`.
|
| 880 |
+
use_cache (:obj:`bool`, `optional`):
|
| 881 |
+
If set to :obj:`True`, :obj:`past_key_values` key value states are returned and can be used to speed up
|
| 882 |
+
decoding (see :obj:`past_key_values`).
|
| 883 |
+
"""
|
| 884 |
+
output_attentions = (
|
| 885 |
+
output_attentions
|
| 886 |
+
if output_attentions is not None
|
| 887 |
+
else self.config.output_attentions
|
| 888 |
+
)
|
| 889 |
+
output_hidden_states = (
|
| 890 |
+
output_hidden_states
|
| 891 |
+
if output_hidden_states is not None
|
| 892 |
+
else self.config.output_hidden_states
|
| 893 |
+
)
|
| 894 |
+
return_dict = (
|
| 895 |
+
return_dict if return_dict is not None else self.config.use_return_dict
|
| 896 |
+
)
|
| 897 |
+
|
| 898 |
+
if is_decoder:
|
| 899 |
+
use_cache = use_cache if use_cache is not None else self.config.use_cache
|
| 900 |
+
else:
|
| 901 |
+
use_cache = False
|
| 902 |
+
|
| 903 |
+
if input_ids is not None and inputs_embeds is not None:
|
| 904 |
+
raise ValueError(
|
| 905 |
+
"You cannot specify both input_ids and inputs_embeds at the same time"
|
| 906 |
+
)
|
| 907 |
+
elif input_ids is not None:
|
| 908 |
+
input_shape = input_ids.size()
|
| 909 |
+
batch_size, seq_length = input_shape
|
| 910 |
+
device = input_ids.device
|
| 911 |
+
elif inputs_embeds is not None:
|
| 912 |
+
input_shape = inputs_embeds.size()[:-1]
|
| 913 |
+
batch_size, seq_length = input_shape
|
| 914 |
+
device = inputs_embeds.device
|
| 915 |
+
elif encoder_embeds is not None:
|
| 916 |
+
input_shape = encoder_embeds.size()[:-1]
|
| 917 |
+
batch_size, seq_length = input_shape
|
| 918 |
+
device = encoder_embeds.device
|
| 919 |
+
else:
|
| 920 |
+
raise ValueError(
|
| 921 |
+
"You have to specify either input_ids or inputs_embeds or encoder_embeds"
|
| 922 |
+
)
|
| 923 |
+
|
| 924 |
+
# past_key_values_length
|
| 925 |
+
past_key_values_length = (
|
| 926 |
+
past_key_values[0][0].shape[2] if past_key_values is not None else 0
|
| 927 |
+
)
|
| 928 |
+
|
| 929 |
+
if attention_mask is None:
|
| 930 |
+
attention_mask = torch.ones(
|
| 931 |
+
((batch_size, seq_length + past_key_values_length)), device=device
|
| 932 |
+
)
|
| 933 |
+
|
| 934 |
+
# We can provide a self-attention mask of dimensions [batch_size, from_seq_length, to_seq_length]
|
| 935 |
+
# ourselves in which case we just need to make it broadcastable to all heads.
|
| 936 |
+
extended_attention_mask: torch.Tensor = self.get_extended_attention_mask(
|
| 937 |
+
attention_mask, input_shape, device, is_decoder
|
| 938 |
+
)
|
| 939 |
+
|
| 940 |
+
# If a 2D or 3D attention mask is provided for the cross-attention
|
| 941 |
+
# we need to make broadcastable to [batch_size, num_heads, seq_length, seq_length]
|
| 942 |
+
if encoder_hidden_states is not None:
|
| 943 |
+
if type(encoder_hidden_states) == list:
|
| 944 |
+
encoder_batch_size, encoder_sequence_length, _ = encoder_hidden_states[
|
| 945 |
+
0
|
| 946 |
+
].size()
|
| 947 |
+
else:
|
| 948 |
+
(
|
| 949 |
+
encoder_batch_size,
|
| 950 |
+
encoder_sequence_length,
|
| 951 |
+
_,
|
| 952 |
+
) = encoder_hidden_states.size()
|
| 953 |
+
encoder_hidden_shape = (encoder_batch_size, encoder_sequence_length)
|
| 954 |
+
|
| 955 |
+
if type(encoder_attention_mask) == list:
|
| 956 |
+
encoder_extended_attention_mask = [
|
| 957 |
+
self.invert_attention_mask(mask) for mask in encoder_attention_mask
|
| 958 |
+
]
|
| 959 |
+
elif encoder_attention_mask is None:
|
| 960 |
+
encoder_attention_mask = torch.ones(encoder_hidden_shape, device=device)
|
| 961 |
+
encoder_extended_attention_mask = self.invert_attention_mask(
|
| 962 |
+
encoder_attention_mask
|
| 963 |
+
)
|
| 964 |
+
else:
|
| 965 |
+
encoder_extended_attention_mask = self.invert_attention_mask(
|
| 966 |
+
encoder_attention_mask
|
| 967 |
+
)
|
| 968 |
+
else:
|
| 969 |
+
encoder_extended_attention_mask = None
|
| 970 |
+
|
| 971 |
+
# Prepare head mask if needed
|
| 972 |
+
# 1.0 in head_mask indicate we keep the head
|
| 973 |
+
# attention_probs has shape bsz x n_heads x N x N
|
| 974 |
+
# input head_mask has shape [num_heads] or [num_hidden_layers x num_heads]
|
| 975 |
+
# and head_mask is converted to shape [num_hidden_layers x batch x num_heads x seq_length x seq_length]
|
| 976 |
+
head_mask = self.get_head_mask(head_mask, self.config.num_hidden_layers)
|
| 977 |
+
|
| 978 |
+
if encoder_embeds is None:
|
| 979 |
+
embedding_output = self.embeddings(
|
| 980 |
+
input_ids=input_ids,
|
| 981 |
+
position_ids=position_ids,
|
| 982 |
+
token_type_ids=token_type_ids,
|
| 983 |
+
inputs_embeds=inputs_embeds,
|
| 984 |
+
past_key_values_length=past_key_values_length,
|
| 985 |
+
)
|
| 986 |
+
else:
|
| 987 |
+
embedding_output = encoder_embeds
|
| 988 |
+
|
| 989 |
+
encoder_outputs = self.encoder(
|
| 990 |
+
embedding_output,
|
| 991 |
+
attention_mask=extended_attention_mask,
|
| 992 |
+
head_mask=head_mask,
|
| 993 |
+
encoder_hidden_states=encoder_hidden_states,
|
| 994 |
+
encoder_attention_mask=encoder_extended_attention_mask,
|
| 995 |
+
past_key_values=past_key_values,
|
| 996 |
+
use_cache=use_cache,
|
| 997 |
+
output_attentions=output_attentions,
|
| 998 |
+
output_hidden_states=output_hidden_states,
|
| 999 |
+
return_dict=return_dict,
|
| 1000 |
+
mode=mode,
|
| 1001 |
+
)
|
| 1002 |
+
sequence_output = encoder_outputs[0]
|
| 1003 |
+
pooled_output = (
|
| 1004 |
+
self.pooler(sequence_output) if self.pooler is not None else None
|
| 1005 |
+
)
|
| 1006 |
+
|
| 1007 |
+
if not return_dict:
|
| 1008 |
+
return (sequence_output, pooled_output) + encoder_outputs[1:]
|
| 1009 |
+
|
| 1010 |
+
return BaseModelOutputWithPoolingAndCrossAttentions(
|
| 1011 |
+
last_hidden_state=sequence_output,
|
| 1012 |
+
pooler_output=pooled_output,
|
| 1013 |
+
past_key_values=encoder_outputs.past_key_values,
|
| 1014 |
+
hidden_states=encoder_outputs.hidden_states,
|
| 1015 |
+
attentions=encoder_outputs.attentions,
|
| 1016 |
+
cross_attentions=encoder_outputs.cross_attentions,
|
| 1017 |
+
)
|
| 1018 |
+
|
| 1019 |
+
|
| 1020 |
+
class BertForMaskedLM(BertPreTrainedModel):
|
| 1021 |
+
|
| 1022 |
+
_keys_to_ignore_on_load_unexpected = [r"pooler"]
|
| 1023 |
+
_keys_to_ignore_on_load_missing = [r"position_ids", r"predictions.decoder.bias"]
|
| 1024 |
+
|
| 1025 |
+
def __init__(self, config):
|
| 1026 |
+
super().__init__(config)
|
| 1027 |
+
|
| 1028 |
+
self.bert = BertModel(config, add_pooling_layer=False)
|
| 1029 |
+
self.cls = BertOnlyMLMHead(config)
|
| 1030 |
+
|
| 1031 |
+
self.init_weights()
|
| 1032 |
+
|
| 1033 |
+
def get_output_embeddings(self):
|
| 1034 |
+
return self.cls.predictions.decoder
|
| 1035 |
+
|
| 1036 |
+
def set_output_embeddings(self, new_embeddings):
|
| 1037 |
+
self.cls.predictions.decoder = new_embeddings
|
| 1038 |
+
|
| 1039 |
+
def forward(
|
| 1040 |
+
self,
|
| 1041 |
+
input_ids=None,
|
| 1042 |
+
attention_mask=None,
|
| 1043 |
+
# token_type_ids=None,
|
| 1044 |
+
position_ids=None,
|
| 1045 |
+
head_mask=None,
|
| 1046 |
+
inputs_embeds=None,
|
| 1047 |
+
encoder_embeds=None,
|
| 1048 |
+
encoder_hidden_states=None,
|
| 1049 |
+
encoder_attention_mask=None,
|
| 1050 |
+
labels=None,
|
| 1051 |
+
output_attentions=None,
|
| 1052 |
+
output_hidden_states=None,
|
| 1053 |
+
return_dict=None,
|
| 1054 |
+
is_decoder=False,
|
| 1055 |
+
mode="multimodal",
|
| 1056 |
+
soft_labels=None,
|
| 1057 |
+
alpha=0,
|
| 1058 |
+
return_logits=False,
|
| 1059 |
+
):
|
| 1060 |
+
r"""
|
| 1061 |
+
labels (:obj:`torch.LongTensor` of shape :obj:`(batch_size, sequence_length)`, `optional`):
|
| 1062 |
+
Labels for computing the masked language modeling loss. Indices should be in ``[-100, 0, ...,
|
| 1063 |
+
config.vocab_size]`` (see ``input_ids`` docstring) Tokens with indices set to ``-100`` are ignored
|
| 1064 |
+
(masked), the loss is only computed for the tokens with labels in ``[0, ..., config.vocab_size]``
|
| 1065 |
+
"""
|
| 1066 |
+
|
| 1067 |
+
return_dict = (
|
| 1068 |
+
return_dict if return_dict is not None else self.config.use_return_dict
|
| 1069 |
+
)
|
| 1070 |
+
|
| 1071 |
+
outputs = self.bert(
|
| 1072 |
+
input_ids,
|
| 1073 |
+
attention_mask=attention_mask,
|
| 1074 |
+
# token_type_ids=token_type_ids,
|
| 1075 |
+
position_ids=position_ids,
|
| 1076 |
+
head_mask=head_mask,
|
| 1077 |
+
inputs_embeds=inputs_embeds,
|
| 1078 |
+
encoder_embeds=encoder_embeds,
|
| 1079 |
+
encoder_hidden_states=encoder_hidden_states,
|
| 1080 |
+
encoder_attention_mask=encoder_attention_mask,
|
| 1081 |
+
output_attentions=output_attentions,
|
| 1082 |
+
output_hidden_states=output_hidden_states,
|
| 1083 |
+
return_dict=return_dict,
|
| 1084 |
+
is_decoder=is_decoder,
|
| 1085 |
+
mode=mode,
|
| 1086 |
+
)
|
| 1087 |
+
|
| 1088 |
+
sequence_output = outputs[0]
|
| 1089 |
+
prediction_scores = self.cls(sequence_output)
|
| 1090 |
+
|
| 1091 |
+
if return_logits:
|
| 1092 |
+
return prediction_scores
|
| 1093 |
+
|
| 1094 |
+
masked_lm_loss = None
|
| 1095 |
+
if labels is not None:
|
| 1096 |
+
loss_fct = CrossEntropyLoss() # -100 index = padding token
|
| 1097 |
+
masked_lm_loss = loss_fct(
|
| 1098 |
+
prediction_scores.view(-1, self.config.vocab_size), labels.view(-1)
|
| 1099 |
+
)
|
| 1100 |
+
|
| 1101 |
+
if soft_labels is not None:
|
| 1102 |
+
loss_distill = -torch.sum(
|
| 1103 |
+
F.log_softmax(prediction_scores, dim=-1) * soft_labels, dim=-1
|
| 1104 |
+
)
|
| 1105 |
+
loss_distill = loss_distill[labels != -100].mean()
|
| 1106 |
+
masked_lm_loss = (1 - alpha) * masked_lm_loss + alpha * loss_distill
|
| 1107 |
+
|
| 1108 |
+
if not return_dict:
|
| 1109 |
+
output = (prediction_scores,) + outputs[2:]
|
| 1110 |
+
return (
|
| 1111 |
+
((masked_lm_loss,) + output) if masked_lm_loss is not None else output
|
| 1112 |
+
)
|
| 1113 |
+
|
| 1114 |
+
return MaskedLMOutput(
|
| 1115 |
+
loss=masked_lm_loss,
|
| 1116 |
+
logits=prediction_scores,
|
| 1117 |
+
hidden_states=outputs.hidden_states,
|
| 1118 |
+
attentions=outputs.attentions,
|
| 1119 |
+
)
|
| 1120 |
+
|
| 1121 |
+
def prepare_inputs_for_generation(
|
| 1122 |
+
self, input_ids, attention_mask=None, **model_kwargs
|
| 1123 |
+
):
|
| 1124 |
+
input_shape = input_ids.shape
|
| 1125 |
+
effective_batch_size = input_shape[0]
|
| 1126 |
+
|
| 1127 |
+
# add a dummy token
|
| 1128 |
+
assert (
|
| 1129 |
+
self.config.pad_token_id is not None
|
| 1130 |
+
), "The PAD token should be defined for generation"
|
| 1131 |
+
attention_mask = torch.cat(
|
| 1132 |
+
[attention_mask, attention_mask.new_zeros((attention_mask.shape[0], 1))],
|
| 1133 |
+
dim=-1,
|
| 1134 |
+
)
|
| 1135 |
+
dummy_token = torch.full(
|
| 1136 |
+
(effective_batch_size, 1),
|
| 1137 |
+
self.config.pad_token_id,
|
| 1138 |
+
dtype=torch.long,
|
| 1139 |
+
device=input_ids.device,
|
| 1140 |
+
)
|
| 1141 |
+
input_ids = torch.cat([input_ids, dummy_token], dim=1)
|
| 1142 |
+
|
| 1143 |
+
return {"input_ids": input_ids, "attention_mask": attention_mask}
|
| 1144 |
+
|
| 1145 |
+
|
| 1146 |
+
class BertLMHeadModel(BertPreTrainedModel):
|
| 1147 |
+
|
| 1148 |
+
_keys_to_ignore_on_load_unexpected = [r"pooler"]
|
| 1149 |
+
_keys_to_ignore_on_load_missing = [r"position_ids", r"predictions.decoder.bias"]
|
| 1150 |
+
|
| 1151 |
+
def __init__(self, config):
|
| 1152 |
+
super().__init__(config)
|
| 1153 |
+
|
| 1154 |
+
self.bert = BertModel(config, add_pooling_layer=False)
|
| 1155 |
+
self.cls = BertOnlyMLMHead(config)
|
| 1156 |
+
|
| 1157 |
+
self.init_weights()
|
| 1158 |
+
|
| 1159 |
+
def get_output_embeddings(self):
|
| 1160 |
+
return self.cls.predictions.decoder
|
| 1161 |
+
|
| 1162 |
+
def set_output_embeddings(self, new_embeddings):
|
| 1163 |
+
self.cls.predictions.decoder = new_embeddings
|
| 1164 |
+
|
| 1165 |
+
def forward(
|
| 1166 |
+
self,
|
| 1167 |
+
input_ids=None,
|
| 1168 |
+
attention_mask=None,
|
| 1169 |
+
position_ids=None,
|
| 1170 |
+
head_mask=None,
|
| 1171 |
+
inputs_embeds=None,
|
| 1172 |
+
encoder_hidden_states=None,
|
| 1173 |
+
encoder_attention_mask=None,
|
| 1174 |
+
labels=None,
|
| 1175 |
+
past_key_values=None,
|
| 1176 |
+
use_cache=None,
|
| 1177 |
+
output_attentions=None,
|
| 1178 |
+
output_hidden_states=None,
|
| 1179 |
+
return_dict=None,
|
| 1180 |
+
return_logits=False,
|
| 1181 |
+
is_decoder=True,
|
| 1182 |
+
reduction="mean",
|
| 1183 |
+
mode="multimodal",
|
| 1184 |
+
soft_labels=None,
|
| 1185 |
+
alpha=0,
|
| 1186 |
+
):
|
| 1187 |
+
r"""
|
| 1188 |
+
encoder_hidden_states (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, sequence_length, hidden_size)`, `optional`):
|
| 1189 |
+
Sequence of hidden-states at the output of the last layer of the encoder. Used in the cross-attention if
|
| 1190 |
+
the model is configured as a decoder.
|
| 1191 |
+
encoder_attention_mask (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, sequence_length)`, `optional`):
|
| 1192 |
+
Mask to avoid performing attention on the padding token indices of the encoder input. This mask is used in
|
| 1193 |
+
the cross-attention if the model is configured as a decoder. Mask values selected in ``[0, 1]``:
|
| 1194 |
+
- 1 for tokens that are **not masked**,
|
| 1195 |
+
- 0 for tokens that are **masked**.
|
| 1196 |
+
labels (:obj:`torch.LongTensor` of shape :obj:`(batch_size, sequence_length)`, `optional`):
|
| 1197 |
+
Labels for computing the left-to-right language modeling loss (next word prediction). Indices should be in
|
| 1198 |
+
``[-100, 0, ..., config.vocab_size]`` (see ``input_ids`` docstring) Tokens with indices set to ``-100`` are
|
| 1199 |
+
ignored (masked), the loss is only computed for the tokens with labels n ``[0, ..., config.vocab_size]``
|
| 1200 |
+
past_key_values (:obj:`tuple(tuple(torch.FloatTensor))` of length :obj:`config.n_layers` with each tuple having 4 tensors of shape :obj:`(batch_size, num_heads, sequence_length - 1, embed_size_per_head)`):
|
| 1201 |
+
Contains precomputed key and value hidden states of the attention blocks. Can be used to speed up decoding.
|
| 1202 |
+
If :obj:`past_key_values` are used, the user can optionally input only the last :obj:`decoder_input_ids`
|
| 1203 |
+
(those that don't have their past key value states given to this model) of shape :obj:`(batch_size, 1)`
|
| 1204 |
+
instead of all :obj:`decoder_input_ids` of shape :obj:`(batch_size, sequence_length)`.
|
| 1205 |
+
use_cache (:obj:`bool`, `optional`):
|
| 1206 |
+
If set to :obj:`True`, :obj:`past_key_values` key value states are returned and can be used to speed up
|
| 1207 |
+
decoding (see :obj:`past_key_values`).
|
| 1208 |
+
Returns:
|
| 1209 |
+
Example::
|
| 1210 |
+
>>> from transformers import BertTokenizer, BertLMHeadModel, BertConfig
|
| 1211 |
+
>>> import torch
|
| 1212 |
+
>>> tokenizer = BertTokenizer.from_pretrained('bert-base-cased')
|
| 1213 |
+
>>> config = BertConfig.from_pretrained("bert-base-cased")
|
| 1214 |
+
>>> model = BertLMHeadModel.from_pretrained('bert-base-cased', config=config)
|
| 1215 |
+
>>> inputs = tokenizer("Hello, my dog is cute", return_tensors="pt")
|
| 1216 |
+
>>> outputs = model(**inputs)
|
| 1217 |
+
>>> prediction_logits = outputs.logits
|
| 1218 |
+
"""
|
| 1219 |
+
return_dict = (
|
| 1220 |
+
return_dict if return_dict is not None else self.config.use_return_dict
|
| 1221 |
+
)
|
| 1222 |
+
if labels is not None:
|
| 1223 |
+
use_cache = False
|
| 1224 |
+
|
| 1225 |
+
outputs = self.bert(
|
| 1226 |
+
input_ids,
|
| 1227 |
+
attention_mask=attention_mask,
|
| 1228 |
+
position_ids=position_ids,
|
| 1229 |
+
head_mask=head_mask,
|
| 1230 |
+
inputs_embeds=inputs_embeds,
|
| 1231 |
+
encoder_hidden_states=encoder_hidden_states,
|
| 1232 |
+
encoder_attention_mask=encoder_attention_mask,
|
| 1233 |
+
past_key_values=past_key_values,
|
| 1234 |
+
use_cache=use_cache,
|
| 1235 |
+
output_attentions=output_attentions,
|
| 1236 |
+
output_hidden_states=output_hidden_states,
|
| 1237 |
+
return_dict=return_dict,
|
| 1238 |
+
is_decoder=is_decoder,
|
| 1239 |
+
mode=mode,
|
| 1240 |
+
)
|
| 1241 |
+
|
| 1242 |
+
sequence_output = outputs[0]
|
| 1243 |
+
prediction_scores = self.cls(sequence_output)
|
| 1244 |
+
|
| 1245 |
+
if return_logits:
|
| 1246 |
+
return prediction_scores[:, :-1, :].contiguous()
|
| 1247 |
+
|
| 1248 |
+
lm_loss = None
|
| 1249 |
+
if labels is not None:
|
| 1250 |
+
# we are doing next-token prediction; shift prediction scores and input ids by one
|
| 1251 |
+
shifted_prediction_scores = prediction_scores[:, :-1, :].contiguous()
|
| 1252 |
+
labels = labels[:, 1:].contiguous()
|
| 1253 |
+
loss_fct = CrossEntropyLoss(reduction=reduction, label_smoothing=0.1)
|
| 1254 |
+
lm_loss = loss_fct(
|
| 1255 |
+
shifted_prediction_scores.view(-1, self.config.vocab_size),
|
| 1256 |
+
labels.view(-1),
|
| 1257 |
+
)
|
| 1258 |
+
if reduction == "none":
|
| 1259 |
+
lm_loss = lm_loss.view(prediction_scores.size(0), -1).sum(1)
|
| 1260 |
+
|
| 1261 |
+
if soft_labels is not None:
|
| 1262 |
+
loss_distill = -torch.sum(
|
| 1263 |
+
F.log_softmax(shifted_prediction_scores, dim=-1) * soft_labels, dim=-1
|
| 1264 |
+
)
|
| 1265 |
+
loss_distill = (loss_distill * (labels != -100)).sum(1)
|
| 1266 |
+
lm_loss = (1 - alpha) * lm_loss + alpha * loss_distill
|
| 1267 |
+
|
| 1268 |
+
if not return_dict:
|
| 1269 |
+
output = (prediction_scores,) + outputs[2:]
|
| 1270 |
+
return ((lm_loss,) + output) if lm_loss is not None else output
|
| 1271 |
+
|
| 1272 |
+
# for embedding of the desc in the decoder, return the last_hidden_state (embeds)
|
| 1273 |
+
if (type(self).__name__ == 'XBertLMHeadDecoder') and (outputs.last_hidden_state.shape[1] == 200):
|
| 1274 |
+
return [CausalLMOutputWithCrossAttentions(
|
| 1275 |
+
loss=lm_loss,
|
| 1276 |
+
logits=prediction_scores,
|
| 1277 |
+
past_key_values=outputs.past_key_values,
|
| 1278 |
+
hidden_states=outputs.hidden_states,
|
| 1279 |
+
attentions=outputs.attentions,
|
| 1280 |
+
cross_attentions=outputs.cross_attentions,
|
| 1281 |
+
), outputs.last_hidden_state]
|
| 1282 |
+
else:
|
| 1283 |
+
return CausalLMOutputWithCrossAttentions(
|
| 1284 |
+
loss=lm_loss,
|
| 1285 |
+
logits=prediction_scores,
|
| 1286 |
+
past_key_values=outputs.past_key_values,
|
| 1287 |
+
hidden_states=outputs.hidden_states,
|
| 1288 |
+
attentions=outputs.attentions,
|
| 1289 |
+
cross_attentions=outputs.cross_attentions,
|
| 1290 |
+
)
|
| 1291 |
+
|
| 1292 |
+
def prepare_inputs_for_generation(
|
| 1293 |
+
self, input_ids, past=None, attention_mask=None, **model_kwargs
|
| 1294 |
+
):
|
| 1295 |
+
input_shape = input_ids.shape
|
| 1296 |
+
# if model is used as a decoder in encoder-decoder model, the decoder attention mask is created on the fly
|
| 1297 |
+
if attention_mask is None:
|
| 1298 |
+
attention_mask = input_ids.new_ones(input_shape)
|
| 1299 |
+
|
| 1300 |
+
# cut decoder_input_ids if past is used
|
| 1301 |
+
if past is not None:
|
| 1302 |
+
input_ids = input_ids[:, -1:]
|
| 1303 |
+
|
| 1304 |
+
return {
|
| 1305 |
+
"input_ids": input_ids,
|
| 1306 |
+
"attention_mask": attention_mask,
|
| 1307 |
+
"past_key_values": past,
|
| 1308 |
+
"encoder_hidden_states": model_kwargs.get("encoder_hidden_states", None),
|
| 1309 |
+
"encoder_attention_mask": model_kwargs.get("encoder_attention_mask", None),
|
| 1310 |
+
"is_decoder": True,
|
| 1311 |
+
}
|
| 1312 |
+
|
| 1313 |
+
def _reorder_cache(self, past, beam_idx):
|
| 1314 |
+
reordered_past = ()
|
| 1315 |
+
for layer_past in past:
|
| 1316 |
+
reordered_past += (
|
| 1317 |
+
tuple(
|
| 1318 |
+
past_state.index_select(0, beam_idx) for past_state in layer_past
|
| 1319 |
+
),
|
| 1320 |
+
)
|
| 1321 |
+
return reordered_past
|
| 1322 |
+
|
| 1323 |
+
|
| 1324 |
+
class XBertLMHeadDecoder(BertLMHeadModel):
|
| 1325 |
+
"""
|
| 1326 |
+
This class decouples the decoder forward logic from the VL model.
|
| 1327 |
+
In this way, different VL models can share this decoder as long as
|
| 1328 |
+
they feed encoder_embeds as required.
|
| 1329 |
+
"""
|
| 1330 |
+
|
| 1331 |
+
@classmethod
|
| 1332 |
+
def from_config(cls, cfg, from_pretrained=False):
|
| 1333 |
+
med_config_path = get_abs_path(cfg.get("med_config_path"))
|
| 1334 |
+
med_config = BertConfig.from_json_file(med_config_path)
|
| 1335 |
+
|
| 1336 |
+
if from_pretrained:
|
| 1337 |
+
return cls.from_pretrained("bert-base-uncased", config=med_config)
|
| 1338 |
+
else:
|
| 1339 |
+
return cls(config=med_config)
|
| 1340 |
+
|
| 1341 |
+
def generate_from_encoder(
|
| 1342 |
+
self,
|
| 1343 |
+
tokenized_prompt,
|
| 1344 |
+
visual_embeds,
|
| 1345 |
+
sep_token_id,
|
| 1346 |
+
pad_token_id,
|
| 1347 |
+
use_nucleus_sampling=False,
|
| 1348 |
+
num_beams=3,
|
| 1349 |
+
max_length=30,
|
| 1350 |
+
min_length=10,
|
| 1351 |
+
top_p=0.9,
|
| 1352 |
+
repetition_penalty=1.0,
|
| 1353 |
+
**kwargs
|
| 1354 |
+
):
|
| 1355 |
+
|
| 1356 |
+
if not use_nucleus_sampling:
|
| 1357 |
+
num_beams = num_beams
|
| 1358 |
+
visual_embeds = visual_embeds.repeat_interleave(num_beams, dim=0)
|
| 1359 |
+
|
| 1360 |
+
image_atts = torch.ones(visual_embeds.size()[:-1], dtype=torch.long).to(
|
| 1361 |
+
self.device
|
| 1362 |
+
)
|
| 1363 |
+
|
| 1364 |
+
model_kwargs = {
|
| 1365 |
+
"encoder_hidden_states": visual_embeds,
|
| 1366 |
+
"encoder_attention_mask": image_atts,
|
| 1367 |
+
}
|
| 1368 |
+
|
| 1369 |
+
if use_nucleus_sampling:
|
| 1370 |
+
# nucleus sampling
|
| 1371 |
+
outputs = self.generate(
|
| 1372 |
+
input_ids=tokenized_prompt.input_ids,
|
| 1373 |
+
max_length=max_length,
|
| 1374 |
+
min_length=min_length,
|
| 1375 |
+
do_sample=True,
|
| 1376 |
+
top_p=top_p,
|
| 1377 |
+
num_return_sequences=1,
|
| 1378 |
+
eos_token_id=sep_token_id,
|
| 1379 |
+
pad_token_id=pad_token_id,
|
| 1380 |
+
repetition_penalty=1.1,
|
| 1381 |
+
**model_kwargs
|
| 1382 |
+
)
|
| 1383 |
+
else:
|
| 1384 |
+
# beam search
|
| 1385 |
+
outputs = self.generate(
|
| 1386 |
+
input_ids=tokenized_prompt.input_ids,
|
| 1387 |
+
max_length=max_length,
|
| 1388 |
+
min_length=min_length,
|
| 1389 |
+
num_beams=num_beams,
|
| 1390 |
+
eos_token_id=sep_token_id,
|
| 1391 |
+
pad_token_id=pad_token_id,
|
| 1392 |
+
repetition_penalty=repetition_penalty,
|
| 1393 |
+
**model_kwargs
|
| 1394 |
+
)
|
| 1395 |
+
|
| 1396 |
+
return outputs
|
| 1397 |
+
|
| 1398 |
+
|
| 1399 |
+
class XBertEncoder(BertModel, BaseEncoder):
|
| 1400 |
+
def __init__(self, config, add_pooling_layer):
|
| 1401 |
+
super().__init__(config, add_pooling_layer)
|
| 1402 |
+
|
| 1403 |
+
self.cls = BertOnlyMLMHead(config)
|
| 1404 |
+
self.init_weights()
|
| 1405 |
+
|
| 1406 |
+
def get_output_embeddings(self):
|
| 1407 |
+
return self.cls.predictions.decoder
|
| 1408 |
+
|
| 1409 |
+
@classmethod
|
| 1410 |
+
def from_config(cls, cfg, from_pretrained=False):
|
| 1411 |
+
configs_root = os.environ.get("CONFIGS_ROOT", "../ckpt")
|
| 1412 |
+
med_config_path = os.path.join(configs_root, "bert-base-chinese/config.json")
|
| 1413 |
+
med_config = BertConfig.from_json_file(med_config_path)
|
| 1414 |
+
|
| 1415 |
+
if from_pretrained:
|
| 1416 |
+
model = cls.from_pretrained(
|
| 1417 |
+
os.path.join(configs_root, "bert-base-chinese"),
|
| 1418 |
+
config=med_config,
|
| 1419 |
+
add_pooling_layer=False
|
| 1420 |
+
)
|
| 1421 |
+
else:
|
| 1422 |
+
model = cls(config=med_config, add_pooling_layer=False)
|
| 1423 |
+
|
| 1424 |
+
for name, param in model.named_parameters():
|
| 1425 |
+
if 'cls.predictions' in name:
|
| 1426 |
+
param.requires_grad = False
|
| 1427 |
+
|
| 1428 |
+
return model
|
| 1429 |
+
|
| 1430 |
+
def forward_automask(self, tokenized_text, visual_embeds, **kwargs):
|
| 1431 |
+
image_atts = torch.ones(visual_embeds.size()[:-1], dtype=torch.long).to(
|
| 1432 |
+
self.device
|
| 1433 |
+
)
|
| 1434 |
+
|
| 1435 |
+
text = tokenized_text
|
| 1436 |
+
text_output = super().forward(
|
| 1437 |
+
text.input_ids,
|
| 1438 |
+
attention_mask=text.attention_mask,
|
| 1439 |
+
encoder_hidden_states=visual_embeds,
|
| 1440 |
+
encoder_attention_mask=image_atts,
|
| 1441 |
+
return_dict=True,
|
| 1442 |
+
)
|
| 1443 |
+
|
| 1444 |
+
return text_output
|
| 1445 |
+
|
| 1446 |
+
def forward_text(self, tokenized_text, **kwargs):
|
| 1447 |
+
text = tokenized_text
|
| 1448 |
+
token_type_ids = kwargs.get("token_type_ids", None)
|
| 1449 |
+
|
| 1450 |
+
text_output = super().forward(
|
| 1451 |
+
text.input_ids,
|
| 1452 |
+
attention_mask=text.attention_mask,
|
| 1453 |
+
token_type_ids=token_type_ids,
|
| 1454 |
+
return_dict=True,
|
| 1455 |
+
mode="text",
|
| 1456 |
+
)
|
| 1457 |
+
|
| 1458 |
+
return text_output
|
| 1459 |
+
|
| 1460 |
+
def forward_mlm(self,
|
| 1461 |
+
input_ids,
|
| 1462 |
+
attention_mask,
|
| 1463 |
+
encoder_hidden_states,
|
| 1464 |
+
encoder_attention_mask,
|
| 1465 |
+
return_dict,
|
| 1466 |
+
labels
|
| 1467 |
+
):
|
| 1468 |
+
|
| 1469 |
+
# slef.bert()
|
| 1470 |
+
outputs = super().forward(
|
| 1471 |
+
input_ids=input_ids,
|
| 1472 |
+
attention_mask=attention_mask,
|
| 1473 |
+
encoder_hidden_states=encoder_hidden_states,
|
| 1474 |
+
encoder_attention_mask=encoder_attention_mask,
|
| 1475 |
+
return_dict=return_dict
|
| 1476 |
+
)
|
| 1477 |
+
|
| 1478 |
+
sequence_output = outputs[0]
|
| 1479 |
+
prediction_scores = self.cls(sequence_output)
|
| 1480 |
+
|
| 1481 |
+
# if return_logits:
|
| 1482 |
+
# return prediction_scores
|
| 1483 |
+
|
| 1484 |
+
masked_lm_loss = None
|
| 1485 |
+
if labels is not None:
|
| 1486 |
+
loss_fct = CrossEntropyLoss() # -100 index = padding token
|
| 1487 |
+
masked_lm_loss = loss_fct(
|
| 1488 |
+
prediction_scores.view(-1, self.config.vocab_size), labels.view(-1)
|
| 1489 |
+
)
|
| 1490 |
+
|
| 1491 |
+
if not return_dict:
|
| 1492 |
+
output = (prediction_scores,) + outputs[2:]
|
| 1493 |
+
return (
|
| 1494 |
+
((masked_lm_loss,) + output) if masked_lm_loss is not None else output
|
| 1495 |
+
)
|
| 1496 |
+
|
| 1497 |
+
return MaskedLMOutput(
|
| 1498 |
+
loss=masked_lm_loss,
|
| 1499 |
+
logits=prediction_scores,
|
| 1500 |
+
hidden_states=outputs.hidden_states,
|
| 1501 |
+
attentions=outputs.attentions,
|
| 1502 |
+
)
|
RADAR_inference/dynamic_network_architectures/vision_branch.py
ADDED
|
@@ -0,0 +1,160 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright (c) MONAI Consortium
|
| 2 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 3 |
+
# you may not use this file except in compliance with the License.
|
| 4 |
+
# You may obtain a copy of the License at
|
| 5 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 6 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 7 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 8 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 9 |
+
# See the License for the specific language governing permissions and
|
| 10 |
+
# limitations under the License.
|
| 11 |
+
|
| 12 |
+
from __future__ import annotations
|
| 13 |
+
|
| 14 |
+
from collections.abc import Sequence
|
| 15 |
+
|
| 16 |
+
import torch
|
| 17 |
+
import torch.nn as nn
|
| 18 |
+
import sys
|
| 19 |
+
from monai.utils import deprecated_arg
|
| 20 |
+
import pydoc
|
| 21 |
+
import warnings
|
| 22 |
+
from typing import Union
|
| 23 |
+
import torch.nn.functional as F
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
class VisionBranch(nn.Module):
|
| 27 |
+
def __init__(
|
| 28 |
+
self,
|
| 29 |
+
in_channels=1
|
| 30 |
+
) -> None:
|
| 31 |
+
super().__init__()
|
| 32 |
+
|
| 33 |
+
# define LightDecoderUNet
|
| 34 |
+
self.UNet = self.get_network_from_plans(
|
| 35 |
+
arch_class_name="dynamic_network_architectures.architectures.unet_lightdecoder.PlainConvUNetLightD",
|
| 36 |
+
arch_kwargs={
|
| 37 |
+
"n_stages": 6,
|
| 38 |
+
"features_per_stage": [32, 64, 128, 256, 320, 320],
|
| 39 |
+
"conv_op": "torch.nn.modules.conv.Conv3d",
|
| 40 |
+
"kernel_sizes": [[1, 3, 3], [1, 3, 3], [3, 3, 3], [3, 3, 3], [3, 3, 3], [3, 3, 3]],
|
| 41 |
+
"strides": [[1, 1, 1], [1, 2, 2], [1, 2, 2], [2, 2, 2], [2, 2, 2], [2, 2, 2]],
|
| 42 |
+
"n_conv_per_stage": [2, 2, 2, 2, 2, 2],
|
| 43 |
+
"n_conv_per_stage_decoder": [1, 1, 1, 1, 1],
|
| 44 |
+
"conv_bias": True,
|
| 45 |
+
"norm_op": "torch.nn.BatchNorm3d",
|
| 46 |
+
"norm_op_kwargs": {},
|
| 47 |
+
"dropout_op": None,
|
| 48 |
+
"dropout_op_kwargs": None,
|
| 49 |
+
"nonlin": "torch.nn.ReLU",
|
| 50 |
+
"nonlin_kwargs": {"inplace": True},
|
| 51 |
+
},
|
| 52 |
+
arch_kwargs_req_import=["conv_op", "norm_op", "dropout_op", "nonlin"],
|
| 53 |
+
input_channels=1,
|
| 54 |
+
output_channels=37,
|
| 55 |
+
allow_init=True,
|
| 56 |
+
deep_supervision=True,
|
| 57 |
+
)
|
| 58 |
+
|
| 59 |
+
self.proj1 = nn.Conv3d(320, 256, kernel_size=1)
|
| 60 |
+
self.proj2 = nn.Conv3d(320, 256, kernel_size=1)
|
| 61 |
+
self.proj3 = nn.Conv3d(256, 256, kernel_size=1)
|
| 62 |
+
|
| 63 |
+
self.organs = [
|
| 64 |
+
'肾上腺', '主动脉', '竖脊肌', '脑', '锁骨', '大肠', '十二指肠', '食管', '面部', '股骨',
|
| 65 |
+
'胆囊', "臀肌", '心脏', '髋关节', '肱骨', '髂动脉', '髂静脉', '髂腰肌', '下腔静脉', '肾',
|
| 66 |
+
'肝', '肺', '胰腺', '门静脉', '肺动脉', '肋骨', '骶骨', '肩胛骨', '小肠', '脾',
|
| 67 |
+
'胃', '气管', '膀胱', '颈椎', '腰椎', '胸椎'
|
| 68 |
+
]
|
| 69 |
+
# organ_dict = {
|
| 70 |
+
# "肾上腺": "adrenal gland", "主动脉": "aorta", "竖脊肌": "erector spinae muscle", "脑": "brain", "锁骨": "clavicle", "大肠": "large bowel", "十二指肠": "duodenum",
|
| 71 |
+
# "食管": "esophagus", "面部": "face", "股骨": "femur", "胆囊": "gallbladder", "臀肌": "gluteus muscle", "心脏": "heart", "髋关节": "hip joint", "肱骨": "humerus",
|
| 72 |
+
# "髂动脉": "iliac artery", "髂静脉": "iliac vena", "髂腰肌": "iliopsoas muscle", "下腔静脉": "inferior vena cava", "肾": "kidney", "肝": "liver", "肺": "lung",
|
| 73 |
+
# "胰腺": "pancreas", "门静脉": "portal vein", "肺动脉": "pulmonary artery", "肋骨": "rib", "骶骨": "sacrum", "肩胛骨": "scapula", "小肠": "small bowel", "脾": "spleen",
|
| 74 |
+
# "胃": "stomach", "气管": "trachea", "膀胱": "bladder", "颈椎": "cervical vertebrae", "腰椎": "lumbar vertebrae", "胸椎": "thoracic vertebrae"
|
| 75 |
+
# }
|
| 76 |
+
|
| 77 |
+
def get_network_from_plans(sefl, arch_class_name, arch_kwargs, arch_kwargs_req_import, input_channels, output_channels,
|
| 78 |
+
allow_init=True, deep_supervision: Union[bool, None] = None):
|
| 79 |
+
network_class = arch_class_name
|
| 80 |
+
architecture_kwargs = dict(**arch_kwargs)
|
| 81 |
+
for ri in arch_kwargs_req_import:
|
| 82 |
+
if architecture_kwargs[ri] is not None:
|
| 83 |
+
architecture_kwargs[ri] = pydoc.locate(architecture_kwargs[ri])
|
| 84 |
+
|
| 85 |
+
nw_class = pydoc.locate(network_class)
|
| 86 |
+
|
| 87 |
+
if deep_supervision is not None:
|
| 88 |
+
architecture_kwargs['deep_supervision'] = deep_supervision
|
| 89 |
+
|
| 90 |
+
network = nw_class(
|
| 91 |
+
input_channels=input_channels,
|
| 92 |
+
num_classes=output_channels,
|
| 93 |
+
**architecture_kwargs
|
| 94 |
+
)
|
| 95 |
+
|
| 96 |
+
if hasattr(network, 'initialize') and allow_init:
|
| 97 |
+
network.apply(network.initialize)
|
| 98 |
+
|
| 99 |
+
return network
|
| 100 |
+
|
| 101 |
+
def forward(self, x, y):
|
| 102 |
+
skips, segs = self.UNet(x)
|
| 103 |
+
|
| 104 |
+
scale1 = self.proj1(skips[-1])
|
| 105 |
+
scale2 = self.proj2(skips[-2])
|
| 106 |
+
scale3 = self.proj3(skips[-3])
|
| 107 |
+
|
| 108 |
+
# process mask
|
| 109 |
+
pred_logit = segs[0]
|
| 110 |
+
seg_probs = torch.softmax(pred_logit, 1)
|
| 111 |
+
target_size = [pred_logit.shape[-3], pred_logit.shape[-2]*2, pred_logit.shape[-1]*2]
|
| 112 |
+
pred_logit = F.interpolate(pred_logit, size=target_size, mode='trilinear', align_corners=False)
|
| 113 |
+
pred_mask = torch.softmax(pred_logit, 1)
|
| 114 |
+
pred_mask = pred_mask.argmax(1)
|
| 115 |
+
y = pred_mask
|
| 116 |
+
|
| 117 |
+
res_x1 = scale1.flatten(2).transpose(1, 2)
|
| 118 |
+
res_x2 = scale2.flatten(2).transpose(1, 2)
|
| 119 |
+
res_x3 = scale3.flatten(2).transpose(1, 2)
|
| 120 |
+
|
| 121 |
+
B, L1, _ = res_x1.size()
|
| 122 |
+
B, L2, _ = res_x2.size()
|
| 123 |
+
B, L3, _ = res_x3.size()
|
| 124 |
+
|
| 125 |
+
with torch.no_grad():
|
| 126 |
+
organ_token_flags1 = torch.zeros(B, len(self.organs), L1, dtype=bool).to(x.device)
|
| 127 |
+
organ_token_flags2 = torch.zeros(B, len(self.organs), L2, dtype=bool).to(x.device)
|
| 128 |
+
organ_token_flags3 = torch.zeros(B, len(self.organs), L3, dtype=bool).to(x.device)
|
| 129 |
+
|
| 130 |
+
b = x.size(0)
|
| 131 |
+
for i in range(b):
|
| 132 |
+
unique_values = torch.unique(y[i])
|
| 133 |
+
unique_values = unique_values[unique_values != 0]
|
| 134 |
+
if unique_values.tolist() == []:
|
| 135 |
+
continue
|
| 136 |
+
masks = torch.stack([torch.eq(y[i], uv) for uv in unique_values]).float()
|
| 137 |
+
|
| 138 |
+
highlight_tokens3 = F.max_pool3d(
|
| 139 |
+
masks.unsqueeze(1),
|
| 140 |
+
kernel_size=(2, 8, 8),
|
| 141 |
+
stride=(2, 8, 8)
|
| 142 |
+
).flatten(1) > 0
|
| 143 |
+
|
| 144 |
+
highlight_tokens2 = F.max_pool3d(
|
| 145 |
+
masks.unsqueeze(1),
|
| 146 |
+
kernel_size=(4, 16, 16),
|
| 147 |
+
stride=(4, 16, 16)
|
| 148 |
+
).flatten(1) > 0
|
| 149 |
+
|
| 150 |
+
highlight_tokens1 = F.max_pool3d(
|
| 151 |
+
masks.unsqueeze(1),
|
| 152 |
+
kernel_size=(8, 32, 32),
|
| 153 |
+
stride=(8, 32, 32)
|
| 154 |
+
).flatten(1) > 0
|
| 155 |
+
|
| 156 |
+
organ_token_flags1[i][unique_values.long() - 1] = highlight_tokens1 > 0
|
| 157 |
+
organ_token_flags2[i][unique_values.long() - 1] = highlight_tokens2 > 0
|
| 158 |
+
organ_token_flags3[i][unique_values.long() - 1] = highlight_tokens3 > 0
|
| 159 |
+
|
| 160 |
+
return seg_probs, pred_mask, res_x1, res_x2, res_x3, organ_token_flags1, organ_token_flags2, organ_token_flags3
|
RADAR_inference/inference_demo.py
ADDED
|
@@ -0,0 +1,630 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import re
|
| 3 |
+
import json
|
| 4 |
+
import argparse
|
| 5 |
+
import numpy as np
|
| 6 |
+
import pandas as pd
|
| 7 |
+
import torch
|
| 8 |
+
import torch.nn as nn
|
| 9 |
+
from tqdm import tqdm
|
| 10 |
+
from torch.utils.data import Dataset, DataLoader, DistributedSampler
|
| 11 |
+
from torch.utils.data.dataloader import default_collate
|
| 12 |
+
from monai import transforms
|
| 13 |
+
from monai.data.utils import dense_patch_slices
|
| 14 |
+
from typing import Any, Callable, List, Sequence, Tuple, Union
|
| 15 |
+
import datetime
|
| 16 |
+
import SimpleITK as sitk
|
| 17 |
+
import torch.nn.functional as F
|
| 18 |
+
from pathlib import Path
|
| 19 |
+
from copy import deepcopy
|
| 20 |
+
import torch.distributed as dist
|
| 21 |
+
from dynamic_network_architectures.med import XBertEncoder, XBertLMHeadDecoder
|
| 22 |
+
from dynamic_network_architectures.vision_branch import VisionBranch
|
| 23 |
+
from transformers import BertTokenizer
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
model_root = os.environ.get("MODEL_ROOT", "../ckpt")
|
| 27 |
+
configs_root = os.environ.get("CONFIGS_ROOT", "../ckpt")
|
| 28 |
+
|
| 29 |
+
def masks_to_boxes_3d(masks):
|
| 30 |
+
"""Compute the bounding boxes around the provided 3D masks
|
| 31 |
+
|
| 32 |
+
The masks should be in format [N, D, H, W] where N is the number of masks, (D, H, W) are the spatial dimensions.
|
| 33 |
+
|
| 34 |
+
Returns a [N, 6] tensor, with the boxes in min_x, min_y, min_z, max_x, max_y, max_z format
|
| 35 |
+
"""
|
| 36 |
+
if masks.numel() == 0:
|
| 37 |
+
return torch.zeros((0, 6), device=masks.device)
|
| 38 |
+
|
| 39 |
+
d, h, w = masks.shape[-3:]
|
| 40 |
+
|
| 41 |
+
z = torch.arange(0, d, dtype=torch.float, device=masks.device)
|
| 42 |
+
y = torch.arange(0, h, dtype=torch.float, device=masks.device)
|
| 43 |
+
x = torch.arange(0, w, dtype=torch.float, device=masks.device)
|
| 44 |
+
|
| 45 |
+
z, y, x = torch.meshgrid(z, y, x, indexing='ij')
|
| 46 |
+
|
| 47 |
+
x_mask = (masks * x.unsqueeze(0))
|
| 48 |
+
x_max = x_mask.flatten(1).max(-1).values
|
| 49 |
+
x_min = x_mask.masked_fill(~masks.bool(), float('inf')).flatten(1).min(-1).values
|
| 50 |
+
|
| 51 |
+
y_mask = (masks * y.unsqueeze(0))
|
| 52 |
+
y_max = y_mask.flatten(1).max(-1).values
|
| 53 |
+
y_min = y_mask.masked_fill(~masks.bool(), float('inf')).flatten(1).min(-1).values
|
| 54 |
+
|
| 55 |
+
z_mask = (masks * z.unsqueeze(0))
|
| 56 |
+
z_max = z_mask.flatten(1).max(-1).values
|
| 57 |
+
z_min = z_mask.masked_fill(~masks.bool(), float('inf')).flatten(1).min(-1).values
|
| 58 |
+
|
| 59 |
+
return torch.stack([x_min, y_min, z_min, x_max, y_max, z_max], dim=1)
|
| 60 |
+
|
| 61 |
+
def collate_fn(batch):
|
| 62 |
+
return batch[0]
|
| 63 |
+
|
| 64 |
+
@torch.no_grad()
|
| 65 |
+
def all_gather(data):
|
| 66 |
+
world_size = dist.get_world_size()
|
| 67 |
+
if world_size == 1:
|
| 68 |
+
return [data]
|
| 69 |
+
data_list = [None] * world_size
|
| 70 |
+
dist.all_gather_object(data_list, data)
|
| 71 |
+
return data_list
|
| 72 |
+
|
| 73 |
+
def _get_scan_interval(
|
| 74 |
+
image_size: Sequence[int], roi_size: Sequence[int], num_spatial_dims: int, overlap: float
|
| 75 |
+
) -> Tuple[int, ...]:
|
| 76 |
+
"""
|
| 77 |
+
Compute scan interval according to the image size, roi size and overlap.
|
| 78 |
+
Scan interval will be `int((1 - overlap) * roi_size)`, if interval is 0,
|
| 79 |
+
use 1 instead to make sure sliding window works.
|
| 80 |
+
|
| 81 |
+
"""
|
| 82 |
+
if len(image_size) != num_spatial_dims:
|
| 83 |
+
raise ValueError("image coord different from spatial dims.")
|
| 84 |
+
if len(roi_size) != num_spatial_dims:
|
| 85 |
+
raise ValueError("roi coord different from spatial dims.")
|
| 86 |
+
|
| 87 |
+
scan_interval = []
|
| 88 |
+
for i in range(num_spatial_dims):
|
| 89 |
+
if roi_size[i] == image_size[i]:
|
| 90 |
+
scan_interval.append(int(roi_size[i]))
|
| 91 |
+
else:
|
| 92 |
+
interval = int(roi_size[i] * (1 - overlap))
|
| 93 |
+
scan_interval.append(interval if interval > 0 else 1)
|
| 94 |
+
return tuple(scan_interval)
|
| 95 |
+
|
| 96 |
+
def center_crop(image, mask, crop_size):
|
| 97 |
+
x_min, y_min, z_min, x_max, y_max, z_max = masks_to_boxes_3d(mask)[0].long()
|
| 98 |
+
|
| 99 |
+
crop_d, crop_h, crop_w = max(crop_size[0], z_max - z_min), max(crop_size[1], y_max - y_min), max(crop_size[2], x_max - x_min)
|
| 100 |
+
|
| 101 |
+
cx = (x_min + x_max) // 2
|
| 102 |
+
cy = (y_min + y_max) // 2
|
| 103 |
+
cz = (z_min + z_max) // 2
|
| 104 |
+
|
| 105 |
+
d, h, w = image.shape[-3:]
|
| 106 |
+
|
| 107 |
+
x_start = max(0, cx - crop_w // 2)
|
| 108 |
+
x_end = min(w, x_start + crop_w)
|
| 109 |
+
if x_end - x_start < crop_w:
|
| 110 |
+
x_start = max(0, x_end - crop_w)
|
| 111 |
+
|
| 112 |
+
y_start = max(0, cy - crop_h // 2)
|
| 113 |
+
y_end = min(h, y_start + crop_h)
|
| 114 |
+
if y_end - y_start < crop_h:
|
| 115 |
+
y_start = max(0, y_end - crop_h)
|
| 116 |
+
|
| 117 |
+
z_start = max(0, cz - crop_d // 2)
|
| 118 |
+
z_end = min(d, z_start + crop_d)
|
| 119 |
+
if z_end - z_start < crop_d:
|
| 120 |
+
z_start = max(0, z_end - crop_d)
|
| 121 |
+
|
| 122 |
+
return image[..., z_start:z_end, y_start:y_end, x_start:x_end], mask[..., z_start:z_end, y_start:y_end, x_start:x_end]
|
| 123 |
+
|
| 124 |
+
class DataFolder(Dataset):
|
| 125 |
+
def __init__(self, img_dir):
|
| 126 |
+
super().__init__()
|
| 127 |
+
|
| 128 |
+
patient_list = os.listdir(img_dir)
|
| 129 |
+
self.img_paths = [
|
| 130 |
+
os.path.join(img_dir, p)
|
| 131 |
+
for p in patient_list
|
| 132 |
+
]
|
| 133 |
+
|
| 134 |
+
self.pad_func = transforms.SpatialPadd(
|
| 135 |
+
keys=["image"],
|
| 136 |
+
spatial_size=(96, 256, 384),
|
| 137 |
+
mode='constant',
|
| 138 |
+
constant_values=0,
|
| 139 |
+
method="end"
|
| 140 |
+
)
|
| 141 |
+
|
| 142 |
+
self.organs = [
|
| 143 |
+
'肾上腺', '主动脉', '竖脊肌', '脑', '锁骨', '大肠', '十二指肠', '食管', '面部', '股骨',
|
| 144 |
+
'胆囊', "臀肌", '心脏', '髋关节', '肱骨', '髂动脉', '髂静脉', '髂腰肌', '下腔静脉', '肾',
|
| 145 |
+
'肝', '肺', '胰腺', '门静脉', '肺动脉', '肋骨', '骶骨', '肩胛骨', '小肠', '脾',
|
| 146 |
+
'胃', '气管', '膀胱', '颈椎', '腰椎', '胸椎'
|
| 147 |
+
]
|
| 148 |
+
# organ_dict = {
|
| 149 |
+
# "肾上腺": "adrenal gland", "主动脉": "aorta", "竖脊肌": "erector spinae muscle", "脑": "brain", "锁骨": "clavicle", "大肠": "large bowel", "十二指肠": "duodenum",
|
| 150 |
+
# "食管": "esophagus", "面部": "face", "股骨": "femur", "胆囊": "gallbladder", "臀肌": "gluteus muscle", "心脏": "heart", "髋关节": "hip joint", "肱骨": "humerus",
|
| 151 |
+
# "髂动脉": "iliac artery", "髂静脉": "iliac vena", "髂腰肌": "iliopsoas muscle", "下腔静脉": "inferior vena cava", "肾": "kidney", "肝": "liver", "肺": "lung",
|
| 152 |
+
# "胰腺": "pancreas", "门静脉": "portal vein", "肺动脉": "pulmonary artery", "肋骨": "rib", "骶骨": "sacrum", "肩胛骨": "scapula", "小肠": "small bowel", "脾": "spleen",
|
| 153 |
+
# "胃": "stomach", "气管": "trachea", "膀胱": "bladder", "颈椎": "cervical vertebrae", "腰椎": "lumbar vertebrae", "胸椎": "thoracic vertebrae"
|
| 154 |
+
# }
|
| 155 |
+
|
| 156 |
+
self.test_items = ['主动脉_主动脉夹层', '主动脉_主动脉瘤', '主动脉_粥样硬化', '主动脉_钙化', '十二指肠_占位', '十二指肠_囊袋状突出影', '十二指肠_憩室', '十二指肠_梗阻', '十二指肠_溃疡', '大肠_克罗恩病', '大肠_大肠(壁)钙化', '大肠_急慢性(结)肠炎', '大肠_浆膜面毛糙', '大肠_溃疡性结肠炎', '大肠_直肠癌', '大肠_积液积气', '大肠_结肠癌', '大肠_肠壁毛糙', '大肠_肠壁水肿', '大肠_肠套叠', '大肠_肠憩室', '大肠_肠梗阻', '大肠_肠穿孔', '大肠_肠道扩张', '大肠_脂肪间隙模糊', '大肠_阑尾炎', '大肠_阑尾粪石', '小肠_克罗恩病', '小肠_套叠', '小肠_扭转', '小肠_梗阻', '小肠_淋巴瘤', '小肠_积气积液', '小肠_系膜指膜炎', '小肠_系膜淋巴结肿大', '小肠_肠壁增厚', '小肠_肠管扩张', '小肠_脂肪瘤', '小肠_间质瘤(胃肠间质瘤-gist)', '小肠_(急慢性)小肠炎', '心脏_心包积液', '心脏_心影(脏)增大', '肋骨_转移瘤(乳腺癌 骨转移)', '肋骨_骨折', '肋骨_骨质破坏', '肝_低密度影', '肝_格林森鞘积液', '肝_比例失调', '肝_波浪状改变', '肝_硬化', '肝_结节状强化', '肝_肝内胆管扩张', '肝_肝内胆管结石', '肝_肝内钙化灶', '肝_肝囊肿', '肝_肝细胞癌', '肝_肝胆管内高密度影', '肝_肝血管瘤', '肝_胆管癌', '肝_脂肪肝', '肝_脓肿', '肝_转移瘤', '肝_边缘不规则', '肺_斑片影', '肺_气胸', '肺_结节', '肺_肺占位', '肺_肺萎陷', '肺_胸腔积液', '肺_膨胀不全', '肺_转移瘤', '肺_钙化灶', '肺_高密度影', '肾_低密度影', '肾_囊肿', '肾_多囊肾', '肾_实质变薄', '肾_无强化囊性灶', '肾_肾动脉瘤', '肾_肾盂扩张', '肾_肾盂癌', '肾_肾盂积水', '肾_肾细胞癌(透明细胞癌)', '肾_肾萎缩', '肾_肾血管平滑肌脂肪瘤', '肾_肾(盂)结石', '肾_高密度影', '肾上腺_增生', '肾上腺_结节', '肾上腺_脂肪瘤', '肾上腺_腺瘤', '肾上腺_转移瘤', '肾上腺_钙化', '胃_壁水肿', '胃_扩张', '胃_胃底静脉曲张', '胃_胃溃疡', '胃_胃癌', '胃_间质瘤(gist)', '胆囊_结石', '胆囊_结节状致密影', '胆囊_胆囊增大', '胆囊_胆囊炎', '胆囊_胆囊癌', '胆囊_胆囊腺肌症', '胆囊_胆管壁增厚', '胆囊_胆管扩张', '胆囊_胆管炎', '胆囊_胆管癌', '胆囊_胆管积气', '胆囊_胆管结石', '胆囊_高密度影', '胆囊_黄色肉芽肿', '胰腺_低密度影', '胰腺_囊肿', '胰腺_围脂肪间隙模糊', '胰腺_肿瘤或胰腺癌', '胰腺_胰周假性囊肿', '胰腺_胰管扩张', '胰腺_胰管结石', '胰腺_胰腺炎', '胰腺_胰腺饱满', '胰腺_萎缩', '脾_低密度灶', '脾_副脾', '脾_囊肿', '脾_梗死', '脾_片状低密度区', '脾_脾大', '脾_脾脏淋巴瘤', '脾_钙化', '膀胱_憩室', '膀胱_结石', '膀胱_膀胱壁毛糙', '膀胱_膀胱炎', '膀胱_膀胱癌', '膀胱_软组织密度影', '门静脉_增宽', '门静脉_栓塞', '门静脉_高压', '食管_增粗迂曲血管影', '食管_管壁增厚', '食管_裂孔疝', '食管_静脉扩张迂曲', '食管_静脉曲张', '骶骨_骨炎']
|
| 157 |
+
self.english_mapping = {'主动脉_主动脉夹层': 'Aorta_Aortic dissection', '主动脉_主动脉瘤': 'Aorta_Aortic aneurysm', '主动脉_粥样硬化': 'Aorta_Atherosclerosis', '主动脉_钙化': 'Aorta_Calcification', '十二指肠_占位': 'Duodenum_Mass', '十二指肠_囊袋状突出影': 'Duodenum_Saccular outpouching', '十二指肠_憩室': 'Duodenum_Diverticulum', '十二指肠_梗阻': 'Duodenum_Obstruction', '十二指肠_溃疡': 'Duodenum_Ulcer', '大肠_克罗恩病': "Large bowel_Crohn's disease", '大肠_大肠(壁)钙化': 'Large bowel_Mural calcification', '大肠_急慢性(结)肠炎': 'Large bowel_Colitis', '大肠_浆膜面毛糙': 'Large bowel_Serosal surface irregularity', '大肠_溃疡性结肠炎': 'Large bowel_Ulcerative colitis', '大肠_直肠癌': 'Large bowel_Rectal cancer', '大肠_积液积气': 'Large bowel_Gas and fluid accumulation', '大肠_结肠癌': 'Large bowel_Colon cancer', '大肠_肠壁毛糙': 'Large bowel_Wall irregularity', '大肠_肠壁水肿': 'Large bowel_Wall edema', '大肠_肠套叠': 'Large bowel_Intussusception', '大肠_肠憩室': 'Large bowel_Diverticulum', '大肠_肠梗阻': 'Large bowel_Obstruction', '大肠_肠穿孔': 'Large bowel_Perforation', '大肠_肠道扩张': 'Large bowel_Dilatation', '大肠_脂肪间隙模糊': 'Large bowel_Blurring of fat planes', '大肠_阑尾炎': 'Large bowel_Appendicitis', '大肠_阑尾粪石': 'Large bowel_Appendicolith', '小肠_克罗恩病': "Small bowel_Crohn's disease", '小肠_套叠': 'Small bowel_Intussusception', '小肠_扭转': 'Small bowel_Volvulus', '小肠_梗阻': 'Small bowel_Obstruction', '小肠_淋巴瘤': 'Small bowel_Lymphoma', '小肠_积气积液': 'Small bowel_Gas and fluid accumulation', '小肠_系膜指膜炎': 'Small bowel_Mesenteric panniculitis', '小肠_系膜淋巴结肿大': 'Small bowel_Mesenteric lymphadenopathy', '小肠_肠壁增厚': 'Small bowel_Wall thickening', '小肠_肠管扩张': 'Small bowel_Dilatation', '小肠_脂肪瘤': 'Small bowel_Lipoma', '小肠_间质瘤(胃肠间质瘤-gist)': 'Small bowel_Gastrointestinal stromal tumor', '小肠_(急慢性)小肠炎': 'Small bowel_Enteritis', '心脏_心包积液': 'Heart_Pericardial effusion', '心脏_心影(脏)增大': 'Heart_Cardiomegaly', '肋骨_转移瘤(乳腺癌 骨转移)': 'Rib_Metastasis', '肋骨_骨折': 'Rib_Fracture', '肋骨_骨质破坏': 'Rib_Bone destruction', '肝_低密度影': 'Liver_Hypoattenuating lesion', '肝_格林森鞘积液': 'Liver_Periportal edema', '肝_比例失调': 'Liver_Lobar volume disproportion', '肝_波浪状改变': 'Liver_Undulating contour', '肝_硬化': 'Liver_Cirrhosis', '肝_结节状强化': 'Liver_Nodular enhancement', '肝_肝内胆管扩张': 'Liver_Intrahepatic bile duct dilatation', '肝_肝内胆管结石': 'Liver_Hepatolithiasis', '肝_肝内钙化灶': 'Liver_Intrahepatic calcification', '肝_肝囊肿': 'Liver_Cyst', '肝_肝细胞癌': 'Liver_Hepatocellular carcinoma', '肝_肝胆管内高密度影': 'Liver_Hyperattenuating lesion in intrahepatic bile ducts', '肝_肝血管瘤': 'Liver_Hemangioma', '肝_胆管癌': 'Liver_Intrahepatic cholangiocarcinoma', '肝_脂肪肝': 'Liver_Steatotic liver disease', '肝_脓肿': 'Liver_Abscess', '肝_转移瘤': 'Liver_Metastasis', '肝_边缘不规则': 'Liver_Irregular margin', '肺_斑片影': 'Lung_Patchy opacity', '肺_气胸': 'Lung_Pneumothorax', '肺_结节': 'Lung_Nodule', '肺_肺占位': 'Lung_Mass', '肺_肺萎陷': 'Lung_Pulmonary collapse', '肺_胸腔积液': 'Lung_Pleural effusion', '肺_膨胀不全': 'Lung_Atelectasis', '肺_转移瘤': 'Lung_Metastasis', '肺_钙化灶': 'Lung_Calcification', '肺_高密度影': 'Lung_Hyperattenuating opacity', '肾_低密度影': 'Kidney_Hypoattenuating lesion', '肾_囊肿': 'Kidney_Cyst', '肾_多囊肾': 'Kidney_Polycystic kidney disease', '肾_实质变薄': 'Kidney_Parenchymal thinning', '肾_无强化囊性灶': 'Kidney_Nonenhancing cystic lesion', '肾_肾动脉瘤': 'Kidney_Renal artery aneurysm', '肾_肾盂扩张': 'Kidney_Renal pelvic dilatation', '肾_肾盂癌': 'Kidney_Renal pelvic cancer', '肾_肾盂积水': 'Kidney_Hydronephrosis', '肾_肾细胞癌(透明细胞癌)': 'Kidney_Renal cell carcinoma', '肾_肾萎缩': 'Kidney_Atrophy', '肾_肾血管平滑肌脂肪瘤': 'Kidney_Angiomyolipoma', '肾_肾(盂)结石': 'Kidney_Nephrolithiasis', '肾_高密度影': 'Kidney_Hyperattenuating lesion', '肾上腺_增生': 'Adrenal gland_Hyperplasia', '肾上腺_结节': 'Adrenal gland_Nodule', '肾上腺_脂肪瘤': 'Adrenal gland_Lipoma', '肾上腺_腺瘤': 'Adrenal gland_Adenoma', '肾上腺_转移瘤': 'Adrenal gland_Metastasis', '肾上腺_钙化': 'Adrenal gland_Calcification', '胃_壁水肿': 'Stomach_Wall edema', '胃_扩张': 'Stomach_Dilatation', '胃_胃底静脉曲张': 'Stomach_Gastric fundal varices', '胃_胃溃疡': 'Stomach_Ulcer', '胃_胃癌': 'Stomach_Gastric cancer', '胃_间质瘤(gist)': 'Stomach_Gastrointestinal stromal tumor (GIST)', '胆囊_结石': 'Gallbladder_Cholecystolithiasis', '胆囊_结节状致密影': 'Gallbladder_Nodular stone-like hyperattenuating lesion', '胆囊_胆囊增大': 'Gallbladder_Distention', '胆囊_胆囊炎': 'Gallbladder_Cholecystitis', '胆囊_胆囊癌': 'Gallbladder_Gallbladder cancer', '胆囊_胆囊腺肌症': 'Gallbladder_Adenomyomatosis', '胆囊_胆管壁增厚': 'Gallbladder_Extrahepatic bile duct wall thickening', '胆囊_胆管扩张': 'Gallbladder_Extrahepatic bile duct dilatation', '胆囊_胆管炎': 'Gallbladder_Cholangitis', '胆囊_胆管癌': 'Gallbladder_Cholangiocarcinoma', '胆囊_胆管积气': 'Gallbladder_Pneumobilia', '胆囊_胆管结石': 'Gallbladder_Extrahepatic bile duct stone', '胆囊_高密度影': 'Gallbladder_Hyperattenuating lesion', '胆囊_黄色肉芽肿': 'Gallbladder_Xanthogranuloma', '胰腺_低密度影': 'Pancreas_Low-density lesion', '胰腺_囊肿': 'Pancreas_Cyst', '胰腺_围脂肪间隙模糊': 'Pancreas_Blurring of peripancreatic fat planes', '胰腺_肿瘤或胰腺癌': 'Pancreas_Pancreatic cancer', '胰腺_胰周假性囊肿': 'Pancreas_Peripancreatic pseudocyst', '胰腺_胰管扩张': 'Pancreas_Pancreatic duct dilatation', '胰腺_胰管结石': 'Pancreas_Pancreatic duct calculus', '胰腺_胰腺炎': 'Pancreas_Pancreatitis', '胰腺_胰腺饱满': 'Pancreas_Enlargement', '胰腺_萎缩': 'Pancreas_Atrophy', '脾_低密度灶': 'Spleen_Hypoattenuating lesion', '脾_副脾': 'Spleen_Accessory spleen', '脾_囊肿': 'Spleen_Cyst', '脾_梗死': 'Spleen_Infarction', '脾_片状低密度区': 'Spleen_Patchy hypoattenuating lesion', '脾_脾大': 'Spleen_Splenomegaly', '脾_脾脏淋巴瘤': 'Spleen_Lymphoma', '脾_钙化': 'Spleen_Calcification', '膀胱_憩室': 'Bladder_Diverticulum', '膀胱_结石': 'Bladder_Stone', '膀胱_膀胱壁毛糙': 'Bladder_Wall irregularity', '膀胱_膀胱炎': 'Bladder_Cystitis', '膀胱_膀胱癌': 'Bladder_Bladder cancer', '膀胱_软组织密度影': 'Bladder_Soft-tissue attenuation lesion', '门静脉_增宽': 'Portal vein_Dilatation', '门静脉_栓塞': 'Portal vein_Thrombosis', '门静脉_高压': 'Portal vein_Hypertension', '食管_增粗迂曲血管影': 'Esophagus_Dilated and tortuous tubular opacities', '食管_管壁增厚': 'Esophagus_Wall thickening', '食管_裂孔疝': 'Esophagus_Hiatal hernia', '食管_静脉扩张迂曲': 'Esophagus_Dilated and tortuous veins', '食管_静脉曲张': 'Esophagus_Varices', '骶骨_骨炎': 'Sacrum_Osteitis'}
|
| 158 |
+
self.test_organs = list(set([item.split('_')[0] for item in self.test_items]))
|
| 159 |
+
|
| 160 |
+
def __len__(self):
|
| 161 |
+
return len(self.img_paths)
|
| 162 |
+
|
| 163 |
+
def __getitem__(self, index):
|
| 164 |
+
# load image
|
| 165 |
+
image_path = self.img_paths[index]
|
| 166 |
+
data = {"image": image_path}
|
| 167 |
+
res = transforms.LoadImaged(keys=["image"], image_only=False, ensure_channel_first=True)(data)
|
| 168 |
+
image = res["image"]
|
| 169 |
+
|
| 170 |
+
affine = res["image_meta_dict"]["affine"]
|
| 171 |
+
spacing = (
|
| 172 |
+
abs(affine[0, 0].item()),
|
| 173 |
+
abs(affine[1, 1].item()),
|
| 174 |
+
abs(affine[2, 2].item())
|
| 175 |
+
)
|
| 176 |
+
_, h, w, d = image.shape
|
| 177 |
+
orig_shape_hwd = (h, w, d)
|
| 178 |
+
|
| 179 |
+
ref_spacing = (1.0, 1.0, 5.0)
|
| 180 |
+
scale = [spacing[i] / ref_spacing[i] for i in range(3)]
|
| 181 |
+
target_size = [int(h * scale[1]), int(w * scale[0]), int(d * scale[2])] # [H', W', D']
|
| 182 |
+
|
| 183 |
+
trans = transforms.Compose(
|
| 184 |
+
[
|
| 185 |
+
transforms.Resized(spatial_size=target_size, keys=["image"], mode="trilinear"),
|
| 186 |
+
transforms.Transposed(keys=["image"], indices=(0, 3, 2, 1)),
|
| 187 |
+
]
|
| 188 |
+
)
|
| 189 |
+
resized_data = trans(res)
|
| 190 |
+
|
| 191 |
+
img_resized = resized_data["image"] # [C, D', W', H']
|
| 192 |
+
image = img_resized
|
| 193 |
+
image[image > 400] = 400
|
| 194 |
+
image[image < -300] = -300
|
| 195 |
+
image = (image - image.min()) / (image.max() - image.min() + 1e-8)
|
| 196 |
+
img = image
|
| 197 |
+
|
| 198 |
+
# crop non-zero region in image
|
| 199 |
+
roi_coords = np.nonzero(img[0].cpu().numpy())
|
| 200 |
+
min_dhw = torch.from_numpy(np.min(roi_coords, axis=1))
|
| 201 |
+
max_dhw = torch.from_numpy(np.max(roi_coords, axis=1))
|
| 202 |
+
|
| 203 |
+
extend_d = 5
|
| 204 |
+
extend_hw = 20
|
| 205 |
+
|
| 206 |
+
min_dhw = torch.max(
|
| 207 |
+
min_dhw - torch.tensor([extend_d, extend_hw, extend_hw]),
|
| 208 |
+
torch.tensor([0, 0, 0]),
|
| 209 |
+
)
|
| 210 |
+
max_dhw = torch.min(
|
| 211 |
+
max_dhw + torch.tensor([extend_d, extend_hw, extend_hw]),
|
| 212 |
+
torch.tensor([img.shape[1], img.shape[2], img.shape[3]]),
|
| 213 |
+
)
|
| 214 |
+
|
| 215 |
+
cropped_image = img[
|
| 216 |
+
:,
|
| 217 |
+
min_dhw[0]: max_dhw[0],
|
| 218 |
+
min_dhw[1]: max_dhw[1],
|
| 219 |
+
min_dhw[2]: max_dhw[2]
|
| 220 |
+
]
|
| 221 |
+
crop_shape_dhw = tuple(cropped_image.shape[1:])
|
| 222 |
+
|
| 223 |
+
# pad data to [96, 256, 384] if smaller
|
| 224 |
+
data["image"] = cropped_image
|
| 225 |
+
data_pad = self.pad_func(data)
|
| 226 |
+
data = data_pad
|
| 227 |
+
|
| 228 |
+
file_name = image_path.split('/')[-1]
|
| 229 |
+
patient_id = file_name.split('_')[0]
|
| 230 |
+
test_organ_names = self.test_organs
|
| 231 |
+
|
| 232 |
+
meta_info = {
|
| 233 |
+
'file_name': file_name,
|
| 234 |
+
'img_path': image_path,
|
| 235 |
+
'patient_id': patient_id,
|
| 236 |
+
'test_organ_names': test_organ_names,
|
| 237 |
+
'letter': 'None',
|
| 238 |
+
}
|
| 239 |
+
return data['image'].as_tensor(), self.test_items, meta_info
|
| 240 |
+
|
| 241 |
+
|
| 242 |
+
class RADAR(nn.Module):
|
| 243 |
+
def __init__(
|
| 244 |
+
self,
|
| 245 |
+
image_encoder,
|
| 246 |
+
text_encoder,
|
| 247 |
+
text_decoder=None,
|
| 248 |
+
queue_size=1234,
|
| 249 |
+
alpha=0.4,
|
| 250 |
+
embed_dim=256,
|
| 251 |
+
momentum=0.995,
|
| 252 |
+
tie_enc_dec_weights=True,
|
| 253 |
+
max_txt_len=175,
|
| 254 |
+
):
|
| 255 |
+
super().__init__()
|
| 256 |
+
|
| 257 |
+
self.tokenizer = BertTokenizer.from_pretrained(os.path.join(configs_root, "bert-base-chinese"))
|
| 258 |
+
|
| 259 |
+
text_encoder.resize_token_embeddings(len(self.tokenizer))
|
| 260 |
+
self.visual_encoder = image_encoder
|
| 261 |
+
self.text_encoder = text_encoder
|
| 262 |
+
|
| 263 |
+
text_width = text_encoder.config.hidden_size
|
| 264 |
+
vision_width = 256
|
| 265 |
+
|
| 266 |
+
self.text_proj = nn.Linear(text_width, embed_dim)
|
| 267 |
+
|
| 268 |
+
self.queue_size = queue_size
|
| 269 |
+
self.momentum = momentum
|
| 270 |
+
self.temp = nn.Parameter(0.07 * torch.ones([]))
|
| 271 |
+
|
| 272 |
+
self.alpha = alpha
|
| 273 |
+
self.max_txt_len = max_txt_len
|
| 274 |
+
|
| 275 |
+
self.organs = [
|
| 276 |
+
'肾上腺', '主动脉', '竖脊肌', '脑', '锁骨', '大肠', '十二指肠', '食管', '面部', '股骨',
|
| 277 |
+
'胆囊', "臀肌", '心脏', '髋关节', '肱骨', '髂动脉', '髂静脉', '髂腰肌', '下腔静脉', '肾',
|
| 278 |
+
'肝', '肺', '胰腺', '门静脉', '肺动脉', '肋骨', '骶骨', '肩胛骨', '小肠', '脾',
|
| 279 |
+
'胃', '气管', '膀胱', '颈椎', '腰椎', '胸椎'
|
| 280 |
+
]
|
| 281 |
+
# organ_dict = {
|
| 282 |
+
# "肾上腺": "adrenal gland", "主动脉": "aorta", "竖脊肌": "erector spinae muscle", "脑": "brain", "锁骨": "clavicle", "大肠": "large bowel", "十二指肠": "duodenum",
|
| 283 |
+
# "食管": "esophagus", "面部": "face", "股骨": "femur", "胆囊": "gallbladder", "臀肌": "gluteus muscle", "心脏": "heart", "髋关节": "hip joint", "肱骨": "humerus",
|
| 284 |
+
# "髂动脉": "iliac artery", "髂静脉": "iliac vena", "髂腰肌": "iliopsoas muscle", "下腔静脉": "inferior vena cava", "肾": "kidney", "肝": "liver", "肺": "lung",
|
| 285 |
+
# "胰腺": "pancreas", "门静脉": "portal vein", "肺动脉": "pulmonary artery", "肋骨": "rib", "骶骨": "sacrum", "肩胛骨": "scapula", "小肠": "small bowel", "脾": "spleen",
|
| 286 |
+
# "胃": "stomach", "气管": "trachea", "膀胱": "bladder", "颈椎": "cervical vertebrae", "腰椎": "lumbar vertebrae", "胸椎": "thoracic vertebrae"
|
| 287 |
+
# }
|
| 288 |
+
|
| 289 |
+
self.attention = nn.MultiheadAttention(
|
| 290 |
+
embed_dim=vision_width,
|
| 291 |
+
num_heads=4,
|
| 292 |
+
dropout=0.1,
|
| 293 |
+
batch_first=True
|
| 294 |
+
)
|
| 295 |
+
|
| 296 |
+
self.vision_projs = nn.ModuleList([nn.Linear(vision_width, embed_dim) for _ in range(len(self.organs))])
|
| 297 |
+
self.query_tokens = nn.Parameter(torch.zeros(len(self.organs), vision_width))
|
| 298 |
+
|
| 299 |
+
@torch.inference_mode()
|
| 300 |
+
def forward_test_win(
|
| 301 |
+
self,
|
| 302 |
+
images,
|
| 303 |
+
masks,
|
| 304 |
+
organ_logits,
|
| 305 |
+
test_organs,
|
| 306 |
+
text_feat_dict,
|
| 307 |
+
organ_feat_dict,
|
| 308 |
+
whole_organ_sizes,
|
| 309 |
+
skip_organ=None
|
| 310 |
+
):
|
| 311 |
+
seg_probs, seg, image_embeds1, image_embeds2, image_embeds3, organ_token_flags1, organ_token_flags2, organ_token_flags3 = self.visual_encoder(images, None)
|
| 312 |
+
|
| 313 |
+
margin = 2
|
| 314 |
+
masks = seg
|
| 315 |
+
|
| 316 |
+
for i, (embed1, embed2, embed3, mask) in enumerate(zip(image_embeds1, image_embeds2, image_embeds3, masks)):
|
| 317 |
+
boundaries = []
|
| 318 |
+
for d in range(mask.dim()):
|
| 319 |
+
start_slice = [slice(None)] * mask.dim()
|
| 320 |
+
end_slice = [slice(None)] * mask.dim()
|
| 321 |
+
|
| 322 |
+
start_slice[d] = slice(None, margin)
|
| 323 |
+
end_slice[d] = slice(-margin, None)
|
| 324 |
+
|
| 325 |
+
boundaries.append(mask[tuple(start_slice)][mask[tuple(start_slice)] > 0])
|
| 326 |
+
boundaries.append(mask[tuple(end_slice)][mask[tuple(end_slice)] > 0])
|
| 327 |
+
boundaries = torch.cat(boundaries)
|
| 328 |
+
|
| 329 |
+
boundary_values = boundaries[boundaries > 0].flatten()
|
| 330 |
+
boundary_organs = torch.unique(boundary_values)
|
| 331 |
+
|
| 332 |
+
if skip_organ is not None:
|
| 333 |
+
boundary_organs = boundary_organs[boundary_organs != skip_organ + 1]
|
| 334 |
+
|
| 335 |
+
organ_ids, organ_counts = torch.unique(mask, return_counts=True)
|
| 336 |
+
organ_ids = organ_ids.long()
|
| 337 |
+
organ_counts = organ_counts[organ_ids != 0]
|
| 338 |
+
organ_ids = organ_ids[organ_ids != 0]
|
| 339 |
+
|
| 340 |
+
# organs not touch boundary
|
| 341 |
+
intact_organ_ids = [organ_id for organ_id, organ_count in zip(organ_ids, organ_counts) if organ_id not in boundary_organs]
|
| 342 |
+
intact_organ_ids = torch.tensor(intact_organ_ids, device=masks.device).long()
|
| 343 |
+
intact_organ_ids = intact_organ_ids - 1
|
| 344 |
+
|
| 345 |
+
if not len(intact_organ_ids):
|
| 346 |
+
continue
|
| 347 |
+
|
| 348 |
+
organ_sizes = dict(zip([self.organs[organ_id] for organ_id in intact_organ_ids], [organ_counts[organ_ids == organ_id + 1].item() for organ_id in intact_organ_ids]))
|
| 349 |
+
|
| 350 |
+
for organ_id in intact_organ_ids:
|
| 351 |
+
organ_name = self.organs[organ_id.item()]
|
| 352 |
+
if organ_name not in test_organs:
|
| 353 |
+
continue
|
| 354 |
+
|
| 355 |
+
if organ_name in organ_feat_dict:
|
| 356 |
+
continue
|
| 357 |
+
|
| 358 |
+
tokens1 = organ_token_flags1[i, organ_id, :]
|
| 359 |
+
tokens2 = organ_token_flags2[i, organ_id, :]
|
| 360 |
+
tokens3 = organ_token_flags3[i, organ_id, :]
|
| 361 |
+
|
| 362 |
+
query = self.query_tokens[organ_id].unsqueeze(0).unsqueeze(0)
|
| 363 |
+
key1 = embed1[tokens1].unsqueeze(0)
|
| 364 |
+
key2 = embed2[tokens2].unsqueeze(0)
|
| 365 |
+
key3 = embed3[tokens3].unsqueeze(0)
|
| 366 |
+
|
| 367 |
+
key = value = torch.cat([key1, key2, key3], dim=1)
|
| 368 |
+
|
| 369 |
+
updated_query_token, _ = self.attention(query, key, value)
|
| 370 |
+
updated_query_token = updated_query_token.squeeze(0)
|
| 371 |
+
|
| 372 |
+
image_feat = F.normalize(self.vision_projs[organ_id](updated_query_token), dim=-1)
|
| 373 |
+
|
| 374 |
+
organ_feat_dict[organ_name] = image_feat.cpu().tolist()
|
| 375 |
+
|
| 376 |
+
for item in organ_logits.keys():
|
| 377 |
+
if isinstance(item, str):
|
| 378 |
+
item_organ_name = item.split('_')[0]
|
| 379 |
+
else:
|
| 380 |
+
item_organ_name = item[0]
|
| 381 |
+
if item_organ_name != organ_name:
|
| 382 |
+
continue
|
| 383 |
+
|
| 384 |
+
text_feat = text_feat_dict[item]
|
| 385 |
+
|
| 386 |
+
logits = image_feat @ text_feat.t() / self.temp
|
| 387 |
+
probs = logits.softmax(-1)
|
| 388 |
+
organ_logits[item].append(probs.cpu().tolist())
|
| 389 |
+
|
| 390 |
+
return organ_logits, seg_probs
|
| 391 |
+
|
| 392 |
+
@torch.inference_mode()
|
| 393 |
+
def evaluate(pad_func, model, img_dir, save_dir, save_tag):
|
| 394 |
+
|
| 395 |
+
datafolder = DataFolder(img_dir)
|
| 396 |
+
dataloader = DataLoader(
|
| 397 |
+
datafolder,
|
| 398 |
+
batch_size=1,
|
| 399 |
+
shuffle=False,
|
| 400 |
+
num_workers=0, # Space patch: single-volume inference; avoids forking 12 torch workers in the CPU container.
|
| 401 |
+
drop_last=False,
|
| 402 |
+
collate_fn=collate_fn
|
| 403 |
+
)
|
| 404 |
+
|
| 405 |
+
sw_batch_size = 1
|
| 406 |
+
overlap = 0.25
|
| 407 |
+
roi_size = (96, 256, 384)
|
| 408 |
+
|
| 409 |
+
miss_num = 0
|
| 410 |
+
results = []
|
| 411 |
+
organ_status = {}
|
| 412 |
+
|
| 413 |
+
# load pos/neg ensembled prompt embeddings
|
| 414 |
+
# Space patch: resolve via MODEL_ROOT (absolute in the Space) instead of cwd-relative '../ckpt'.
|
| 415 |
+
text_feat_dict = torch.load(os.path.join(model_root, 'infer_text_embedding_radar.pt'))
|
| 416 |
+
organ_feat_dict = {}
|
| 417 |
+
save_path = os.path.join(save_dir, f'RADAR_infer_results_{save_tag}.csv')
|
| 418 |
+
os.makedirs(save_dir, exist_ok=True)
|
| 419 |
+
|
| 420 |
+
for i, (image, test_items, meta_info) in enumerate(tqdm(dataloader, desc='Infer')):
|
| 421 |
+
torch.cuda.empty_cache()
|
| 422 |
+
skip_case = False
|
| 423 |
+
for tmp_s in image.shape[1:]:
|
| 424 |
+
if tmp_s > 1000:
|
| 425 |
+
skip_case = True
|
| 426 |
+
break
|
| 427 |
+
if skip_case:
|
| 428 |
+
continue
|
| 429 |
+
|
| 430 |
+
fid = meta_info['file_name']
|
| 431 |
+
organ_feat_dict[fid] = {}
|
| 432 |
+
|
| 433 |
+
image = image[None].cuda()
|
| 434 |
+
|
| 435 |
+
test_organs = meta_info['test_organ_names']
|
| 436 |
+
|
| 437 |
+
image_size = list(image.shape[2:])
|
| 438 |
+
num_spatial_dims = len(image.shape) - 2
|
| 439 |
+
|
| 440 |
+
scan_interval = _get_scan_interval(
|
| 441 |
+
image_size, roi_size, num_spatial_dims, overlap
|
| 442 |
+
)
|
| 443 |
+
slices = dense_patch_slices(image_size, roi_size, scan_interval)
|
| 444 |
+
num_win = len(slices)
|
| 445 |
+
organ_logits = dict(zip(test_items, [[] for _ in test_items]))
|
| 446 |
+
# organ_logits.pop('胆囊_术后胆囊缺失') # surgically_absent_gallbladder
|
| 447 |
+
|
| 448 |
+
# get full mask
|
| 449 |
+
full_mask = torch.zeros((1, 37) + tuple(image_size)).cuda()
|
| 450 |
+
count_map = torch.zeros_like(full_mask).cuda()
|
| 451 |
+
|
| 452 |
+
for slice_g in range(0, num_win, sw_batch_size):
|
| 453 |
+
slice_range = range(slice_g, min(slice_g + sw_batch_size, num_win))
|
| 454 |
+
unravel_slice = [
|
| 455 |
+
[slice(int(idx / num_win), int(idx / num_win) + 1), slice(None)] + list(slices[idx % num_win])
|
| 456 |
+
for idx in slice_range
|
| 457 |
+
]
|
| 458 |
+
|
| 459 |
+
window_patches = torch.cat([image[win_slice] for win_slice in unravel_slice]).cuda()
|
| 460 |
+
|
| 461 |
+
organ_logits, pred_window_seg_prob = model.forward_test_win(
|
| 462 |
+
window_patches,
|
| 463 |
+
None,
|
| 464 |
+
organ_logits,
|
| 465 |
+
test_organs,
|
| 466 |
+
text_feat_dict,
|
| 467 |
+
organ_feat_dict[fid],
|
| 468 |
+
None
|
| 469 |
+
)
|
| 470 |
+
|
| 471 |
+
# interpolate
|
| 472 |
+
interpolated_seg_prob = F.interpolate(pred_window_seg_prob, size=window_patches.shape[2:], mode='trilinear')
|
| 473 |
+
|
| 474 |
+
for ii, slice_idx in enumerate(slice_range):
|
| 475 |
+
full_slice = unravel_slice[ii]
|
| 476 |
+
full_mask[full_slice] += interpolated_seg_prob[ii]
|
| 477 |
+
count_map[full_slice] += 1
|
| 478 |
+
|
| 479 |
+
# Avoid division by zero by ensuring count_map is at least 1 everywhere
|
| 480 |
+
count_map = torch.clamp(count_map, min=1)
|
| 481 |
+
stitched_mask = full_mask / count_map # argmax
|
| 482 |
+
stitched_mask = stitched_mask.argmax(1).unsqueeze(0)
|
| 483 |
+
|
| 484 |
+
margin = 2
|
| 485 |
+
boundaries = []
|
| 486 |
+
squeeze_stitched_mask = stitched_mask.squeeze(0).squeeze(0)
|
| 487 |
+
for d in range(squeeze_stitched_mask.dim()):
|
| 488 |
+
start_slice = [slice(None)] * squeeze_stitched_mask.dim()
|
| 489 |
+
end_slice = [slice(None)] * squeeze_stitched_mask.dim()
|
| 490 |
+
|
| 491 |
+
start_slice[d] = slice(None, margin)
|
| 492 |
+
end_slice[d] = slice(-margin, None)
|
| 493 |
+
|
| 494 |
+
boundaries.append(squeeze_stitched_mask[tuple(start_slice)][squeeze_stitched_mask[tuple(start_slice)] > 0])
|
| 495 |
+
boundaries.append(squeeze_stitched_mask[tuple(end_slice)][squeeze_stitched_mask[tuple(end_slice)] > 0])
|
| 496 |
+
boundaries = torch.cat(boundaries)
|
| 497 |
+
|
| 498 |
+
boundary_values = boundaries[boundaries > 0].flatten()
|
| 499 |
+
boundary_organs = torch.unique(boundary_values)
|
| 500 |
+
|
| 501 |
+
organ_ids, organ_counts = torch.unique(squeeze_stitched_mask, return_counts=True)
|
| 502 |
+
organ_ids = organ_ids.long()
|
| 503 |
+
organ_counts = organ_counts[organ_ids != 0]
|
| 504 |
+
organ_ids = organ_ids[organ_ids != 0]
|
| 505 |
+
|
| 506 |
+
# organs not touch boundary
|
| 507 |
+
intact_organ_ids = [organ_id for organ_id, organ_count in zip(organ_ids, organ_counts) if organ_id not in boundary_organs]
|
| 508 |
+
intact_organ_ids = torch.tensor(intact_organ_ids, device=squeeze_stitched_mask.device).long()
|
| 509 |
+
intact_organ_ids = intact_organ_ids - 1
|
| 510 |
+
|
| 511 |
+
# for melrin data, we just infer all organs
|
| 512 |
+
# organ_logits = {k:v for k,v in organ_logits.items() if datafolder.organs.index(k[0]) in intact_organ_ids}
|
| 513 |
+
|
| 514 |
+
for k, v in organ_logits.items():
|
| 515 |
+
if not len(v):
|
| 516 |
+
organ_name = k.split('_')[0]
|
| 517 |
+
organ_id = datafolder.organs.index(organ_name)
|
| 518 |
+
|
| 519 |
+
window_patch, window_mask = center_crop(
|
| 520 |
+
image,
|
| 521 |
+
torch.eq(stitched_mask, organ_id + 1),
|
| 522 |
+
crop_size=roi_size
|
| 523 |
+
)
|
| 524 |
+
window_mask = window_mask.float()
|
| 525 |
+
window_mask[window_mask == 1] = organ_id + 1
|
| 526 |
+
|
| 527 |
+
pad_data = pad_func({'image': window_patch[0], 'label': window_mask[0]})
|
| 528 |
+
window_patch, window_mask = pad_data['image'], pad_data['label']
|
| 529 |
+
|
| 530 |
+
organ_logits, _ = model.forward_test_win(
|
| 531 |
+
window_patch[None],
|
| 532 |
+
None,
|
| 533 |
+
organ_logits,
|
| 534 |
+
test_organs,
|
| 535 |
+
text_feat_dict,
|
| 536 |
+
organ_feat_dict[fid],
|
| 537 |
+
None,
|
| 538 |
+
skip_organ=organ_id
|
| 539 |
+
)
|
| 540 |
+
|
| 541 |
+
res = [meta_info['file_name']] + [''] * len(datafolder.test_items)
|
| 542 |
+
organ_logits = {item: probs for item, probs in organ_logits.items() if len(probs) > 0}
|
| 543 |
+
|
| 544 |
+
for item, probs in organ_logits.items():
|
| 545 |
+
res[datafolder.test_items.index(item) + 1] = np.concatenate(probs).mean(0)[1] # get average of one organ in multi-widows
|
| 546 |
+
results.append(res)
|
| 547 |
+
|
| 548 |
+
if dist.is_initialized():
|
| 549 |
+
results = np.concatenate(all_gather(results), axis=0)
|
| 550 |
+
else:
|
| 551 |
+
results = results
|
| 552 |
+
|
| 553 |
+
pd.DataFrame(
|
| 554 |
+
results,
|
| 555 |
+
columns=['file_name'] + [f'{k} ({datafolder.english_mapping[k]})' for k in datafolder.test_items]
|
| 556 |
+
).to_csv(save_path, index=False, encoding='utf-8-sig')
|
| 557 |
+
|
| 558 |
+
def initialize():
|
| 559 |
+
"""
|
| 560 |
+
Returns: transforms.DivisiblePadd, RADAR
|
| 561 |
+
"""
|
| 562 |
+
print('\n--> Start initializing...')
|
| 563 |
+
pad_func = transforms.DivisiblePadd(
|
| 564 |
+
keys=["image", "label"],
|
| 565 |
+
k=32,
|
| 566 |
+
mode='constant',
|
| 567 |
+
constant_values=0,
|
| 568 |
+
method="end"
|
| 569 |
+
)
|
| 570 |
+
|
| 571 |
+
vision_encoder = VisionBranch()
|
| 572 |
+
text_encoder = XBertEncoder.from_config({}, from_pretrained=True)
|
| 573 |
+
|
| 574 |
+
model = RADAR(
|
| 575 |
+
image_encoder=vision_encoder,
|
| 576 |
+
text_encoder=text_encoder,
|
| 577 |
+
)
|
| 578 |
+
|
| 579 |
+
ckpt_path = os.path.join(model_root, "checkpoint_radar_pretrain.pth")
|
| 580 |
+
print('--> ckpt_path: ', ckpt_path)
|
| 581 |
+
ckpt = torch.load(ckpt_path, map_location='cpu', weights_only=False)
|
| 582 |
+
|
| 583 |
+
msg = model.load_state_dict(ckpt['model'], strict=False)
|
| 584 |
+
|
| 585 |
+
model.eval()
|
| 586 |
+
model.cuda()
|
| 587 |
+
|
| 588 |
+
print('\n--> Initialize done')
|
| 589 |
+
|
| 590 |
+
return pad_func, model
|
| 591 |
+
|
| 592 |
+
def inference(initialize_returns, img_dir, save_dir, save_tag):
|
| 593 |
+
"""
|
| 594 |
+
Args:
|
| 595 |
+
initialize_returns: pad_func, model
|
| 596 |
+
img_dir: see argparse
|
| 597 |
+
save_dir: see argparse
|
| 598 |
+
"""
|
| 599 |
+
print('\n--> Start inference.')
|
| 600 |
+
pad_func, model = initialize_returns
|
| 601 |
+
evaluate(pad_func, model, img_dir, save_dir, save_tag)
|
| 602 |
+
csv_file = os.path.join(save_dir, f'RADAR_infer_results_{save_tag}.csv')
|
| 603 |
+
print(f'evaluate done, save result_csv to {csv_file}.')
|
| 604 |
+
|
| 605 |
+
# TODO: compute metrics
|
| 606 |
+
|
| 607 |
+
|
| 608 |
+
def parse_args():
|
| 609 |
+
parser = argparse.ArgumentParser()
|
| 610 |
+
parser.add_argument('--img_dir', type=str, default='../data/demo_cases', help='The path to inference image folder.')
|
| 611 |
+
parser.add_argument('--save_dir', type=str, default='../results', help='The path to save folder.')
|
| 612 |
+
parser.add_argument('--save_tag', type=str, default='demo', help='Save tag.')
|
| 613 |
+
|
| 614 |
+
args = parser.parse_args()
|
| 615 |
+
return args
|
| 616 |
+
|
| 617 |
+
|
| 618 |
+
def main():
|
| 619 |
+
args = parse_args()
|
| 620 |
+
initialize_returns = initialize()
|
| 621 |
+
|
| 622 |
+
# infer
|
| 623 |
+
inference(initialize_returns, args.img_dir, args.save_dir, args.save_tag)
|
| 624 |
+
|
| 625 |
+
|
| 626 |
+
if __name__ == '__main__':
|
| 627 |
+
main()
|
| 628 |
+
|
| 629 |
+
|
| 630 |
+
|
README.md
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
title: RADAR Abdominal CT
|
| 3 |
+
emoji: 🩻
|
| 4 |
+
colorFrom: blue
|
| 5 |
+
colorTo: green
|
| 6 |
+
sdk: gradio
|
| 7 |
+
app_file: app.py
|
| 8 |
+
pinned: false
|
| 9 |
+
license: cc-by-nc-sa-4.0
|
| 10 |
+
python_version: "3.12"
|
| 11 |
+
preload_from_hub:
|
| 12 |
+
- radar-generalist/RADAR checkpoint_radar_pretrain.pth,bert-base-chinese/config.json,bert-base-chinese/config_decoder.json,bert-base-chinese/pytorch_model.bin,bert-base-chinese/tokenizer.json,bert-base-chinese/tokenizer_config.json,bert-base-chinese/vocab.txt
|
| 13 |
+
---
|
| 14 |
+
|
| 15 |
+
RADAR abdominal-CT findings demo: upload a contrast-enhanced abdominal CT volume (`.nii` / `.nii.gz`), get per-finding scores plus a ranked report.
|
| 16 |
+
|
| 17 |
+
- Paper/model: [radar-generalist/RADAR](https://huggingface.co/radar-generalist/RADAR) (CC BY-NC-SA 4.0, non-commercial research use only)
|
| 18 |
+
- Source: vendored inference code from the RADAR release (Zenodo `damo-radar.zip`, record `21271172`); weights load at runtime from the public HuggingFace repo, no token needed
|
| 19 |
+
- Runtime: ZeroGPU (`large`), single `@spaces.GPU(duration=180)` call per volume; preprocessing (1x1x5mm resample, 96x256x384 pad/crop) is the upstream MONAI chain, reused verbatim
|
| 20 |
+
|
| 21 |
+
Research assistance tool, not a medical diagnosis.
|
app.py
ADDED
|
@@ -0,0 +1,133 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""RADAR abdominal-CT ZeroGPU Space.
|
| 2 |
+
|
| 3 |
+
Reuse map (all inference logic is vendored, not reimplemented):
|
| 4 |
+
- `RADAR_inference/inference_demo.py::initialize` builds the RADAR model and
|
| 5 |
+
loads `ckpt/checkpoint_radar_pretrain.pth` (strict=False), then `.cuda()`.
|
| 6 |
+
- `RADAR_inference/inference_demo.py::evaluate` runs the full single-volume
|
| 7 |
+
pipeline: MONAI resample to 1x1x5mm -> HU clip [-300,400] -> min-max norm ->
|
| 8 |
+
non-zero ROI crop -> pad (96,256,384) -> sliding-window forward with
|
| 9 |
+
`inference_demo.RADAR.forward_test_win` -> per-organ center-crop second
|
| 10 |
+
pass -> CSV of mean positive-class scores for 146 organ_finding pairs.
|
| 11 |
+
- `RADAR_inference/inference_demo.py::DataFolder` owns every preprocessing
|
| 12 |
+
transform. This file adds no transforms and no report prose: it bridges
|
| 13 |
+
Gradio upload -> temp dir -> evaluate -> (Label, Dataframe of raw scores).
|
| 14 |
+
|
| 15 |
+
Space layout notes:
|
| 16 |
+
- `MODEL_ROOT`/`CONFIGS_ROOT` must be absolute before importing the vendored
|
| 17 |
+
module (it reads them at import time). Weights arrive preloaded at build
|
| 18 |
+
time (`preload_from_hub`, same HF cache `snapshot_download` reads) and are
|
| 19 |
+
symlinked into `ckpt/`; prompt embeddings (`infer_text_embedding_radar.pt`,
|
| 20 |
+
340KB) ship in git because they are absent from the HuggingFace repo.
|
| 21 |
+
"""
|
| 22 |
+
|
| 23 |
+
import os
|
| 24 |
+
import re
|
| 25 |
+
import shutil
|
| 26 |
+
import tempfile
|
| 27 |
+
|
| 28 |
+
ROOT = os.path.dirname(os.path.abspath(__file__))
|
| 29 |
+
CKPT_DIR = os.path.join(ROOT, "ckpt")
|
| 30 |
+
os.environ.setdefault("MODEL_ROOT", CKPT_DIR)
|
| 31 |
+
os.environ.setdefault("CONFIGS_ROOT", CKPT_DIR)
|
| 32 |
+
|
| 33 |
+
import sys
|
| 34 |
+
sys.path.insert(0, os.path.join(ROOT, "RADAR_inference")) # vendored absolute imports (dynamic_network_architectures) resolve from here
|
| 35 |
+
import pandas as pd # noqa: E402
|
| 36 |
+
import spaces # noqa: E402
|
| 37 |
+
import gradio as gr # noqa: E402
|
| 38 |
+
from huggingface_hub import snapshot_download # noqa: E402
|
| 39 |
+
from inference_demo import initialize, evaluate # noqa: E402
|
| 40 |
+
|
| 41 |
+
REPO_ID = "radar-generalist/RADAR"
|
| 42 |
+
CHECKPOINT_NAME = "checkpoint_radar_pretrain.pth"
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
def _ensure_ckpts() -> None:
|
| 46 |
+
snap = snapshot_download(
|
| 47 |
+
REPO_ID,
|
| 48 |
+
allow_patterns=[CHECKPOINT_NAME, "bert-base-chinese/*"],
|
| 49 |
+
)
|
| 50 |
+
targets = [CHECKPOINT_NAME, "bert-base-chinese"]
|
| 51 |
+
for name in targets:
|
| 52 |
+
src = os.path.join(snap, name)
|
| 53 |
+
dst = os.path.join(CKPT_DIR, name)
|
| 54 |
+
if not os.path.exists(src):
|
| 55 |
+
raise RuntimeError(f"{name} absent from {REPO_ID} snapshot {snap}")
|
| 56 |
+
if os.path.lexists(dst):
|
| 57 |
+
continue
|
| 58 |
+
os.symlink(src, dst)
|
| 59 |
+
ckpt_path = os.path.join(CKPT_DIR, CHECKPOINT_NAME)
|
| 60 |
+
if not os.path.exists(ckpt_path):
|
| 61 |
+
raise RuntimeError(
|
| 62 |
+
f"{CHECKPOINT_NAME} missing after download; "
|
| 63 |
+
"check Space logs/network and re-run."
|
| 64 |
+
)
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
_ensure_ckpts()
|
| 68 |
+
pad_func, model = initialize() # module level; .cuda() here is intentional (ZeroGPU pattern)
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
def _english_name(column: str) -> str:
|
| 72 |
+
m = re.search(r"\((.+)\)\s*$", column)
|
| 73 |
+
return m.group(1) if m else column
|
| 74 |
+
|
| 75 |
+
|
| 76 |
+
@spaces.GPU(duration=180)
|
| 77 |
+
def diagnose(ct_file):
|
| 78 |
+
"""Run RADAR on one uploaded contrast-enhanced abdominal CT volume."""
|
| 79 |
+
path = ct_file if isinstance(ct_file, str) else ct_file.name
|
| 80 |
+
if not path.endswith((".nii", ".nii.gz")):
|
| 81 |
+
raise gr.Error("upload a .nii or .nii.gz CT volume")
|
| 82 |
+
try:
|
| 83 |
+
import nibabel as nib
|
| 84 |
+
|
| 85 |
+
nib.load(path)
|
| 86 |
+
except Exception:
|
| 87 |
+
raise gr.Error("upload a .nii or .nii.gz CT volume")
|
| 88 |
+
|
| 89 |
+
tmpdir = tempfile.mkdtemp(prefix="radar_case_")
|
| 90 |
+
outdir = tempfile.mkdtemp(prefix="radar_out_")
|
| 91 |
+
try:
|
| 92 |
+
fname = os.path.basename(path)
|
| 93 |
+
if not fname.endswith((".nii", ".nii.gz")):
|
| 94 |
+
fname += ".nii.gz"
|
| 95 |
+
shutil.copy(path, os.path.join(tmpdir, fname))
|
| 96 |
+
try:
|
| 97 |
+
evaluate(pad_func, model, tmpdir, outdir, "space")
|
| 98 |
+
except (OSError, ValueError, RuntimeError) as exc:
|
| 99 |
+
raise gr.Error(f"could not process this volume: {exc}")
|
| 100 |
+
csv_path = os.path.join(outdir, "RADAR_infer_results_space.csv")
|
| 101 |
+
df = pd.read_csv(csv_path, encoding="utf-8-sig")
|
| 102 |
+
if df.empty:
|
| 103 |
+
raise gr.Error("model skipped this volume (check dimensions/spacing)")
|
| 104 |
+
row = df.iloc[0]
|
| 105 |
+
scores = {
|
| 106 |
+
_english_name(col): float(row[col])
|
| 107 |
+
for col in df.columns[1:]
|
| 108 |
+
if str(row[col]).strip() != ""
|
| 109 |
+
}
|
| 110 |
+
ranked = sorted(scores.items(), key=lambda kv: kv[1], reverse=True)
|
| 111 |
+
table = pd.DataFrame(ranked, columns=["Finding", "Score"])
|
| 112 |
+
return scores, table
|
| 113 |
+
finally:
|
| 114 |
+
shutil.rmtree(tmpdir, ignore_errors=True)
|
| 115 |
+
shutil.rmtree(outdir, ignore_errors=True)
|
| 116 |
+
|
| 117 |
+
|
| 118 |
+
demo = gr.Interface(
|
| 119 |
+
fn=diagnose,
|
| 120 |
+
inputs=gr.File(
|
| 121 |
+
label="Abdominal CT (.nii.gz)", file_types=[".nii", ".nii.gz"]
|
| 122 |
+
),
|
| 123 |
+
outputs=[gr.Label(label="Top findings", num_top_classes=10), gr.Dataframe(label="All finding scores")],
|
| 124 |
+
title="RADAR Abdominal CT",
|
| 125 |
+
description=(
|
| 126 |
+
"Expert-level generalist AI for contrast-enhanced abdominal CT "
|
| 127 |
+
"(18 structures, 146 findings). Non-commercial research demo "
|
| 128 |
+
"(CC BY-NC-SA 4.0); assistance tool, not a diagnosis."
|
| 129 |
+
),
|
| 130 |
+
)
|
| 131 |
+
|
| 132 |
+
if __name__ == "__main__":
|
| 133 |
+
demo.launch()
|
ckpt/infer_text_embedding_radar.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:0ac960b0300c2ba10526e3bf500cd0c879fe5d9d15b9c49c17c34f0fc40e6cbe
|
| 3 |
+
size 346348
|
requirements.txt
ADDED
|
@@ -0,0 +1,45 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
contexttimer
|
| 2 |
+
decord
|
| 3 |
+
diffusers<=0.16.0
|
| 4 |
+
einops>=0.4.1
|
| 5 |
+
fairscale==0.4.4
|
| 6 |
+
ftfy
|
| 7 |
+
iopath
|
| 8 |
+
ipython
|
| 9 |
+
omegaconf
|
| 10 |
+
# opencv 4.5.5.64 ships no cp312 files: keep upstream pin below 3.12 only.
|
| 11 |
+
opencv-python-headless==4.5.5.64; python_version < '3.12'
|
| 12 |
+
opencv-python-headless; python_version >= '3.12'
|
| 13 |
+
opendatasets
|
| 14 |
+
packaging
|
| 15 |
+
pandas
|
| 16 |
+
plotly
|
| 17 |
+
pre-commit
|
| 18 |
+
pycocoevalcap
|
| 19 |
+
pycocotools
|
| 20 |
+
python-magic
|
| 21 |
+
scikit-image
|
| 22 |
+
sentencepiece
|
| 23 |
+
spacy
|
| 24 |
+
streamlit
|
| 25 |
+
timm==0.4.12
|
| 26 |
+
torch>=1.10.0
|
| 27 |
+
torchvision
|
| 28 |
+
tqdm
|
| 29 |
+
# transformers 4.25 pulls tokenizers 0.13.3 (no cp312 wheel, needs Rust): keep it below 3.12 only.
|
| 30 |
+
transformers==4.25; python_version < '3.12'
|
| 31 |
+
transformers==4.48.3; python_version >= '3.12'
|
| 32 |
+
webdataset
|
| 33 |
+
wheel
|
| 34 |
+
h5py
|
| 35 |
+
monai
|
| 36 |
+
batchgenerators
|
| 37 |
+
SimpleITK
|
| 38 |
+
nibabel
|
| 39 |
+
matplotlib
|
| 40 |
+
nltk
|
| 41 |
+
gradio>=4
|
| 42 |
+
huggingface_hub
|
| 43 |
+
nibabel
|
| 44 |
+
monai
|
| 45 |
+
torch>=2.8,<2.14
|