isaac commited on
Commit
af2b273
·
0 Parent(s):

RADAR ZeroGPU Space

Browse files
Files changed (28) hide show
  1. .gitattributes +1 -0
  2. .gitignore +17 -0
  3. RADAR_inference/dynamic_network_architectures/__init__.py +0 -0
  4. RADAR_inference/dynamic_network_architectures/architectures/__init__.py +0 -0
  5. RADAR_inference/dynamic_network_architectures/architectures/resnet.py +236 -0
  6. RADAR_inference/dynamic_network_architectures/architectures/resnet_vl.py +225 -0
  7. RADAR_inference/dynamic_network_architectures/architectures/unet.py +218 -0
  8. RADAR_inference/dynamic_network_architectures/architectures/unet_lightdecoder.py +218 -0
  9. RADAR_inference/dynamic_network_architectures/architectures/vgg.py +85 -0
  10. RADAR_inference/dynamic_network_architectures/building_blocks/__init__.py +0 -0
  11. RADAR_inference/dynamic_network_architectures/building_blocks/helper.py +242 -0
  12. RADAR_inference/dynamic_network_architectures/building_blocks/plain_conv_encoder.py +105 -0
  13. RADAR_inference/dynamic_network_architectures/building_blocks/regularization.py +86 -0
  14. RADAR_inference/dynamic_network_architectures/building_blocks/residual.py +371 -0
  15. RADAR_inference/dynamic_network_architectures/building_blocks/residual_encoders.py +172 -0
  16. RADAR_inference/dynamic_network_architectures/building_blocks/simple_conv_blocks.py +167 -0
  17. RADAR_inference/dynamic_network_architectures/building_blocks/unet_decoder.py +154 -0
  18. RADAR_inference/dynamic_network_architectures/building_blocks/unet_decoder_light.py +154 -0
  19. RADAR_inference/dynamic_network_architectures/building_blocks/unet_residual_decoder.py +155 -0
  20. RADAR_inference/dynamic_network_architectures/initialization/__init__.py +0 -0
  21. RADAR_inference/dynamic_network_architectures/initialization/weight_init.py +34 -0
  22. RADAR_inference/dynamic_network_architectures/med.py +1502 -0
  23. RADAR_inference/dynamic_network_architectures/vision_branch.py +160 -0
  24. RADAR_inference/inference_demo.py +630 -0
  25. README.md +21 -0
  26. app.py +133 -0
  27. ckpt/infer_text_embedding_radar.pt +3 -0
  28. 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