Team Ai
Apppublic

radames/Text2Human-API

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
unet_arch.py694 linesDownload Raw Back to archs
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