CoolFace
Apppublic

xdecoder/Instruct-X-Decoder

sourceHugging Faceafl-3.0updated 3y agoView on Hugging Face
163likes
transformer_encoder_fpn.py324 linesDownload Raw Back to encoder
1# Copyright (c) Facebook, Inc. and its affiliates.2import logging3import numpy as np4from typing import Callable, Dict, List, Optional, Tuple, Union5 6import torch7from torch import nn8from torch.nn import functional as F9from torch.nn.init import xavier_uniform_, constant_, uniform_, normal_10from torch.cuda.amp import autocast11 12import fvcore.nn.weight_init as weight_init13from detectron2.layers import Conv2d, DeformConv, ShapeSpec, get_norm14 15from .registry import register_encoder16from ..transformer_blocks import TransformerEncoder, TransformerEncoderLayer, _get_clones, _get_activation_fn17from ...modules import PositionEmbeddingSine18from ...utils import configurable19 20# from ..layers import Conv2d, DeformConv, ShapeSpec, get_norm21 22# This is a modified FPN decoder.23class BasePixelDecoder(nn.Module):24    def __init__(25        self,26        input_shape: Dict[str, ShapeSpec],27        *,28        conv_dim: int,29        mask_dim: int,30        mask_on: bool,31        norm: Optional[Union[str, Callable]] = None,32    ):33        """34        NOTE: this interface is experimental.35        Args:36            input_shape: shapes (channels and stride) of the input features37            conv_dims: number of output channels for the intermediate conv layers.38            mask_dim: number of output channels for the final conv layer.39            norm (str or callable): normalization for all conv layers40        """41        super().__init__()42 43        input_shape = sorted(input_shape.items(), key=lambda x: x[1].stride)44        self.in_features = [k for k, v in input_shape]  # starting from "res2" to "res5"45        feature_channels = [v.channels for k, v in input_shape]46 47        lateral_convs = []48        output_convs = []49 50        use_bias = norm == ""51        for idx, in_channels in enumerate(feature_channels):52            if idx == len(self.in_features) - 1:53                output_norm = get_norm(norm, conv_dim)54                output_conv = Conv2d(55                    in_channels,56                    conv_dim,57                    kernel_size=3,58                    stride=1,59                    padding=1,60                    bias=use_bias,61                    norm=output_norm,62                    activation=F.relu,63                )64                weight_init.c2_xavier_fill(output_conv)65                self.add_module("layer_{}".format(idx + 1), output_conv)66 67                lateral_convs.append(None)68                output_convs.append(output_conv)69            else:70                lateral_norm = get_norm(norm, conv_dim)71                output_norm = get_norm(norm, conv_dim)72 73                lateral_conv = Conv2d(74                    in_channels, conv_dim, kernel_size=1, bias=use_bias, norm=lateral_norm75                )76                output_conv = Conv2d(77                    conv_dim,78                    conv_dim,79                    kernel_size=3,80                    stride=1,81                    padding=1,82                    bias=use_bias,83                    norm=output_norm,84                    activation=F.relu,85                )86                weight_init.c2_xavier_fill(lateral_conv)87                weight_init.c2_xavier_fill(output_conv)88                self.add_module("adapter_{}".format(idx + 1), lateral_conv)89                self.add_module("layer_{}".format(idx + 1), output_conv)90 91                lateral_convs.append(lateral_conv)92                output_convs.append(output_conv)93        # Place convs into top-down order (from low to high resolution)94        # to make the top-down computation in forward clearer.95        self.lateral_convs = lateral_convs[::-1]96        self.output_convs = output_convs[::-1]97 98        self.mask_on = mask_on99        if self.mask_on:100            self.mask_dim = mask_dim101            self.mask_features = Conv2d(102                conv_dim,103                mask_dim,104                kernel_size=3,105                stride=1,106                padding=1,107            )108            weight_init.c2_xavier_fill(self.mask_features)109 110        self.maskformer_num_feature_levels = 3  # always use 3 scales111 112    @classmethod113    def from_config(cls, cfg, input_shape: Dict[str, ShapeSpec]):114        enc_cfg = cfg['MODEL']['ENCODER']115        ret = {}116        ret["input_shape"] = {117            k: v for k, v in input_shape.items() if k in enc_cfg['IN_FEATURES']118        }119        ret["conv_dim"] = enc_cfg['CONVS_DIM']120        ret["mask_dim"] = enc_cfg['MASK_DIM']121        ret["norm"] = enc_cfg['NORM']122        return ret123 124    def forward_features(self, features):125        multi_scale_features = []126        num_cur_levels = 0127        # Reverse feature maps into top-down order (from low to high resolution)128        for idx, f in enumerate(self.in_features[::-1]):129            x = features[f]130            lateral_conv = self.lateral_convs[idx]131            output_conv = self.output_convs[idx]132            if lateral_conv is None:133                y = output_conv(x)134            else:135                cur_fpn = lateral_conv(x)136                # Following FPN implementation, we use nearest upsampling here137                y = cur_fpn + F.interpolate(y, size=cur_fpn.shape[-2:], mode="nearest")138                y = output_conv(y)139            if num_cur_levels < self.maskformer_num_feature_levels:140                multi_scale_features.append(y)141                num_cur_levels += 1142        143        mask_features = self.mask_features(y) if self.mask_on else None144        return mask_features, None, multi_scale_features145 146    def forward(self, features, targets=None):147        logger = logging.getLogger(__name__)148        logger.warning("Calling forward() may cause unpredicted behavior of PixelDecoder module.")149        return self.forward_features(features)150 151 152class TransformerEncoderOnly(nn.Module):153    def __init__(154        self,155        d_model=512,156        nhead=8,157        num_encoder_layers=6,158        dim_feedforward=2048,159        dropout=0.1,160        activation="relu",161        normalize_before=False,162    ):163        super().__init__()164 165        encoder_layer = TransformerEncoderLayer(166            d_model, nhead, dim_feedforward, dropout, activation, normalize_before167        )168        encoder_norm = nn.LayerNorm(d_model) if normalize_before else None169        self.encoder = TransformerEncoder(encoder_layer, num_encoder_layers, encoder_norm)170 171        self._reset_parameters()172 173        self.d_model = d_model174        self.nhead = nhead175 176    def _reset_parameters(self):177        for p in self.parameters():178            if p.dim() > 1:179                nn.init.xavier_uniform_(p)180 181    def forward(self, src, mask, pos_embed):182        # flatten NxCxHxW to HWxNxC183        bs, c, h, w = src.shape184        src = src.flatten(2).permute(2, 0, 1)185        pos_embed = pos_embed.flatten(2).permute(2, 0, 1)186        if mask is not None:187            mask = mask.flatten(1)188 189        memory = self.encoder(src, src_key_padding_mask=mask, pos=pos_embed)190        return memory.permute(1, 2, 0).view(bs, c, h, w)191 192 193# This is a modified FPN decoder with extra Transformer encoder that processes the lowest-resolution feature map.194class TransformerEncoderPixelDecoder(BasePixelDecoder):195    @configurable196    def __init__(197        self,198        input_shape: Dict[str, ShapeSpec],199        *,200        transformer_dropout: float,201        transformer_nheads: int,202        transformer_dim_feedforward: int,203        transformer_enc_layers: int,204        transformer_pre_norm: bool,205        conv_dim: int,206        mask_dim: int,207        mask_on: int,208        norm: Optional[Union[str, Callable]] = None,209    ):210        """211        NOTE: this interface is experimental.212        Args:213            input_shape: shapes (channels and stride) of the input features214            transformer_dropout: dropout probability in transformer215            transformer_nheads: number of heads in transformer216            transformer_dim_feedforward: dimension of feedforward network217            transformer_enc_layers: number of transformer encoder layers218            transformer_pre_norm: whether to use pre-layernorm or not219            conv_dims: number of output channels for the intermediate conv layers.220            mask_dim: number of output channels for the final conv layer.221            norm (str or callable): normalization for all conv layers222        """223        super().__init__(input_shape, conv_dim=conv_dim, mask_dim=mask_dim, norm=norm, mask_on=mask_on)224 225        input_shape = sorted(input_shape.items(), key=lambda x: x[1].stride)226        self.in_features = [k for k, v in input_shape]  # starting from "res2" to "res5"227        feature_strides = [v.stride for k, v in input_shape]228        feature_channels = [v.channels for k, v in input_shape]229 230        in_channels = feature_channels[len(self.in_features) - 1]231        self.input_proj = Conv2d(in_channels, conv_dim, kernel_size=1)232        weight_init.c2_xavier_fill(self.input_proj)233        self.transformer = TransformerEncoderOnly(234            d_model=conv_dim,235            dropout=transformer_dropout,236            nhead=transformer_nheads,237            dim_feedforward=transformer_dim_feedforward,238            num_encoder_layers=transformer_enc_layers,239            normalize_before=transformer_pre_norm,240        )241        N_steps = conv_dim // 2242        self.pe_layer = PositionEmbeddingSine(N_steps, normalize=True)243 244        # update layer245        use_bias = norm == ""246        output_norm = get_norm(norm, conv_dim)247        output_conv = Conv2d(248            conv_dim,249            conv_dim,250            kernel_size=3,251            stride=1,252            padding=1,253            bias=use_bias,254            norm=output_norm,255            activation=F.relu,256        )257        weight_init.c2_xavier_fill(output_conv)258        delattr(self, "layer_{}".format(len(self.in_features)))259        self.add_module("layer_{}".format(len(self.in_features)), output_conv)260        self.output_convs[0] = output_conv261 262    @classmethod263    def from_config(cls, cfg, input_shape: Dict[str, ShapeSpec]):264        enc_cfg = cfg['MODEL']['ENCODER']265        dec_cfg = cfg['MODEL']['DECODER']266 267        ret = super().from_config(cfg, input_shape)268        ret["transformer_dropout"] = dec_cfg['DROPOUT']269        ret["transformer_nheads"] = dec_cfg['NHEADS']270        ret["transformer_dim_feedforward"] = dec_cfg['DIM_FEEDFORWARD']271        ret["transformer_enc_layers"] = enc_cfg['TRANSFORMER_ENC_LAYERS']  # a separate config272        ret["transformer_pre_norm"] = dec_cfg['PRE_NORM']273 274        ret['mask_on'] = cfg['MODEL']['DECODER']['MASK']275        return ret276 277    def forward_features(self, features):278        multi_scale_features = []279        num_cur_levels = 0280        281        # Reverse feature maps into top-down order (from low to high resolution)282        for idx, f in enumerate(self.in_features[::-1]):283            x = features[f]284            lateral_conv = self.lateral_convs[idx]285            output_conv = self.output_convs[idx]286            if lateral_conv is None:287                transformer = self.input_proj(x)288                pos = self.pe_layer(x)289                transformer = self.transformer(transformer, None, pos)290                y = output_conv(transformer)291                # save intermediate feature as input to Transformer decoder292                transformer_encoder_features = transformer293            else:294                cur_fpn = lateral_conv(x)295                # Following FPN implementation, we use nearest upsampling here296                y = cur_fpn + F.interpolate(y, size=cur_fpn.shape[-2:], mode="nearest")297                y = output_conv(y)298            if num_cur_levels < self.maskformer_num_feature_levels:299                multi_scale_features.append(y)300                num_cur_levels += 1301 302        mask_features = self.mask_features(y) if self.mask_on else None303        return mask_features, transformer_encoder_features, multi_scale_features304 305    def forward(self, features, targets=None):306        logger = logging.getLogger(__name__)307        logger.warning("Calling forward() may cause unpredicted behavior of PixelDecoder module.")308        return self.forward_features(features)309 310 311 312@register_encoder313def get_transformer_encoder_fpn(cfg, input_shape):314    """315    Build a pixel decoder from `cfg.MODEL.MASK_FORMER.PIXEL_DECODER_NAME`.316    """317    model = TransformerEncoderPixelDecoder(cfg, input_shape)    318    forward_features = getattr(model, "forward_features", None)319    if not callable(forward_features):320        raise ValueError(321            "Only SEM_SEG_HEADS with forward_features method can be used as pixel decoder. "322            f"Please implement forward_features for {name} to only return mask features."323        )324    return model