xdecoder/Instruct-X-Decoder
163
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