Team Ai
Apppublic

xdecoder/Instruct-X-Decoder

sourceHugging Faceafl-3.0updated 3y agoView on Hugging Face
163likes
resnet.py732 linesDownload Raw Back to backbone
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