radames/Text2Human-API
1
1import torch2import torch.nn as nn3import torch.utils.checkpoint as cp4from mmcv.cnn import (UPSAMPLE_LAYERS, ConvModule, build_activation_layer,5 build_norm_layer, build_upsample_layer, constant_init,6 kaiming_init)7from mmcv.runner import load_checkpoint8from mmcv.utils.parrots_wrapper import _BatchNorm9from mmseg.utils import get_root_logger10 11 12class UpConvBlock(nn.Module):13 """Upsample convolution block in decoder for UNet.14 15 This upsample convolution block consists of one upsample module16 followed by one convolution block. The upsample module expands the17 high-level low-resolution feature map and the convolution block fuses18 the upsampled high-level low-resolution feature map and the low-level19 high-resolution feature map from encoder.20 21 Args:22 conv_block (nn.Sequential): Sequential of convolutional layers.23 in_channels (int): Number of input channels of the high-level24 skip_channels (int): Number of input channels of the low-level25 high-resolution feature map from encoder.26 out_channels (int): Number of output channels.27 num_convs (int): Number of convolutional layers in the conv_block.28 Default: 2.29 stride (int): Stride of convolutional layer in conv_block. Default: 1.30 dilation (int): Dilation rate of convolutional layer in conv_block.31 Default: 1.32 with_cp (bool): Use checkpoint or not. Using checkpoint will save some33 memory while slowing down the training speed. Default: False.34 conv_cfg (dict | None): Config dict for convolution layer.35 Default: None.36 norm_cfg (dict | None): Config dict for normalization layer.37 Default: dict(type='BN').38 act_cfg (dict | None): Config dict for activation layer in ConvModule.39 Default: dict(type='ReLU').40 upsample_cfg (dict): The upsample config of the upsample module in41 decoder. Default: dict(type='InterpConv'). If the size of42 high-level feature map is the same as that of skip feature map43 (low-level feature map from encoder), it does not need upsample the44 high-level feature map and the upsample_cfg is None.45 dcn (bool): Use deformable convoluton in convolutional layer or not.46 Default: None.47 plugins (dict): plugins for convolutional layers. Default: None.48 """49 50 def __init__(self,51 conv_block,52 in_channels,53 skip_channels,54 out_channels,55 num_convs=2,56 stride=1,57 dilation=1,58 with_cp=False,59 conv_cfg=None,60 norm_cfg=dict(type='BN'),61 act_cfg=dict(type='ReLU'),62 upsample_cfg=dict(type='InterpConv'),63 dcn=None,64 plugins=None):65 super(UpConvBlock, self).__init__()66 assert dcn is None, 'Not implemented yet.'67 assert plugins is None, 'Not implemented yet.'68 69 self.conv_block = conv_block(70 in_channels=2 * skip_channels,71 out_channels=out_channels,72 num_convs=num_convs,73 stride=stride,74 dilation=dilation,75 with_cp=with_cp,76 conv_cfg=conv_cfg,77 norm_cfg=norm_cfg,78 act_cfg=act_cfg,79 dcn=None,80 plugins=None)81 if upsample_cfg is not None:82 self.upsample = build_upsample_layer(83 cfg=upsample_cfg,84 in_channels=in_channels,85 out_channels=skip_channels,86 with_cp=with_cp,87 norm_cfg=norm_cfg,88 act_cfg=act_cfg)89 else:90 self.upsample = ConvModule(91 in_channels,92 skip_channels,93 kernel_size=1,94 stride=1,95 padding=0,96 conv_cfg=conv_cfg,97 norm_cfg=norm_cfg,98 act_cfg=act_cfg)99 100 def forward(self, skip, x):101 """Forward function."""102 103 x = self.upsample(x)104 out = torch.cat([skip, x], dim=1)105 out = self.conv_block(out)106 107 return out108 109 110class BasicConvBlock(nn.Module):111 """Basic convolutional block for UNet.112 113 This module consists of several plain convolutional layers.114 115 Args:116 in_channels (int): Number of input channels.117 out_channels (int): Number of output channels.118 num_convs (int): Number of convolutional layers. Default: 2.119 stride (int): Whether use stride convolution to downsample120 the input feature map. If stride=2, it only uses stride convolution121 in the first convolutional layer to downsample the input feature122 map. Options are 1 or 2. Default: 1.123 dilation (int): Whether use dilated convolution to expand the124 receptive field. Set dilation rate of each convolutional layer and125 the dilation rate of the first convolutional layer is always 1.126 Default: 1.127 with_cp (bool): Use checkpoint or not. Using checkpoint will save some128 memory while slowing down the training speed. Default: False.129 conv_cfg (dict | None): Config dict for convolution layer.130 Default: None.131 norm_cfg (dict | None): Config dict for normalization layer.132 Default: dict(type='BN').133 act_cfg (dict | None): Config dict for activation layer in ConvModule.134 Default: dict(type='ReLU').135 dcn (bool): Use deformable convoluton in convolutional layer or not.136 Default: None.137 plugins (dict): plugins for convolutional layers. Default: None.138 """139 140 def __init__(self,141 in_channels,142 out_channels,143 num_convs=2,144 stride=1,145 dilation=1,146 with_cp=False,147 conv_cfg=None,148 norm_cfg=dict(type='BN'),149 act_cfg=dict(type='ReLU'),150 dcn=None,151 plugins=None):152 super(BasicConvBlock, self).__init__()153 assert dcn is None, 'Not implemented yet.'154 assert plugins is None, 'Not implemented yet.'155 156 self.with_cp = with_cp157 convs = []158 for i in range(num_convs):159 convs.append(160 ConvModule(161 in_channels=in_channels if i == 0 else out_channels,162 out_channels=out_channels,163 kernel_size=3,164 stride=stride if i == 0 else 1,165 dilation=1 if i == 0 else dilation,166 padding=1 if i == 0 else dilation,167 conv_cfg=conv_cfg,168 norm_cfg=norm_cfg,169 act_cfg=act_cfg))170 171 self.convs = nn.Sequential(*convs)172 173 def forward(self, x):174 """Forward function."""175 176 if self.with_cp and x.requires_grad:177 out = cp.checkpoint(self.convs, x)178 else:179 out = self.convs(x)180 return out181 182 183class DeconvModule(nn.Module):184 """Deconvolution upsample module in decoder for UNet (2X upsample).185 186 This module uses deconvolution to upsample feature map in the decoder187 of UNet.188 189 Args:190 in_channels (int): Number of input channels.191 out_channels (int): Number of output channels.192 with_cp (bool): Use checkpoint or not. Using checkpoint will save some193 memory while slowing down the training speed. Default: False.194 norm_cfg (dict | None): Config dict for normalization layer.195 Default: dict(type='BN').196 act_cfg (dict | None): Config dict for activation layer in ConvModule.197 Default: dict(type='ReLU').198 kernel_size (int): Kernel size of the convolutional layer. Default: 4.199 """200 201 def __init__(self,202 in_channels,203 out_channels,204 with_cp=False,205 norm_cfg=dict(type='BN'),206 act_cfg=dict(type='ReLU'),207 *,208 kernel_size=4,209 scale_factor=2):210 super(DeconvModule, self).__init__()211 212 assert (kernel_size - scale_factor >= 0) and\213 (kernel_size - scale_factor) % 2 == 0,\214 f'kernel_size should be greater than or equal to scale_factor '\215 f'and (kernel_size - scale_factor) should be even numbers, '\216 f'while the kernel size is {kernel_size} and scale_factor is '\217 f'{scale_factor}.'218 219 stride = scale_factor220 padding = (kernel_size - scale_factor) // 2221 self.with_cp = with_cp222 deconv = nn.ConvTranspose2d(223 in_channels,224 out_channels,225 kernel_size=kernel_size,226 stride=stride,227 padding=padding)228 229 norm_name, norm = build_norm_layer(norm_cfg, out_channels)230 activate = build_activation_layer(act_cfg)231 self.deconv_upsamping = nn.Sequential(deconv, norm, activate)232 233 def forward(self, x):234 """Forward function."""235 236 if self.with_cp and x.requires_grad:237 out = cp.checkpoint(self.deconv_upsamping, x)238 else:239 out = self.deconv_upsamping(x)240 return out241 242 243@UPSAMPLE_LAYERS.register_module()244class InterpConv(nn.Module):245 """Interpolation upsample module in decoder for UNet.246 247 This module uses interpolation to upsample feature map in the decoder248 of UNet. It consists of one interpolation upsample layer and one249 convolutional layer. It can be one interpolation upsample layer followed250 by one convolutional layer (conv_first=False) or one convolutional layer251 followed by one interpolation upsample layer (conv_first=True).252 253 Args:254 in_channels (int): Number of input channels.255 out_channels (int): Number of output channels.256 with_cp (bool): Use checkpoint or not. Using checkpoint will save some257 memory while slowing down the training speed. Default: False.258 norm_cfg (dict | None): Config dict for normalization layer.259 Default: dict(type='BN').260 act_cfg (dict | None): Config dict for activation layer in ConvModule.261 Default: dict(type='ReLU').262 conv_cfg (dict | None): Config dict for convolution layer.263 Default: None.264 conv_first (bool): Whether convolutional layer or interpolation265 upsample layer first. Default: False. It means interpolation266 upsample layer followed by one convolutional layer.267 kernel_size (int): Kernel size of the convolutional layer. Default: 1.268 stride (int): Stride of the convolutional layer. Default: 1.269 padding (int): Padding of the convolutional layer. Default: 1.270 upsampe_cfg (dict): Interpolation config of the upsample layer.271 Default: dict(272 scale_factor=2, mode='bilinear', align_corners=False).273 """274 275 def __init__(self,276 in_channels,277 out_channels,278 with_cp=False,279 norm_cfg=dict(type='BN'),280 act_cfg=dict(type='ReLU'),281 *,282 conv_cfg=None,283 conv_first=False,284 kernel_size=1,285 stride=1,286 padding=0,287 upsampe_cfg=dict(288 scale_factor=2, mode='bilinear', align_corners=False)):289 super(InterpConv, self).__init__()290 291 self.with_cp = with_cp292 conv = ConvModule(293 in_channels,294 out_channels,295 kernel_size=kernel_size,296 stride=stride,297 padding=padding,298 conv_cfg=conv_cfg,299 norm_cfg=norm_cfg,300 act_cfg=act_cfg)301 upsample = nn.Upsample(**upsampe_cfg)302 if conv_first:303 self.interp_upsample = nn.Sequential(conv, upsample)304 else:305 self.interp_upsample = nn.Sequential(upsample, conv)306 307 def forward(self, x):308 """Forward function."""309 310 if self.with_cp and x.requires_grad:311 out = cp.checkpoint(self.interp_upsample, x)312 else:313 out = self.interp_upsample(x)314 return out315 316 317class UNet(nn.Module):318 """UNet backbone.319 U-Net: Convolutional Networks for Biomedical Image Segmentation.320 https://arxiv.org/pdf/1505.04597.pdf321 322 Args:323 in_channels (int): Number of input image channels. Default" 3.324 base_channels (int): Number of base channels of each stage.325 The output channels of the first stage. Default: 64.326 num_stages (int): Number of stages in encoder, normally 5. Default: 5.327 strides (Sequence[int 1 | 2]): Strides of each stage in encoder.328 len(strides) is equal to num_stages. Normally the stride of the329 first stage in encoder is 1. If strides[i]=2, it uses stride330 convolution to downsample in the correspondence encoder stage.331 Default: (1, 1, 1, 1, 1).332 enc_num_convs (Sequence[int]): Number of convolutional layers in the333 convolution block of the correspondence encoder stage.334 Default: (2, 2, 2, 2, 2).335 dec_num_convs (Sequence[int]): Number of convolutional layers in the336 convolution block of the correspondence decoder stage.337 Default: (2, 2, 2, 2).338 downsamples (Sequence[int]): Whether use MaxPool to downsample the339 feature map after the first stage of encoder340 (stages: [1, num_stages)). If the correspondence encoder stage use341 stride convolution (strides[i]=2), it will never use MaxPool to342 downsample, even downsamples[i-1]=True.343 Default: (True, True, True, True).344 enc_dilations (Sequence[int]): Dilation rate of each stage in encoder.345 Default: (1, 1, 1, 1, 1).346 dec_dilations (Sequence[int]): Dilation rate of each stage in decoder.347 Default: (1, 1, 1, 1).348 with_cp (bool): Use checkpoint or not. Using checkpoint will save some349 memory while slowing down the training speed. Default: False.350 conv_cfg (dict | None): Config dict for convolution layer.351 Default: None.352 norm_cfg (dict | None): Config dict for normalization layer.353 Default: dict(type='BN').354 act_cfg (dict | None): Config dict for activation layer in ConvModule.355 Default: dict(type='ReLU').356 upsample_cfg (dict): The upsample config of the upsample module in357 decoder. Default: dict(type='InterpConv').358 norm_eval (bool): Whether to set norm layers to eval mode, namely,359 freeze running stats (mean and var). Note: Effect on Batch Norm360 and its variants only. Default: False.361 dcn (bool): Use deformable convolution in convolutional layer or not.362 Default: None.363 plugins (dict): plugins for convolutional layers. Default: None.364 365 Notice:366 The input image size should be devisible by the whole downsample rate367 of the encoder. More detail of the whole downsample rate can be found368 in UNet._check_input_devisible.369 370 """371 372 def __init__(self,373 in_channels=3,374 base_channels=64,375 num_stages=5,376 strides=(1, 1, 1, 1, 1),377 enc_num_convs=(2, 2, 2, 2, 2),378 dec_num_convs=(2, 2, 2, 2),379 downsamples=(True, True, True, True),380 enc_dilations=(1, 1, 1, 1, 1),381 dec_dilations=(1, 1, 1, 1),382 with_cp=False,383 conv_cfg=None,384 norm_cfg=dict(type='BN'),385 act_cfg=dict(type='ReLU'),386 upsample_cfg=dict(type='InterpConv'),387 norm_eval=False,388 dcn=None,389 plugins=None):390 super(UNet, self).__init__()391 assert dcn is None, 'Not implemented yet.'392 assert plugins is None, 'Not implemented yet.'393 assert len(strides) == num_stages, \394 'The length of strides should be equal to num_stages, '\395 f'while the strides is {strides}, the length of '\396 f'strides is {len(strides)}, and the num_stages is '\397 f'{num_stages}.'398 assert len(enc_num_convs) == num_stages, \399 'The length of enc_num_convs should be equal to num_stages, '\400 f'while the enc_num_convs is {enc_num_convs}, the length of '\401 f'enc_num_convs is {len(enc_num_convs)}, and the num_stages is '\402 f'{num_stages}.'403 assert len(dec_num_convs) == (num_stages-1), \404 'The length of dec_num_convs should be equal to (num_stages-1), '\405 f'while the dec_num_convs is {dec_num_convs}, the length of '\406 f'dec_num_convs is {len(dec_num_convs)}, and the num_stages is '\407 f'{num_stages}.'408 assert len(downsamples) == (num_stages-1), \409 'The length of downsamples should be equal to (num_stages-1), '\410 f'while the downsamples is {downsamples}, the length of '\411 f'downsamples is {len(downsamples)}, and the num_stages is '\412 f'{num_stages}.'413 assert len(enc_dilations) == num_stages, \414 'The length of enc_dilations should be equal to num_stages, '\415 f'while the enc_dilations is {enc_dilations}, the length of '\416 f'enc_dilations is {len(enc_dilations)}, and the num_stages is '\417 f'{num_stages}.'418 assert len(dec_dilations) == (num_stages-1), \419 'The length of dec_dilations should be equal to (num_stages-1), '\420 f'while the dec_dilations is {dec_dilations}, the length of '\421 f'dec_dilations is {len(dec_dilations)}, and the num_stages is '\422 f'{num_stages}.'423 self.num_stages = num_stages424 self.strides = strides425 self.downsamples = downsamples426 self.norm_eval = norm_eval427 428 self.encoder = nn.ModuleList()429 self.decoder = nn.ModuleList()430 431 for i in range(num_stages):432 enc_conv_block = []433 if i != 0:434 if strides[i] == 1 and downsamples[i - 1]:435 enc_conv_block.append(nn.MaxPool2d(kernel_size=2))436 upsample = (strides[i] != 1 or downsamples[i - 1])437 self.decoder.append(438 UpConvBlock(439 conv_block=BasicConvBlock,440 in_channels=base_channels * 2**i,441 skip_channels=base_channels * 2**(i - 1),442 out_channels=base_channels * 2**(i - 1),443 num_convs=dec_num_convs[i - 1],444 stride=1,445 dilation=dec_dilations[i - 1],446 with_cp=with_cp,447 conv_cfg=conv_cfg,448 norm_cfg=norm_cfg,449 act_cfg=act_cfg,450 upsample_cfg=upsample_cfg if upsample else None,451 dcn=None,452 plugins=None))453 454 enc_conv_block.append(455 BasicConvBlock(456 in_channels=in_channels,457 out_channels=base_channels * 2**i,458 num_convs=enc_num_convs[i],459 stride=strides[i],460 dilation=enc_dilations[i],461 with_cp=with_cp,462 conv_cfg=conv_cfg,463 norm_cfg=norm_cfg,464 act_cfg=act_cfg,465 dcn=None,466 plugins=None))467 self.encoder.append((nn.Sequential(*enc_conv_block)))468 in_channels = base_channels * 2**i469 470 def forward(self, x):471 enc_outs = []472 473 for enc in self.encoder:474 x = enc(x)475 enc_outs.append(x)476 dec_outs = [x]477 for i in reversed(range(len(self.decoder))):478 x = self.decoder[i](enc_outs[i], x)479 dec_outs.append(x)480 481 return dec_outs482 483 def init_weights(self, pretrained=None):484 """Initialize the weights in backbone.485 486 Args:487 pretrained (str, optional): Path to pre-trained weights.488 Defaults to None.489 """490 if isinstance(pretrained, str):491 logger = get_root_logger()492 load_checkpoint(self, pretrained, strict=False, logger=logger)493 elif pretrained is None:494 for m in self.modules():495 if isinstance(m, nn.Conv2d):496 kaiming_init(m)497 elif isinstance(m, (_BatchNorm, nn.GroupNorm)):498 constant_init(m, 1)499 else:500 raise TypeError('pretrained must be a str or None')501 502 503class ShapeUNet(nn.Module):504 """ShapeUNet backbone with small modifications.505 U-Net: Convolutional Networks for Biomedical Image Segmentation.506 https://arxiv.org/pdf/1505.04597.pdf507 508 Args:509 in_channels (int): Number of input image channels. Default" 3.510 base_channels (int): Number of base channels of each stage.511 The output channels of the first stage. Default: 64.512 num_stages (int): Number of stages in encoder, normally 5. Default: 5.513 strides (Sequence[int 1 | 2]): Strides of each stage in encoder.514 len(strides) is equal to num_stages. Normally the stride of the515 first stage in encoder is 1. If strides[i]=2, it uses stride516 convolution to downsample in the correspondance encoder stage.517 Default: (1, 1, 1, 1, 1).518 enc_num_convs (Sequence[int]): Number of convolutional layers in the519 convolution block of the correspondance encoder stage.520 Default: (2, 2, 2, 2, 2).521 dec_num_convs (Sequence[int]): Number of convolutional layers in the522 convolution block of the correspondance decoder stage.523 Default: (2, 2, 2, 2).524 downsamples (Sequence[int]): Whether use MaxPool to downsample the525 feature map after the first stage of encoder526 (stages: [1, num_stages)). If the correspondance encoder stage use527 stride convolution (strides[i]=2), it will never use MaxPool to528 downsample, even downsamples[i-1]=True.529 Default: (True, True, True, True).530 enc_dilations (Sequence[int]): Dilation rate of each stage in encoder.531 Default: (1, 1, 1, 1, 1).532 dec_dilations (Sequence[int]): Dilation rate of each stage in decoder.533 Default: (1, 1, 1, 1).534 with_cp (bool): Use checkpoint or not. Using checkpoint will save some535 memory while slowing down the training speed. Default: False.536 conv_cfg (dict | None): Config dict for convolution layer.537 Default: None.538 norm_cfg (dict | None): Config dict for normalization layer.539 Default: dict(type='BN').540 act_cfg (dict | None): Config dict for activation layer in ConvModule.541 Default: dict(type='ReLU').542 upsample_cfg (dict): The upsample config of the upsample module in543 decoder. Default: dict(type='InterpConv').544 norm_eval (bool): Whether to set norm layers to eval mode, namely,545 freeze running stats (mean and var). Note: Effect on Batch Norm546 and its variants only. Default: False.547 dcn (bool): Use deformable convoluton in convolutional layer or not.548 Default: None.549 plugins (dict): plugins for convolutional layers. Default: None.550 551 Notice:552 The input image size should be devisible by the whole downsample rate553 of the encoder. More detail of the whole downsample rate can be found554 in UNet._check_input_devisible.555 556 """557 558 def __init__(self,559 in_channels=3,560 base_channels=64,561 num_stages=5,562 attr_embedding=128,563 strides=(1, 1, 1, 1, 1),564 enc_num_convs=(2, 2, 2, 2, 2),565 dec_num_convs=(2, 2, 2, 2),566 downsamples=(True, True, True, True),567 enc_dilations=(1, 1, 1, 1, 1),568 dec_dilations=(1, 1, 1, 1),569 with_cp=False,570 conv_cfg=None,571 norm_cfg=dict(type='BN'),572 act_cfg=dict(type='ReLU'),573 upsample_cfg=dict(type='InterpConv'),574 norm_eval=False,575 dcn=None,576 plugins=None):577 super(ShapeUNet, self).__init__()578 assert dcn is None, 'Not implemented yet.'579 assert plugins is None, 'Not implemented yet.'580 assert len(strides) == num_stages, \581 'The length of strides should be equal to num_stages, '\582 f'while the strides is {strides}, the length of '\583 f'strides is {len(strides)}, and the num_stages is '\584 f'{num_stages}.'585 assert len(enc_num_convs) == num_stages, \586 'The length of enc_num_convs should be equal to num_stages, '\587 f'while the enc_num_convs is {enc_num_convs}, the length of '\588 f'enc_num_convs is {len(enc_num_convs)}, and the num_stages is '\589 f'{num_stages}.'590 assert len(dec_num_convs) == (num_stages-1), \591 'The length of dec_num_convs should be equal to (num_stages-1), '\592 f'while the dec_num_convs is {dec_num_convs}, the length of '\593 f'dec_num_convs is {len(dec_num_convs)}, and the num_stages is '\594 f'{num_stages}.'595 assert len(downsamples) == (num_stages-1), \596 'The length of downsamples should be equal to (num_stages-1), '\597 f'while the downsamples is {downsamples}, the length of '\598 f'downsamples is {len(downsamples)}, and the num_stages is '\599 f'{num_stages}.'600 assert len(enc_dilations) == num_stages, \601 'The length of enc_dilations should be equal to num_stages, '\602 f'while the enc_dilations is {enc_dilations}, the length of '\603 f'enc_dilations is {len(enc_dilations)}, and the num_stages is '\604 f'{num_stages}.'605 assert len(dec_dilations) == (num_stages-1), \606 'The length of dec_dilations should be equal to (num_stages-1), '\607 f'while the dec_dilations is {dec_dilations}, the length of '\608 f'dec_dilations is {len(dec_dilations)}, and the num_stages is '\609 f'{num_stages}.'610 self.num_stages = num_stages611 self.strides = strides612 self.downsamples = downsamples613 self.norm_eval = norm_eval614 615 self.encoder = nn.ModuleList()616 self.decoder = nn.ModuleList()617 618 for i in range(num_stages):619 enc_conv_block = []620 if i != 0:621 if strides[i] == 1 and downsamples[i - 1]:622 enc_conv_block.append(nn.MaxPool2d(kernel_size=2))623 upsample = (strides[i] != 1 or downsamples[i - 1])624 self.decoder.append(625 UpConvBlock(626 conv_block=BasicConvBlock,627 in_channels=base_channels * 2**i,628 skip_channels=base_channels * 2**(i - 1),629 out_channels=base_channels * 2**(i - 1),630 num_convs=dec_num_convs[i - 1],631 stride=1,632 dilation=dec_dilations[i - 1],633 with_cp=with_cp,634 conv_cfg=conv_cfg,635 norm_cfg=norm_cfg,636 act_cfg=act_cfg,637 upsample_cfg=upsample_cfg if upsample else None,638 dcn=None,639 plugins=None))640 641 enc_conv_block.append(642 BasicConvBlock(643 in_channels=in_channels + attr_embedding,644 out_channels=base_channels * 2**i,645 num_convs=enc_num_convs[i],646 stride=strides[i],647 dilation=enc_dilations[i],648 with_cp=with_cp,649 conv_cfg=conv_cfg,650 norm_cfg=norm_cfg,651 act_cfg=act_cfg,652 dcn=None,653 plugins=None))654 self.encoder.append((nn.Sequential(*enc_conv_block)))655 in_channels = base_channels * 2**i656 657 def forward(self, x, attr_embedding):658 enc_outs = []659 Be, Ce = attr_embedding.size()660 for enc in self.encoder:661 _, _, H, W = x.size()662 x = enc(663 torch.cat([664 x,665 attr_embedding.view(Be, Ce, 1, 1).expand((Be, Ce, H, W))666 ],667 dim=1))668 enc_outs.append(x)669 dec_outs = [x]670 for i in reversed(range(len(self.decoder))):671 x = self.decoder[i](enc_outs[i], x)672 dec_outs.append(x)673 674 return dec_outs675 676 def init_weights(self, pretrained=None):677 """Initialize the weights in backbone.678 679 Args:680 pretrained (str, optional): Path to pre-trained weights.681 Defaults to None.682 """683 if isinstance(pretrained, str):684 logger = get_root_logger()685 load_checkpoint(self, pretrained, strict=False, logger=logger)686 elif pretrained is None:687 for m in self.modules():688 if isinstance(m, nn.Conv2d):689 kaiming_init(m)690 elif isinstance(m, (_BatchNorm, nn.GroupNorm)):691 constant_init(m, 1)692 else:693 raise TypeError('pretrained must be a str or None')694 