xdecoder/Instruct-X-Decoder
163
1# Copyright (c) Facebook, Inc. and its affiliates.2import pickle3import numpy as np4from typing import Any, Dict5import fvcore.nn.weight_init as weight_init6import torch7import torch.nn.functional as F8from torch import nn9 10 11from .backbone import Backbone12from .registry import register_backbone13 14from detectron2.layers import (15 CNNBlockBase,16 Conv2d,17 DeformConv,18 ModulatedDeformConv,19 ShapeSpec,20 get_norm,21)22from detectron2.utils.file_io import PathManager23 24__all__ = [25 "ResNetBlockBase",26 "BasicBlock",27 "BottleneckBlock",28 "DeformBottleneckBlock",29 "BasicStem",30 "ResNet",31 "make_stage",32 "get_resnet_backbone",33]34 35 36class BasicBlock(CNNBlockBase):37 """38 The basic residual block for ResNet-18 and ResNet-34 defined in :paper:`ResNet`,39 with two 3x3 conv layers and a projection shortcut if needed.40 """41 42 def __init__(self, in_channels, out_channels, *, stride=1, norm="BN"):43 """44 Args:45 in_channels (int): Number of input channels.46 out_channels (int): Number of output channels.47 stride (int): Stride for the first conv.48 norm (str or callable): normalization for all conv layers.49 See :func:`layers.get_norm` for supported format.50 """51 super().__init__(in_channels, out_channels, stride)52 53 if in_channels != out_channels:54 self.shortcut = Conv2d(55 in_channels,56 out_channels,57 kernel_size=1,58 stride=stride,59 bias=False,60 norm=get_norm(norm, out_channels),61 )62 else:63 self.shortcut = None64 65 self.conv1 = Conv2d(66 in_channels,67 out_channels,68 kernel_size=3,69 stride=stride,70 padding=1,71 bias=False,72 norm=get_norm(norm, out_channels),73 )74 75 self.conv2 = Conv2d(76 out_channels,77 out_channels,78 kernel_size=3,79 stride=1,80 padding=1,81 bias=False,82 norm=get_norm(norm, out_channels),83 )84 85 for layer in [self.conv1, self.conv2, self.shortcut]:86 if layer is not None: # shortcut can be None87 weight_init.c2_msra_fill(layer)88 89 def forward(self, x):90 out = self.conv1(x)91 out = F.relu_(out)92 out = self.conv2(out)93 94 if self.shortcut is not None:95 shortcut = self.shortcut(x)96 else:97 shortcut = x98 99 out += shortcut100 out = F.relu_(out)101 return out102 103 104class BottleneckBlock(CNNBlockBase):105 """106 The standard bottleneck residual block used by ResNet-50, 101 and 152107 defined in :paper:`ResNet`. It contains 3 conv layers with kernels108 1x1, 3x3, 1x1, and a projection shortcut if needed.109 """110 111 def __init__(112 self,113 in_channels,114 out_channels,115 *,116 bottleneck_channels,117 stride=1,118 num_groups=1,119 norm="BN",120 stride_in_1x1=False,121 dilation=1,122 ):123 """124 Args:125 bottleneck_channels (int): number of output channels for the 3x3126 "bottleneck" conv layers.127 num_groups (int): number of groups for the 3x3 conv layer.128 norm (str or callable): normalization for all conv layers.129 See :func:`layers.get_norm` for supported format.130 stride_in_1x1 (bool): when stride>1, whether to put stride in the131 first 1x1 convolution or the bottleneck 3x3 convolution.132 dilation (int): the dilation rate of the 3x3 conv layer.133 """134 super().__init__(in_channels, out_channels, stride)135 136 if in_channels != out_channels:137 self.shortcut = Conv2d(138 in_channels,139 out_channels,140 kernel_size=1,141 stride=stride,142 bias=False,143 norm=get_norm(norm, out_channels),144 )145 else:146 self.shortcut = None147 148 # The original MSRA ResNet models have stride in the first 1x1 conv149 # The subsequent fb.torch.resnet and Caffe2 ResNe[X]t implementations have150 # stride in the 3x3 conv151 stride_1x1, stride_3x3 = (stride, 1) if stride_in_1x1 else (1, stride)152 153 self.conv1 = Conv2d(154 in_channels,155 bottleneck_channels,156 kernel_size=1,157 stride=stride_1x1,158 bias=False,159 norm=get_norm(norm, bottleneck_channels),160 )161 162 self.conv2 = Conv2d(163 bottleneck_channels,164 bottleneck_channels,165 kernel_size=3,166 stride=stride_3x3,167 padding=1 * dilation,168 bias=False,169 groups=num_groups,170 dilation=dilation,171 norm=get_norm(norm, bottleneck_channels),172 )173 174 self.conv3 = Conv2d(175 bottleneck_channels,176 out_channels,177 kernel_size=1,178 bias=False,179 norm=get_norm(norm, out_channels),180 )181 182 for layer in [self.conv1, self.conv2, self.conv3, self.shortcut]:183 if layer is not None: # shortcut can be None184 weight_init.c2_msra_fill(layer)185 186 # Zero-initialize the last normalization in each residual branch,187 # so that at the beginning, the residual branch starts with zeros,188 # and each residual block behaves like an identity.189 # See Sec 5.1 in "Accurate, Large Minibatch SGD: Training ImageNet in 1 Hour":190 # "For BN layers, the learnable scaling coefficient γ is initialized191 # to be 1, except for each residual block's last BN192 # where γ is initialized to be 0."193 194 # nn.init.constant_(self.conv3.norm.weight, 0)195 # TODO this somehow hurts performance when training GN models from scratch.196 # Add it as an option when we need to use this code to train a backbone.197 198 def forward(self, x):199 out = self.conv1(x)200 out = F.relu_(out)201 202 out = self.conv2(out)203 out = F.relu_(out)204 205 out = self.conv3(out)206 207 if self.shortcut is not None:208 shortcut = self.shortcut(x)209 else:210 shortcut = x211 212 out += shortcut213 out = F.relu_(out)214 return out215 216 217class DeformBottleneckBlock(CNNBlockBase):218 """219 Similar to :class:`BottleneckBlock`, but with :paper:`deformable conv <deformconv>`220 in the 3x3 convolution.221 """222 223 def __init__(224 self,225 in_channels,226 out_channels,227 *,228 bottleneck_channels,229 stride=1,230 num_groups=1,231 norm="BN",232 stride_in_1x1=False,233 dilation=1,234 deform_modulated=False,235 deform_num_groups=1,236 ):237 super().__init__(in_channels, out_channels, stride)238 self.deform_modulated = deform_modulated239 240 if in_channels != out_channels:241 self.shortcut = Conv2d(242 in_channels,243 out_channels,244 kernel_size=1,245 stride=stride,246 bias=False,247 norm=get_norm(norm, out_channels),248 )249 else:250 self.shortcut = None251 252 stride_1x1, stride_3x3 = (stride, 1) if stride_in_1x1 else (1, stride)253 254 self.conv1 = Conv2d(255 in_channels,256 bottleneck_channels,257 kernel_size=1,258 stride=stride_1x1,259 bias=False,260 norm=get_norm(norm, bottleneck_channels),261 )262 263 if deform_modulated:264 deform_conv_op = ModulatedDeformConv265 # offset channels are 2 or 3 (if with modulated) * kernel_size * kernel_size266 offset_channels = 27267 else:268 deform_conv_op = DeformConv269 offset_channels = 18270 271 self.conv2_offset = Conv2d(272 bottleneck_channels,273 offset_channels * deform_num_groups,274 kernel_size=3,275 stride=stride_3x3,276 padding=1 * dilation,277 dilation=dilation,278 )279 self.conv2 = deform_conv_op(280 bottleneck_channels,281 bottleneck_channels,282 kernel_size=3,283 stride=stride_3x3,284 padding=1 * dilation,285 bias=False,286 groups=num_groups,287 dilation=dilation,288 deformable_groups=deform_num_groups,289 norm=get_norm(norm, bottleneck_channels),290 )291 292 self.conv3 = Conv2d(293 bottleneck_channels,294 out_channels,295 kernel_size=1,296 bias=False,297 norm=get_norm(norm, out_channels),298 )299 300 for layer in [self.conv1, self.conv2, self.conv3, self.shortcut]:301 if layer is not None: # shortcut can be None302 weight_init.c2_msra_fill(layer)303 304 nn.init.constant_(self.conv2_offset.weight, 0)305 nn.init.constant_(self.conv2_offset.bias, 0)306 307 def forward(self, x):308 out = self.conv1(x)309 out = F.relu_(out)310 311 if self.deform_modulated:312 offset_mask = self.conv2_offset(out)313 offset_x, offset_y, mask = torch.chunk(offset_mask, 3, dim=1)314 offset = torch.cat((offset_x, offset_y), dim=1)315 mask = mask.sigmoid()316 out = self.conv2(out, offset, mask)317 else:318 offset = self.conv2_offset(out)319 out = self.conv2(out, offset)320 out = F.relu_(out)321 322 out = self.conv3(out)323 324 if self.shortcut is not None:325 shortcut = self.shortcut(x)326 else:327 shortcut = x328 329 out += shortcut330 out = F.relu_(out)331 return out332 333 334class BasicStem(CNNBlockBase):335 """336 The standard ResNet stem (layers before the first residual block),337 with a conv, relu and max_pool.338 """339 340 def __init__(self, in_channels=3, out_channels=64, norm="BN"):341 """342 Args:343 norm (str or callable): norm after the first conv layer.344 See :func:`layers.get_norm` for supported format.345 """346 super().__init__(in_channels, out_channels, 4)347 self.in_channels = in_channels348 self.conv1 = Conv2d(349 in_channels,350 out_channels,351 kernel_size=7,352 stride=2,353 padding=3,354 bias=False,355 norm=get_norm(norm, out_channels),356 )357 weight_init.c2_msra_fill(self.conv1)358 359 def forward(self, x):360 x = self.conv1(x)361 x = F.relu_(x)362 x = F.max_pool2d(x, kernel_size=3, stride=2, padding=1)363 return x364 365 366class ResNet(Backbone):367 """368 Implement :paper:`ResNet`.369 """370 371 def __init__(self, stem, stages, num_classes=None, out_features=None, freeze_at=0):372 """373 Args:374 stem (nn.Module): a stem module375 stages (list[list[CNNBlockBase]]): several (typically 4) stages,376 each contains multiple :class:`CNNBlockBase`.377 num_classes (None or int): if None, will not perform classification.378 Otherwise, will create a linear layer.379 out_features (list[str]): name of the layers whose outputs should380 be returned in forward. Can be anything in "stem", "linear", or "res2" ...381 If None, will return the output of the last layer.382 freeze_at (int): The number of stages at the beginning to freeze.383 see :meth:`freeze` for detailed explanation.384 """385 super().__init__()386 self.stem = stem387 self.num_classes = num_classes388 389 current_stride = self.stem.stride390 self._out_feature_strides = {"stem": current_stride}391 self._out_feature_channels = {"stem": self.stem.out_channels}392 393 self.stage_names, self.stages = [], []394 395 if out_features is not None:396 # Avoid keeping unused layers in this module. They consume extra memory397 # and may cause allreduce to fail398 num_stages = max(399 [{"res2": 1, "res3": 2, "res4": 3, "res5": 4}.get(f, 0) for f in out_features]400 )401 stages = stages[:num_stages]402 for i, blocks in enumerate(stages):403 assert len(blocks) > 0, len(blocks)404 for block in blocks:405 assert isinstance(block, CNNBlockBase), block406 407 name = "res" + str(i + 2)408 stage = nn.Sequential(*blocks)409 410 self.add_module(name, stage)411 self.stage_names.append(name)412 self.stages.append(stage)413 414 self._out_feature_strides[name] = current_stride = int(415 current_stride * np.prod([k.stride for k in blocks])416 )417 self._out_feature_channels[name] = curr_channels = blocks[-1].out_channels418 self.stage_names = tuple(self.stage_names) # Make it static for scripting419 420 if num_classes is not None:421 self.avgpool = nn.AdaptiveAvgPool2d((1, 1))422 self.linear = nn.Linear(curr_channels, num_classes)423 424 # Sec 5.1 in "Accurate, Large Minibatch SGD: Training ImageNet in 1 Hour":425 # "The 1000-way fully-connected layer is initialized by426 # drawing weights from a zero-mean Gaussian with standard deviation of 0.01."427 nn.init.normal_(self.linear.weight, std=0.01)428 name = "linear"429 430 if out_features is None:431 out_features = [name]432 self._out_features = out_features433 assert len(self._out_features)434 children = [x[0] for x in self.named_children()]435 for out_feature in self._out_features:436 assert out_feature in children, "Available children: {}".format(", ".join(children))437 self.freeze(freeze_at)438 439 def forward(self, x):440 """441 Args:442 x: Tensor of shape (N,C,H,W). H, W must be a multiple of ``self.size_divisibility``.443 444 Returns:445 dict[str->Tensor]: names and the corresponding features446 """447 assert x.dim() == 4, f"ResNet takes an input of shape (N, C, H, W). Got {x.shape} instead!"448 outputs = {}449 x = self.stem(x)450 if "stem" in self._out_features:451 outputs["stem"] = x452 for name, stage in zip(self.stage_names, self.stages):453 x = stage(x)454 if name in self._out_features:455 outputs[name] = x456 if self.num_classes is not None:457 x = self.avgpool(x)458 x = torch.flatten(x, 1)459 x = self.linear(x)460 if "linear" in self._out_features:461 outputs["linear"] = x462 return outputs463 464 def output_shape(self):465 return {466 name: ShapeSpec(467 channels=self._out_feature_channels[name], stride=self._out_feature_strides[name]468 )469 for name in self._out_features470 }471 472 def freeze(self, freeze_at=0):473 """474 Freeze the first several stages of the ResNet. Commonly used in475 fine-tuning.476 477 Layers that produce the same feature map spatial size are defined as one478 "stage" by :paper:`FPN`.479 480 Args:481 freeze_at (int): number of stages to freeze.482 `1` means freezing the stem. `2` means freezing the stem and483 one residual stage, etc.484 485 Returns:486 nn.Module: this ResNet itself487 """488 if freeze_at >= 1:489 self.stem.freeze()490 for idx, stage in enumerate(self.stages, start=2):491 if freeze_at >= idx:492 for block in stage.children():493 block.freeze()494 return self495 496 @staticmethod497 def make_stage(block_class, num_blocks, *, in_channels, out_channels, **kwargs):498 """499 Create a list of blocks of the same type that forms one ResNet stage.500 501 Args:502 block_class (type): a subclass of CNNBlockBase that's used to create all blocks in this503 stage. A module of this type must not change spatial resolution of inputs unless its504 stride != 1.505 num_blocks (int): number of blocks in this stage506 in_channels (int): input channels of the entire stage.507 out_channels (int): output channels of **every block** in the stage.508 kwargs: other arguments passed to the constructor of509 `block_class`. If the argument name is "xx_per_block", the510 argument is a list of values to be passed to each block in the511 stage. Otherwise, the same argument is passed to every block512 in the stage.513 514 Returns:515 list[CNNBlockBase]: a list of block module.516 517 Examples:518 ::519 stage = ResNet.make_stage(520 BottleneckBlock, 3, in_channels=16, out_channels=64,521 bottleneck_channels=16, num_groups=1,522 stride_per_block=[2, 1, 1],523 dilations_per_block=[1, 1, 2]524 )525 526 Usually, layers that produce the same feature map spatial size are defined as one527 "stage" (in :paper:`FPN`). Under such definition, ``stride_per_block[1:]`` should528 all be 1.529 """530 blocks = []531 for i in range(num_blocks):532 curr_kwargs = {}533 for k, v in kwargs.items():534 if k.endswith("_per_block"):535 assert len(v) == num_blocks, (536 f"Argument '{k}' of make_stage should have the "537 f"same length as num_blocks={num_blocks}."538 )539 newk = k[: -len("_per_block")]540 assert newk not in kwargs, f"Cannot call make_stage with both {k} and {newk}!"541 curr_kwargs[newk] = v[i]542 else:543 curr_kwargs[k] = v544 545 blocks.append(546 block_class(in_channels=in_channels, out_channels=out_channels, **curr_kwargs)547 )548 in_channels = out_channels549 return blocks550 551 @staticmethod552 def make_default_stages(depth, block_class=None, **kwargs):553 """554 Created list of ResNet stages from pre-defined depth (one of 18, 34, 50, 101, 152).555 If it doesn't create the ResNet variant you need, please use :meth:`make_stage`556 instead for fine-grained customization.557 558 Args:559 depth (int): depth of ResNet560 block_class (type): the CNN block class. Has to accept561 `bottleneck_channels` argument for depth > 50.562 By default it is BasicBlock or BottleneckBlock, based on the563 depth.564 kwargs:565 other arguments to pass to `make_stage`. Should not contain566 stride and channels, as they are predefined for each depth.567 568 Returns:569 list[list[CNNBlockBase]]: modules in all stages; see arguments of570 :class:`ResNet.__init__`.571 """572 num_blocks_per_stage = {573 18: [2, 2, 2, 2],574 34: [3, 4, 6, 3],575 50: [3, 4, 6, 3],576 101: [3, 4, 23, 3],577 152: [3, 8, 36, 3],578 }[depth]579 if block_class is None:580 block_class = BasicBlock if depth < 50 else BottleneckBlock581 if depth < 50:582 in_channels = [64, 64, 128, 256]583 out_channels = [64, 128, 256, 512]584 else:585 in_channels = [64, 256, 512, 1024]586 out_channels = [256, 512, 1024, 2048]587 ret = []588 for (n, s, i, o) in zip(num_blocks_per_stage, [1, 2, 2, 2], in_channels, out_channels):589 if depth >= 50:590 kwargs["bottleneck_channels"] = o // 4591 ret.append(592 ResNet.make_stage(593 block_class=block_class,594 num_blocks=n,595 stride_per_block=[s] + [1] * (n - 1),596 in_channels=i,597 out_channels=o,598 **kwargs,599 )600 )601 return ret602 603 604ResNetBlockBase = CNNBlockBase605"""606Alias for backward compatibiltiy.607"""608 609 610def make_stage(*args, **kwargs):611 """612 Deprecated alias for backward compatibiltiy.613 """614 return ResNet.make_stage(*args, **kwargs)615 616 617def _convert_ndarray_to_tensor(state_dict: Dict[str, Any]) -> None:618 """619 In-place convert all numpy arrays in the state_dict to torch tensor.620 Args:621 state_dict (dict): a state-dict to be loaded to the model.622 Will be modified.623 """624 # model could be an OrderedDict with _metadata attribute625 # (as returned by Pytorch's state_dict()). We should preserve these626 # properties.627 for k in list(state_dict.keys()):628 v = state_dict[k]629 if not isinstance(v, np.ndarray) and not isinstance(v, torch.Tensor):630 raise ValueError(631 "Unsupported type found in checkpoint! {}: {}".format(k, type(v))632 )633 if not isinstance(v, torch.Tensor):634 state_dict[k] = torch.from_numpy(v)635 636 637@register_backbone638def get_resnet_backbone(cfg):639 """640 Create a ResNet instance from config.641 642 Returns:643 ResNet: a :class:`ResNet` instance.644 """645 res_cfg = cfg['MODEL']['BACKBONE']['RESNETS']646 647 # need registration of new blocks/stems?648 norm = res_cfg['NORM']649 stem = BasicStem(650 in_channels=res_cfg['STEM_IN_CHANNELS'],651 out_channels=res_cfg['STEM_OUT_CHANNELS'],652 norm=norm,653 )654 655 # fmt: off656 freeze_at = res_cfg['FREEZE_AT']657 out_features = res_cfg['OUT_FEATURES']658 depth = res_cfg['DEPTH']659 num_groups = res_cfg['NUM_GROUPS']660 width_per_group = res_cfg['WIDTH_PER_GROUP']661 bottleneck_channels = num_groups * width_per_group662 in_channels = res_cfg['STEM_OUT_CHANNELS']663 out_channels = res_cfg['RES2_OUT_CHANNELS']664 stride_in_1x1 = res_cfg['STRIDE_IN_1X1']665 res5_dilation = res_cfg['RES5_DILATION']666 deform_on_per_stage = res_cfg['DEFORM_ON_PER_STAGE']667 deform_modulated = res_cfg['DEFORM_MODULATED']668 deform_num_groups = res_cfg['DEFORM_NUM_GROUPS']669 # fmt: on670 assert res5_dilation in {1, 2}, "res5_dilation cannot be {}.".format(res5_dilation)671 672 num_blocks_per_stage = {673 18: [2, 2, 2, 2],674 34: [3, 4, 6, 3],675 50: [3, 4, 6, 3],676 101: [3, 4, 23, 3],677 152: [3, 8, 36, 3],678 }[depth]679 680 if depth in [18, 34]:681 assert out_channels == 64, "Must set MODEL.RESNETS.RES2_OUT_CHANNELS = 64 for R18/R34"682 assert not any(683 deform_on_per_stage684 ), "MODEL.RESNETS.DEFORM_ON_PER_STAGE unsupported for R18/R34"685 assert res5_dilation == 1, "Must set MODEL.RESNETS.RES5_DILATION = 1 for R18/R34"686 assert num_groups == 1, "Must set MODEL.RESNETS.NUM_GROUPS = 1 for R18/R34"687 688 stages = []689 690 for idx, stage_idx in enumerate(range(2, 6)):691 # res5_dilation is used this way as a convention in R-FCN & Deformable Conv paper692 dilation = res5_dilation if stage_idx == 5 else 1693 first_stride = 1 if idx == 0 or (stage_idx == 5 and dilation == 2) else 2694 stage_kargs = {695 "num_blocks": num_blocks_per_stage[idx],696 "stride_per_block": [first_stride] + [1] * (num_blocks_per_stage[idx] - 1),697 "in_channels": in_channels,698 "out_channels": out_channels,699 "norm": norm,700 }701 # Use BasicBlock for R18 and R34.702 if depth in [18, 34]:703 stage_kargs["block_class"] = BasicBlock704 else:705 stage_kargs["bottleneck_channels"] = bottleneck_channels706 stage_kargs["stride_in_1x1"] = stride_in_1x1707 stage_kargs["dilation"] = dilation708 stage_kargs["num_groups"] = num_groups709 if deform_on_per_stage[idx]:710 stage_kargs["block_class"] = DeformBottleneckBlock711 stage_kargs["deform_modulated"] = deform_modulated712 stage_kargs["deform_num_groups"] = deform_num_groups713 else:714 stage_kargs["block_class"] = BottleneckBlock715 blocks = ResNet.make_stage(**stage_kargs)716 in_channels = out_channels717 out_channels *= 2718 bottleneck_channels *= 2719 stages.append(blocks)720 backbone = ResNet(stem, stages, out_features=out_features, freeze_at=freeze_at)721 722 if cfg['MODEL']['BACKBONE']['LOAD_PRETRAINED'] is True:723 filename = cfg['MODEL']['BACKBONE']['PRETRAINED']724 with PathManager.open(filename, "rb") as f:725 ckpt = pickle.load(f, encoding="latin1")['model']726 _convert_ndarray_to_tensor(ckpt)727 ckpt.pop('stem.fc.weight')728 ckpt.pop('stem.fc.bias')729 backbone.load_state_dict(ckpt)730 731 return backbone732 