Team Ai
Apppublic

xdecoder/Instruct-X-Decoder

sourceHugging Faceafl-3.0updated 3y agoView on Hugging Face
163likes
tmp.py664 linesDownload Raw Back to decoder
1# Copyright (c) Facebook, Inc. and its affiliates.2# Modified by Bowen Cheng from: https://github.com/facebookresearch/detr/blob/master/models/detr.py3import logging4from typing import Optional5 6import torch7from torch import nn, Tensor8from torch.nn import functional as F9 10from timm.models.layers import trunc_normal_11from detectron2.layers import Conv2d12import fvcore.nn.weight_init as weight_init13 14from .registry import register_decoder15from ...utils import configurable16from ...modules import PositionEmbeddingSine17 18from image2html.visualizer import VL19 20 21class SelfAttentionLayer(nn.Module):22 23    def __init__(self, d_model, nhead, dropout=0.0,24                 activation="relu", normalize_before=False):25        super().__init__()26        self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)27 28        self.norm = nn.LayerNorm(d_model)29        self.dropout = nn.Dropout(dropout)30 31        self.activation = _get_activation_fn(activation)32        self.normalize_before = normalize_before33 34        self._reset_parameters()35    36    def _reset_parameters(self):37        for p in self.parameters():38            if p.dim() > 1:39                nn.init.xavier_uniform_(p)40 41    def with_pos_embed(self, tensor, pos: Optional[Tensor]):42        return tensor if pos is None else tensor + pos43 44    def forward_post(self, tgt,45                     tgt_mask: Optional[Tensor] = None,46                     tgt_key_padding_mask: Optional[Tensor] = None,47                     query_pos: Optional[Tensor] = None):48        q = k = self.with_pos_embed(tgt, query_pos)49        tgt2 = self.self_attn(q, k, value=tgt, attn_mask=tgt_mask,50                              key_padding_mask=tgt_key_padding_mask)[0]51        tgt = tgt + self.dropout(tgt2)52        tgt = self.norm(tgt)53 54        return tgt55 56    def forward_pre(self, tgt,57                    tgt_mask: Optional[Tensor] = None,58                    tgt_key_padding_mask: Optional[Tensor] = None,59                    query_pos: Optional[Tensor] = None):60        tgt2 = self.norm(tgt)61        q = k = self.with_pos_embed(tgt2, query_pos)62        tgt2 = self.self_attn(q, k, value=tgt2, attn_mask=tgt_mask,63                              key_padding_mask=tgt_key_padding_mask)[0]64        tgt = tgt + self.dropout(tgt2)65        66        return tgt67 68    def forward(self, tgt,69                tgt_mask: Optional[Tensor] = None,70                tgt_key_padding_mask: Optional[Tensor] = None,71                query_pos: Optional[Tensor] = None):72        if self.normalize_before:73            return self.forward_pre(tgt, tgt_mask,74                                    tgt_key_padding_mask, query_pos)75        return self.forward_post(tgt, tgt_mask,76                                 tgt_key_padding_mask, query_pos)77 78 79class CrossAttentionLayer(nn.Module):80 81    def __init__(self, d_model, nhead, dropout=0.0,82                 activation="relu", normalize_before=False):83        super().__init__()84        self.multihead_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)85 86        self.norm = nn.LayerNorm(d_model)87        self.dropout = nn.Dropout(dropout)88 89        self.activation = _get_activation_fn(activation)90        self.normalize_before = normalize_before91 92        self._reset_parameters()93    94    def _reset_parameters(self):95        for p in self.parameters():96            if p.dim() > 1:97                nn.init.xavier_uniform_(p)98 99    def with_pos_embed(self, tensor, pos: Optional[Tensor]):100        return tensor if pos is None else tensor + pos101 102    def forward_post(self, tgt, memory,103                     memory_mask: Optional[Tensor] = None,104                     memory_key_padding_mask: Optional[Tensor] = None,105                     pos: Optional[Tensor] = None,106                     query_pos: Optional[Tensor] = None):107        tgt2, avg_attn = self.multihead_attn(query=self.with_pos_embed(tgt, query_pos),108                                   key=self.with_pos_embed(memory, pos),109                                   value=memory, attn_mask=memory_mask,110                                   key_padding_mask=memory_key_padding_mask)111        tgt = tgt + self.dropout(tgt2)112        tgt = self.norm(tgt)113        return tgt, avg_attn114 115    def forward_pre(self, tgt, memory,116                    memory_mask: Optional[Tensor] = None,117                    memory_key_padding_mask: Optional[Tensor] = None,118                    pos: Optional[Tensor] = None,119                    query_pos: Optional[Tensor] = None):120        tgt2 = self.norm(tgt)121        tgt2, avg_attn = self.multihead_attn(query=self.with_pos_embed(tgt2, query_pos),122                                   key=self.with_pos_embed(memory, pos),123                                   value=memory, attn_mask=memory_mask,124                                   key_padding_mask=memory_key_padding_mask)125        tgt = tgt + self.dropout(tgt2)126 127        return tgt, avg_attn128 129    def forward(self, tgt, memory,130                memory_mask: Optional[Tensor] = None,131                memory_key_padding_mask: Optional[Tensor] = None,132                pos: Optional[Tensor] = None,133                query_pos: Optional[Tensor] = None):134        if self.normalize_before:135            return self.forward_pre(tgt, memory, memory_mask,136                                    memory_key_padding_mask, pos, query_pos)137        return self.forward_post(tgt, memory, memory_mask,138                                 memory_key_padding_mask, pos, query_pos)139 140 141class FFNLayer(nn.Module):142 143    def __init__(self, d_model, dim_feedforward=2048, dropout=0.0,144                 activation="relu", normalize_before=False):145        super().__init__()146        # Implementation of Feedforward model147        self.linear1 = nn.Linear(d_model, dim_feedforward)148        self.dropout = nn.Dropout(dropout)149        self.linear2 = nn.Linear(dim_feedforward, d_model)150 151        self.norm = nn.LayerNorm(d_model)152 153        self.activation = _get_activation_fn(activation)154        self.normalize_before = normalize_before155 156        self._reset_parameters()157    158    def _reset_parameters(self):159        for p in self.parameters():160            if p.dim() > 1:161                nn.init.xavier_uniform_(p)162 163    def with_pos_embed(self, tensor, pos: Optional[Tensor]):164        return tensor if pos is None else tensor + pos165 166    def forward_post(self, tgt):167        tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt))))168        tgt = tgt + self.dropout(tgt2)169        tgt = self.norm(tgt)170        return tgt171 172    def forward_pre(self, tgt):173        tgt2 = self.norm(tgt)174        tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt2))))175        tgt = tgt + self.dropout(tgt2)176        return tgt177 178    def forward(self, tgt):179        if self.normalize_before:180            return self.forward_pre(tgt)181        return self.forward_post(tgt)182 183 184def _get_activation_fn(activation):185    """Return an activation function given a string"""186    if activation == "relu":187        return F.relu188    if activation == "gelu":189        return F.gelu190    if activation == "glu":191        return F.glu192    raise RuntimeError(F"activation should be relu/gelu, not {activation}.")193 194 195class MLP(nn.Module):196    """ Very simple multi-layer perceptron (also called FFN)"""197 198    def __init__(self, input_dim, hidden_dim, output_dim, num_layers):199        super().__init__()200        self.num_layers = num_layers201        h = [hidden_dim] * (num_layers - 1)202        self.layers = nn.ModuleList(nn.Linear(n, k) for n, k in zip([input_dim] + h, h + [output_dim]))203 204    def forward(self, x):205        for i, layer in enumerate(self.layers):206            x = F.relu(layer(x)) if i < self.num_layers - 1 else layer(x)207        return x208 209 210class MultiScaleMaskedTransformerDecoder(nn.Module):211 212    _version = 2213 214    @configurable215    def __init__(216        self,217        lang_encoder: nn.Module,218        in_channels,219        mask_classification=True,220        *,221        hidden_dim: int,222        dim_proj: int,223        num_queries: int,224        contxt_len: int,225        nheads: int,226        dim_feedforward: int,227        dec_layers: int,228        pre_norm: bool,229        mask_dim: int,230        task_switch: dict,231        captioning_step: int,232        enforce_input_project: bool,233    ):234        """235        NOTE: this interface is experimental.236        Args:237            in_channels: channels of the input features238            mask_classification: whether to add mask classifier or not239            num_classes: number of classes240            hidden_dim: Transformer feature dimension241            num_queries: number of queries242            nheads: number of heads243            dim_feedforward: feature dimension in feedforward network244            enc_layers: number of Transformer encoder layers245            dec_layers: number of Transformer decoder layers246            pre_norm: whether to use pre-LayerNorm or not247            mask_dim: mask feature dimension248            enforce_input_project: add input project 1x1 conv even if input249                channels and hidden dim is identical250        """251        super().__init__()252        assert mask_classification, "Only support mask classification model"253        self.mask_classification = mask_classification254 255        # positional encoding256        N_steps = hidden_dim // 2257        self.pe_layer = PositionEmbeddingSine(N_steps, normalize=True)258        259        # define Transformer decoder here260        self.num_heads = nheads261        self.num_layers = dec_layers262        self.contxt_len = contxt_len263        self.transformer_self_attention_layers = nn.ModuleList()264        self.transformer_cross_attention_layers = nn.ModuleList()265        self.transformer_ffn_layers = nn.ModuleList()266 267        for _ in range(self.num_layers):268            self.transformer_self_attention_layers.append(269                SelfAttentionLayer(270                    d_model=hidden_dim,271                    nhead=nheads,272                    dropout=0.0,273                    normalize_before=pre_norm,274                )275            )276 277            self.transformer_cross_attention_layers.append(278                CrossAttentionLayer(279                    d_model=hidden_dim,280                    nhead=nheads,281                    dropout=0.0,282                    normalize_before=pre_norm,283                )284            )285 286            self.transformer_ffn_layers.append(287                FFNLayer(288                    d_model=hidden_dim,289                    dim_feedforward=dim_feedforward,290                    dropout=0.0,291                    normalize_before=pre_norm,292                )293            )294 295        self.decoder_norm = nn.LayerNorm(hidden_dim)296 297        self.num_queries = num_queries298        # learnable query features299        self.query_feat = nn.Embedding(num_queries, hidden_dim)300        # learnable query p.e.301        self.query_embed = nn.Embedding(num_queries, hidden_dim)302        303        # level embedding (we always use 3 scales)304        self.num_feature_levels = 3305        self.level_embed = nn.Embedding(self.num_feature_levels, hidden_dim)306        self.input_proj = nn.ModuleList()307        308        for _ in range(self.num_feature_levels):309            if in_channels != hidden_dim or enforce_input_project:310                self.input_proj.append(Conv2d(in_channels, hidden_dim, kernel_size=1))311                weight_init.c2_xavier_fill(self.input_proj[-1])312            else:313                self.input_proj.append(nn.Sequential())314 315        self.task_switch = task_switch316 317        # output FFNs318        self.lang_encoder = lang_encoder319        if self.task_switch['mask']:320            self.mask_embed = MLP(hidden_dim, hidden_dim, mask_dim, 3)321 322        self.class_embed = nn.Parameter(torch.empty(hidden_dim, dim_proj))323        trunc_normal_(self.class_embed, std=.02)324 325        if task_switch['bbox']:326            self.bbox_embed = MLP(hidden_dim, hidden_dim, 4, 3)327 328        # Caption Project and query329        if task_switch['captioning']:330            self.caping_embed = nn.Parameter(torch.empty(hidden_dim, dim_proj))331            trunc_normal_(self.caping_embed, std=.02)332            self.query_feat_caping = nn.Embedding(contxt_len, hidden_dim)333            self.captioning_step = captioning_step334 335        # register self_attn_mask to avoid information leakage, it includes interaction between object query, class query and caping query336        self_attn_mask = torch.zeros((1, num_queries + contxt_len, num_queries + contxt_len)).bool()337        self_attn_mask[:, :num_queries, num_queries:] = True # object+class query does not attend with caption query.338        self_attn_mask[:, num_queries:, num_queries:] = torch.triu(torch.ones((1, contxt_len, contxt_len)), diagonal=1).bool() # caption query only attend with previous token.339        self_attn_mask[:, :num_queries-1, num_queries-1:num_queries] = True # object query does not attend with class query.340        self_attn_mask[:, num_queries-1:num_queries, :num_queries-1] = True # class query does not attend with object query.341        self.register_buffer("self_attn_mask", self_attn_mask)342 343 344    @classmethod345    def from_config(cls, cfg, in_channels, lang_encoder, mask_classification, extra):346        ret = {}347 348        ret["lang_encoder"] = lang_encoder349        ret["in_channels"] = in_channels350        ret["mask_classification"] = mask_classification351 352        enc_cfg = cfg['MODEL']['ENCODER']353        dec_cfg = cfg['MODEL']['DECODER']354        355        ret["hidden_dim"] = dec_cfg['HIDDEN_DIM']356        ret["dim_proj"] = cfg['MODEL']['DIM_PROJ']357        ret["num_queries"] = dec_cfg['NUM_OBJECT_QUERIES']358        ret["contxt_len"] = cfg['MODEL']['TEXT']['CONTEXT_LENGTH']359        360        # Transformer parameters:361        ret["nheads"] = dec_cfg['NHEADS']362        ret["dim_feedforward"] = dec_cfg['DIM_FEEDFORWARD']363 364        # NOTE: because we add learnable query features which requires supervision,365        # we add minus 1 to decoder layers to be consistent with our loss366        # implementation: that is, number of auxiliary losses is always367        # equal to number of decoder layers. With learnable query features, the number of368        # auxiliary losses equals number of decoders plus 1.369        assert dec_cfg['DEC_LAYERS'] >= 1370        ret["dec_layers"] = dec_cfg['DEC_LAYERS'] - 1371        ret["pre_norm"] = dec_cfg['PRE_NORM']372        ret["enforce_input_project"] = dec_cfg['ENFORCE_INPUT_PROJ']373        ret["mask_dim"] = enc_cfg['MASK_DIM']374 375        ret["task_switch"] = extra['task_switch']376        ret["captioning_step"] = dec_cfg['CAPTIONING'].get('STEP', 50)377 378        return ret379 380    def forward(self, x, mask_features, mask=None, target_queries=None, target_vlp=None, task='seg', extra={}):381        if task == 'captioning_infer':382            return self.forward_captioning(x, mask_features, mask=mask, target_queries=target_queries, target_vlp=target_vlp, task=task, extra=extra)383        # x is a list of multi-scale feature384        assert len(x) == self.num_feature_levels385        src = []386        pos = []387        size_list = []388        389        # disable mask, it does not affect performance390        del mask391        for i in range(self.num_feature_levels):392            size_list.append(x[i].shape[-2:])393            pos.append(self.pe_layer(x[i], None).flatten(2))394            src.append(self.input_proj[i](x[i]).flatten(2) + self.level_embed.weight[i][None, :, None])395 396            # flatten NxCxHxW to HWxNxC397            pos[-1] = pos[-1].permute(2, 0, 1)398            src[-1] = src[-1].permute(2, 0, 1)399 400        _, bs, _ = src[0].shape401 402        # QxNxC403        query_embed = self.query_embed.weight.unsqueeze(1).repeat(1, bs, 1)404        output = self.query_feat.weight.unsqueeze(1).repeat(1, bs, 1)405 406        predictions_class = []407        predictions_mask = []408        predictions_bbox = []409        predictions_caption = []410        predictions_captioning = []411        412        self_tgt_mask = None413        if self.training and task == 'vlp' and self.task_switch['captioning']:414            output = torch.cat((output, self.query_feat_caping.weight.unsqueeze(1).repeat(1, bs, 1)), dim=0) # concat object query, class token and caption token.415            caping_lang_embed = torch.cat([caption['caption_tokens'] for caption in target_vlp], dim=0).transpose(0, 1) # language output416            query_embed = torch.cat((query_embed, caping_lang_embed), dim=0) # may not add at the beginning.417            self_tgt_mask = self.self_attn_mask.repeat(output.shape[1]*self.num_heads, 1, 1)418        elif (((self.training and task == 'seg') or (task == 'grounding_eval')) and self.task_switch['grounding']) \419                or ((self.training and task == 'openimage') and self.task_switch['openimage']['grounding']):420            self_tgt_mask = self.self_attn_mask[:,:self.num_queries,:self.num_queries].repeat(output.shape[1]*self.num_heads, 1, 1)421            grounding_tokens = extra['grounding_tokens']422            _grounding_tokens = grounding_tokens.detach().clone()423            # initialize with negative attention at the beginning.424            pad_tgt_mask = torch.ones((1, self.num_queries + (self.num_queries-1) + len(grounding_tokens), self.num_queries + (self.num_queries-1) + len(grounding_tokens)), device=self_tgt_mask.device).bool().repeat(output.shape[1]*self.num_heads, 1, 1)425            pad_tgt_mask[:,:self.num_queries,:self.num_queries] = self_tgt_mask426            pad_tgt_mask[:,self.num_queries:,self.num_queries:] = False # grounding tokens could attend with eatch other427            self_tgt_mask = pad_tgt_mask428            output = torch.cat((output, output[:-1]), dim=0)429            query_embed = torch.cat((query_embed, query_embed[:-1]), dim=0) # also pad language embdding to fix embedding430        else:431            self_tgt_mask = self.self_attn_mask[:,:self.num_queries,:self.num_queries].repeat(output.shape[1]*self.num_heads, 1, 1)432 433        # prediction heads on learnable query features434        results = self.forward_prediction_heads(output, mask_features, attn_mask_target_size=size_list[0], task=task)435        attn_mask = results["attn_mask"]436        predictions_class.append(results["outputs_class"])437        predictions_mask.append(results["outputs_mask"])438        predictions_bbox.append(results["outputs_bbox"])439        predictions_caption.append(results["outputs_caption"])440        predictions_captioning.append(results["outputs_captionting"])441        442        for i in range(self.num_layers):443            level_index = i % self.num_feature_levels444            attn_mask[torch.where(attn_mask.sum(-1) == attn_mask.shape[-1])] = False445 446            if self.training and task == 'vlp' and self.task_switch['captioning']:447                attn_mask = torch.cat((attn_mask, torch.zeros_like(attn_mask[:, :self.contxt_len, :])), dim=1)448            # attention: cross-attention first449            output, avg_attn = self.transformer_cross_attention_layers[i](450                output, src[level_index],451                memory_mask=attn_mask,452                memory_key_padding_mask=None,  # here we do not apply masking on padded region453                pos=pos[level_index], query_pos=query_embed454            )455 456            if (((self.training and task == 'seg') or (task == 'grounding_eval')) and self.task_switch['grounding']) \457                    or ((self.training and task == 'openimage') and self.task_switch['openimage']['grounding']):458                output = torch.cat((output, _grounding_tokens), dim=0)459                query_embed = torch.cat((query_embed, grounding_tokens), dim=0)460 461            output = self.transformer_self_attention_layers[i](462                output, tgt_mask=self_tgt_mask,463                tgt_key_padding_mask=None,464                query_pos=query_embed465            )466            467            # FFN468            output = self.transformer_ffn_layers[i](469                output470            )471 472            if ((self.training and task == 'seg') or (task == 'grounding_eval')) and self.task_switch['grounding'] \473                    or ((self.training and task == 'openimage') and self.task_switch['openimage']['grounding']):474                _grounding_tokens = output[-len(_grounding_tokens):]475                output = output[:-len(_grounding_tokens)]476                query_embed = query_embed[:-len(_grounding_tokens)]477 478            results = self.forward_prediction_heads(output, mask_features, attn_mask_target_size=size_list[(i + 1) % self.num_feature_levels], layer_id=i, task=task)479            attn_mask = results["attn_mask"]480            predictions_class.append(results["outputs_class"])481            predictions_mask.append(results["outputs_mask"])482            predictions_bbox.append(results["outputs_bbox"])483            predictions_caption.append(results["outputs_caption"])484            predictions_captioning.append(results["outputs_captionting"])485 486        assert len(predictions_class) == self.num_layers + 1487        if task == 'vlp':488            out = {'pred_captionings': predictions_captioning[-1], 489                   'pred_captions': predictions_caption[-1], 490                   'aux_outputs': [{'pred_captionings': x, 'pred_captions': y } for x, y in zip(predictions_captioning[:-1], predictions_caption[:-1])]}491            return out492        else:493            out = {494                'pred_logits': predictions_class[-1],495                'pred_masks': predictions_mask[-1],496                'pred_boxes': predictions_bbox[-1],497                'pred_captions': predictions_caption[-1],498                'aux_outputs': self._set_aux_loss(499                    predictions_class if self.mask_classification else None, predictions_mask, predictions_bbox, predictions_caption500                )501            }502            return out503 504    def forward_captioning(self, x, mask_features, mask = None, target_queries = None, target_vlp = None, task='seg', extra={}):505        # x is a list of multi-scale feature506        assert len(x) == self.num_feature_levels507        src = []508        pos = []509        size_list = []510        511        # disable mask, it does not affect performance512        del mask513        for i in range(self.num_feature_levels):514            size_list.append(x[i].shape[-2:])515            pos.append(self.pe_layer(x[i], None).flatten(2))516            src.append(self.input_proj[i](x[i]).flatten(2) + self.level_embed.weight[i][None, :, None])517 518            # flatten NxCxHxW to HWxNxC519            pos[-1] = pos[-1].permute(2, 0, 1)520            src[-1] = src[-1].permute(2, 0, 1)521 522        _, bs, _ = src[0].shape523 524        # QxNxC525        query_embed_ = self.query_embed.weight.unsqueeze(1).repeat(1, bs, 1)526        query_feat = self.query_feat.weight.unsqueeze(1).repeat(1, bs, 1)        527        caping_lang_token = extra['start_token'].repeat(bs, 1)528        query_feat_caping = self.query_feat_caping.weight.unsqueeze(1).repeat(1, bs, 1)529        530        # prepare token embedding for evaluation531        token_embs = self.lang_encoder.lang_encoder.token_embedding.weight532        # token_embs = (token_embs / token_embs.norm(dim=-1, keepdim=True) + 1e-7)533        534        for cap_idx in range(0, self.captioning_step):535            caping_lang_embed = self.lang_encoder.forward_language_token((caping_lang_token,))[0].transpose(0, 1)536            query_embed = torch.cat((query_embed_, caping_lang_embed), dim=0) # may not add at the beginning.537            output = torch.cat((query_feat, query_feat_caping), dim=0) # concat object query, class token and caption token.538 539            # prediction heads on learnable query features540            results = self.forward_prediction_heads(output, mask_features, attn_mask_target_size=size_list[0], task=task)541            attn_mask = results["attn_mask"]542        543            for i in range(self.num_layers):544                level_index = i % self.num_feature_levels545                attn_mask[torch.where(attn_mask.sum(-1) == attn_mask.shape[-1])] = False546                attn_mask = torch.cat((attn_mask, torch.zeros_like(attn_mask[:, :self.contxt_len, :])), dim=1)547                self_tgt_mask = self.self_attn_mask.repeat(output.shape[1]*self.num_heads, 1, 1)548                549                # attention: cross-attention first550                output, avg_attn = self.transformer_cross_attention_layers[i](551                    output, src[level_index],552                    memory_mask=attn_mask,553                    memory_key_padding_mask=None,  # here we do not apply masking on padded region554                    pos=pos[level_index], query_pos=query_embed555                )556 557                output = self.transformer_self_attention_layers[i](558                    output, tgt_mask=self_tgt_mask,559                    tgt_key_padding_mask=None,560                    query_pos=query_embed561                )562                563                # FFN564                output = self.transformer_ffn_layers[i](565                    output566                )567 568                results = self.forward_prediction_heads(output, mask_features, attn_mask_target_size=size_list[(i + 1) % self.num_feature_levels], layer_id=i, task=task)569                attn_mask = results["attn_mask"]570            571            pred_captions_gen = results['outputs_captionting']572            # pred_captions_gen = (pred_captions_gen / pred_captions_gen.norm(dim=-1, keepdim=True) + 1e-7)573            pred_captions_gen = pred_captions_gen @ token_embs.t()574            caping_lang_token[:,cap_idx+1] = pred_captions_gen[:,cap_idx].max(-1)[1]575 576        out = {'pred_captionings': caping_lang_token,577               'pred_texts': self.lang_encoder.tokenizer.batch_decode(caping_lang_token, skip_special_tokens=True)}578        return out579 580 581    def forward_prediction_heads(self, output, mask_features, attn_mask_target_size, layer_id=-1, task='seg'):582        decoder_output = self.decoder_norm(output)583        decoder_output = decoder_output.transpose(0, 1)584 585        # extract image captioning token from decoder output.586        if self.task_switch['captioning'] and (task == 'vlp' or task == 'captioning_infer'):587            outputs_captionting = decoder_output[:,self.num_queries:] @ self.caping_embed588        else:589            outputs_captionting = None590 591        # recompute class token output.592        norm_decoder_output = decoder_output / (decoder_output.norm(dim=-1, keepdim=True) + 1e-7)593        obj_token = norm_decoder_output[:,:self.num_queries-1]594        cls_token = norm_decoder_output[:,self.num_queries-1:self.num_queries]595 596        sim = (cls_token @ obj_token.transpose(1,2)).softmax(-1)[:,0,:,None] # TODO include class token.597        cls_token = (sim * decoder_output[:,:self.num_queries-1]).sum(dim=1, keepdim=True)598 599        if (((self.training and task == 'seg') or (task == 'grounding_eval')) and self.task_switch['grounding']) \600                or ((self.training and task == 'openimage') and self.task_switch['openimage']['grounding']):601            decoder_output = torch.cat((decoder_output[:,:self.num_queries-1], cls_token, decoder_output[:,self.num_queries:2*self.num_queries-1]), dim=1)602        else:603            decoder_output = torch.cat((decoder_output[:,:self.num_queries-1], cls_token), dim=1)604 605        # compute class, mask and bbox.606        class_embed = decoder_output @ self.class_embed607        # HACK do not compute similarity if mask is not on608        outputs_class = self.lang_encoder.compute_similarity(class_embed, fake=(((not self.task_switch['mask']) and self.training) or (task == 'openimage')))609 610        if self.task_switch['mask'] or self.task_switch['openimage']['mask']:611            mask_embed = self.mask_embed(decoder_output)612            outputs_mask = torch.einsum("bqc,bchw->bqhw", mask_embed, mask_features)613 614            # NOTE: prediction is of higher-resolution615            # [B, Q, H, W] -> [B, Q, H*W] -> [B, h, Q, H*W] -> [B*h, Q, HW]616            attn_mask = F.interpolate(outputs_mask, size=attn_mask_target_size, mode="bilinear", align_corners=False)617 618            # must use bool type619            # If a BoolTensor is provided, positions with ``True`` are not allowed to attend while ``False`` values will be unchanged.620            attn_mask = (attn_mask.sigmoid().flatten(2).unsqueeze(1).repeat(1, self.num_heads, 1, 1).flatten(0, 1) < 0.5).bool()621            attn_mask = attn_mask.detach()622 623            # NOTE: fill False for cls token (JY)624            attn_mask[:, self.num_queries:self.num_queries+1].fill_(False)625        else:626            outputs_mask = None627            attn_mask = torch.zeros((list(decoder_output.shape[:2]) + [attn_mask_target_size[0]*attn_mask_target_size[1]]), device=decoder_output.device).repeat(self.num_heads, 1, 1).bool()628 629        outputs_bbox = [None for i in range(len(decoder_output))]630        if self.task_switch['bbox']:631            outputs_bbox = self.bbox_embed(decoder_output)632 633        outputs_caption = None634        if self.task_switch['caption']:635            outputs_caption = class_embed636            637 638        results = {639            "outputs_class": outputs_class,640            "outputs_mask": outputs_mask,641            "outputs_bbox": outputs_bbox,642            "attn_mask": attn_mask,643            "outputs_caption": outputs_caption,644            "outputs_captionting": outputs_captionting,645        }646        return results647 648    @torch.jit.unused649    def _set_aux_loss(self, outputs_class, outputs_seg_masks, outputs_boxes, outputs_captions):650        # this is a workaround to make torchscript happy, as torchscript651        # doesn't support dictionary with non-homogeneous values, such652        # as a dict having both a Tensor and a list.653        if self.mask_classification:654            return [655                {"pred_logits": a, "pred_masks": b, "pred_boxes": c, "pred_captions": d}656                for a, b, c, d in zip(outputs_class[:-1], outputs_seg_masks[:-1], outputs_boxes[:-1], outputs_captions[:-1])657            ]658        else:659            return [{"pred_masks": b} for b in outputs_seg_masks[:-1]]660 661 662@register_decoder663def get_masked_transformer_decoder(cfg, in_channels, lang_encoder, mask_classification, extra):664    return MultiScaleMaskedTransformerDecoder(cfg, in_channels, lang_encoder, mask_classification, extra)