Team Ai
Apppublic

Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
trident_backbone.py221 linesDownload Raw Back to tridentnet
1# Copyright (c) Facebook, Inc. and its affiliates.2import fvcore.nn.weight_init as weight_init3import torch4import torch.nn.functional as F5 6from detectron2.layers import Conv2d, FrozenBatchNorm2d, get_norm7from detectron2.modeling import BACKBONE_REGISTRY, ResNet, ResNetBlockBase8from detectron2.modeling.backbone.resnet import BasicStem, BottleneckBlock, DeformBottleneckBlock9 10from .trident_conv import TridentConv11 12__all__ = ["TridentBottleneckBlock", "make_trident_stage", "build_trident_resnet_backbone"]13 14 15class TridentBottleneckBlock(ResNetBlockBase):16    def __init__(17        self,18        in_channels,19        out_channels,20        *,21        bottleneck_channels,22        stride=1,23        num_groups=1,24        norm="BN",25        stride_in_1x1=False,26        num_branch=3,27        dilations=(1, 2, 3),28        concat_output=False,29        test_branch_idx=-1,30    ):31        """32        Args:33            num_branch (int): the number of branches in TridentNet.34            dilations (tuple): the dilations of multiple branches in TridentNet.35            concat_output (bool): if concatenate outputs of multiple branches in TridentNet.36                Use 'True' for the last trident block.37        """38        super().__init__(in_channels, out_channels, stride)39 40        assert num_branch == len(dilations)41 42        self.num_branch = num_branch43        self.concat_output = concat_output44        self.test_branch_idx = test_branch_idx45 46        if in_channels != out_channels:47            self.shortcut = Conv2d(48                in_channels,49                out_channels,50                kernel_size=1,51                stride=stride,52                bias=False,53                norm=get_norm(norm, out_channels),54            )55        else:56            self.shortcut = None57 58        stride_1x1, stride_3x3 = (stride, 1) if stride_in_1x1 else (1, stride)59 60        self.conv1 = Conv2d(61            in_channels,62            bottleneck_channels,63            kernel_size=1,64            stride=stride_1x1,65            bias=False,66            norm=get_norm(norm, bottleneck_channels),67        )68 69        self.conv2 = TridentConv(70            bottleneck_channels,71            bottleneck_channels,72            kernel_size=3,73            stride=stride_3x3,74            paddings=dilations,75            bias=False,76            groups=num_groups,77            dilations=dilations,78            num_branch=num_branch,79            test_branch_idx=test_branch_idx,80            norm=get_norm(norm, bottleneck_channels),81        )82 83        self.conv3 = Conv2d(84            bottleneck_channels,85            out_channels,86            kernel_size=1,87            bias=False,88            norm=get_norm(norm, out_channels),89        )90 91        for layer in [self.conv1, self.conv2, self.conv3, self.shortcut]:92            if layer is not None:  # shortcut can be None93                weight_init.c2_msra_fill(layer)94 95    def forward(self, x):96        num_branch = self.num_branch if self.training or self.test_branch_idx == -1 else 197        if not isinstance(x, list):98            x = [x] * num_branch99        out = [self.conv1(b) for b in x]100        out = [F.relu_(b) for b in out]101 102        out = self.conv2(out)103        out = [F.relu_(b) for b in out]104 105        out = [self.conv3(b) for b in out]106 107        if self.shortcut is not None:108            shortcut = [self.shortcut(b) for b in x]109        else:110            shortcut = x111 112        out = [out_b + shortcut_b for out_b, shortcut_b in zip(out, shortcut)]113        out = [F.relu_(b) for b in out]114        if self.concat_output:115            out = torch.cat(out)116        return out117 118 119def make_trident_stage(block_class, num_blocks, **kwargs):120    """121    Create a resnet stage by creating many blocks for TridentNet.122    """123    concat_output = [False] * (num_blocks - 1) + [True]124    kwargs["concat_output_per_block"] = concat_output125    return ResNet.make_stage(block_class, num_blocks, **kwargs)126 127 128@BACKBONE_REGISTRY.register()129def build_trident_resnet_backbone(cfg, input_shape):130    """131    Create a ResNet instance from config for TridentNet.132 133    Returns:134        ResNet: a :class:`ResNet` instance.135    """136    # need registration of new blocks/stems?137    norm = cfg.MODEL.RESNETS.NORM138    stem = BasicStem(139        in_channels=input_shape.channels,140        out_channels=cfg.MODEL.RESNETS.STEM_OUT_CHANNELS,141        norm=norm,142    )143    freeze_at = cfg.MODEL.BACKBONE.FREEZE_AT144 145    if freeze_at >= 1:146        for p in stem.parameters():147            p.requires_grad = False148        stem = FrozenBatchNorm2d.convert_frozen_batchnorm(stem)149 150    # fmt: off151    out_features         = cfg.MODEL.RESNETS.OUT_FEATURES152    depth                = cfg.MODEL.RESNETS.DEPTH153    num_groups           = cfg.MODEL.RESNETS.NUM_GROUPS154    width_per_group      = cfg.MODEL.RESNETS.WIDTH_PER_GROUP155    bottleneck_channels  = num_groups * width_per_group156    in_channels          = cfg.MODEL.RESNETS.STEM_OUT_CHANNELS157    out_channels         = cfg.MODEL.RESNETS.RES2_OUT_CHANNELS158    stride_in_1x1        = cfg.MODEL.RESNETS.STRIDE_IN_1X1159    res5_dilation        = cfg.MODEL.RESNETS.RES5_DILATION160    deform_on_per_stage  = cfg.MODEL.RESNETS.DEFORM_ON_PER_STAGE161    deform_modulated     = cfg.MODEL.RESNETS.DEFORM_MODULATED162    deform_num_groups    = cfg.MODEL.RESNETS.DEFORM_NUM_GROUPS163    num_branch           = cfg.MODEL.TRIDENT.NUM_BRANCH164    branch_dilations     = cfg.MODEL.TRIDENT.BRANCH_DILATIONS165    trident_stage        = cfg.MODEL.TRIDENT.TRIDENT_STAGE166    test_branch_idx      = cfg.MODEL.TRIDENT.TEST_BRANCH_IDX167    # fmt: on168    assert res5_dilation in {1, 2}, "res5_dilation cannot be {}.".format(res5_dilation)169 170    num_blocks_per_stage = {50: [3, 4, 6, 3], 101: [3, 4, 23, 3], 152: [3, 8, 36, 3]}[depth]171 172    stages = []173 174    res_stage_idx = {"res2": 2, "res3": 3, "res4": 4, "res5": 5}175    out_stage_idx = [res_stage_idx[f] for f in out_features]176    trident_stage_idx = res_stage_idx[trident_stage]177    max_stage_idx = max(out_stage_idx)178    for idx, stage_idx in enumerate(range(2, max_stage_idx + 1)):179        dilation = res5_dilation if stage_idx == 5 else 1180        first_stride = 1 if idx == 0 or (stage_idx == 5 and dilation == 2) else 2181        stage_kargs = {182            "num_blocks": num_blocks_per_stage[idx],183            "stride_per_block": [first_stride] + [1] * (num_blocks_per_stage[idx] - 1),184            "in_channels": in_channels,185            "bottleneck_channels": bottleneck_channels,186            "out_channels": out_channels,187            "num_groups": num_groups,188            "norm": norm,189            "stride_in_1x1": stride_in_1x1,190            "dilation": dilation,191        }192        if stage_idx == trident_stage_idx:193            assert not deform_on_per_stage[194                idx195            ], "Not support deformable conv in Trident blocks yet."196            stage_kargs["block_class"] = TridentBottleneckBlock197            stage_kargs["num_branch"] = num_branch198            stage_kargs["dilations"] = branch_dilations199            stage_kargs["test_branch_idx"] = test_branch_idx200            stage_kargs.pop("dilation")201        elif deform_on_per_stage[idx]:202            stage_kargs["block_class"] = DeformBottleneckBlock203            stage_kargs["deform_modulated"] = deform_modulated204            stage_kargs["deform_num_groups"] = deform_num_groups205        else:206            stage_kargs["block_class"] = BottleneckBlock207        blocks = (208            make_trident_stage(**stage_kargs)209            if stage_idx == trident_stage_idx210            else ResNet.make_stage(**stage_kargs)211        )212        in_channels = out_channels213        out_channels *= 2214        bottleneck_channels *= 2215 216        if freeze_at >= stage_idx:217            for block in blocks:218                block.freeze()219        stages.append(blocks)220    return ResNet(stem, stages, out_features=out_features)221