Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
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 