CoolFace
Apppublic

xdecoder/Instruct-X-Decoder

sourceHugging Faceafl-3.0updated 3y agoView on Hugging Face
163likes
xdecoder_head.py123 linesDownload Raw Back to body
1# Copyright (c) Facebook, Inc. and its affiliates.2 3# --------------------------------------------------------4# X-Decoder -- Generalized Decoding for Pixel, Image, and Language5# Copyright (c) 2022 Microsoft6# Licensed under The MIT License [see LICENSE for details]7# Written by Jianwei Yang (jianwyan@microsoft.com), Xueyan Zou (xueyan@cs.wisc.edu)8# --------------------------------------------------------9 10from typing import Dict11 12from torch import nn13 14from detectron2.layers import ShapeSpec15 16from .registry import register_body17from .encoder import build_encoder18from .decoder import build_decoder19from ..utils import configurable20 21 22class XDecoderHead(nn.Module):23 24    @configurable25    def __init__(26        self,27        input_shape: Dict[str, ShapeSpec],28        *,29        num_classes: int,30        pixel_decoder: nn.Module,31        loss_weight: float = 1.0,32        ignore_value: int = -1,33        # extra parameters34        transformer_predictor: nn.Module,35        transformer_in_feature: str,36    ):37        """38        NOTE: this interface is experimental.39        Args:40            input_shape: shapes (channels and stride) of the input features41            num_classes: number of classes to predict42            pixel_decoder: the pixel decoder module43            loss_weight: loss weight44            ignore_value: category id to be ignored during training.45            transformer_predictor: the transformer decoder that makes prediction46            transformer_in_feature: input feature name to the transformer_predictor47        """48        super().__init__()49 50        input_shape = sorted(input_shape.items(), key=lambda x: x[1].stride)51        self.in_features = [k for k, v in input_shape]52        feature_strides = [v.stride for k, v in input_shape]53        feature_channels = [v.channels for k, v in input_shape]54 55        self.ignore_value = ignore_value56        self.common_stride = 457        self.loss_weight = loss_weight58 59        self.pixel_decoder = pixel_decoder60        self.predictor = transformer_predictor61        self.transformer_in_feature = transformer_in_feature62 63        self.num_classes = num_classes64 65    @classmethod66    def from_config(cls, cfg, input_shape: Dict[str, ShapeSpec], lang_encoder: nn.Module, extra: dict):67 68        in_features_type = cfg['MODEL']['DECODER']['TRANSFORMER_IN_FEATURE']69        enc_cfg = cfg['MODEL']['ENCODER']70        dec_cfg = cfg['MODEL']['DECODER']71 72        # figure out in_channels to transformer predictor73        if in_features_type == "transformer_encoder":74            transformer_predictor_in_channels = enc_cfg['CONVS_DIM']75        elif in_features_type == "pixel_embedding":76            transformer_predictor_in_channels = enc_cfg['MASK_DIM']77        elif in_features_type == "multi_scale_pixel_decoder":  # for maskformer278            transformer_predictor_in_channels = enc_cfg['CONVS_DIM']79        else:80            transformer_predictor_in_channels = input_shape[dec_cfg['TRANSFORMER_IN_FEATURE']].channels81 82        return {83            "input_shape": {84                k: v for k, v in input_shape.items() if k in enc_cfg['IN_FEATURES']85            },86            "ignore_value": enc_cfg['IGNORE_VALUE'],87            "num_classes": enc_cfg.get('NUM_CLASSES', None),88            "pixel_decoder": build_encoder(cfg, input_shape),89            "loss_weight": enc_cfg['LOSS_WEIGHT'],90            "transformer_in_feature": dec_cfg['TRANSFORMER_IN_FEATURE'],91            "transformer_predictor": build_decoder(92                cfg,93                transformer_predictor_in_channels,94                lang_encoder,95                mask_classification=True,96                extra=extra,97            ),98        }99 100    def forward(self, features, mask=None, target_queries=None, target_vlp=None, task='seg', extra={}):101        return self.layers(features, mask, target_queries, target_vlp, task, extra)102 103    def layers(self, features, mask=None, target_queries=None, target_vlp=None, task='seg', extra={}):104        mask_features, transformer_encoder_features, multi_scale_features = self.pixel_decoder.forward_features(features)105        106        if self.transformer_in_feature == "multi_scale_pixel_decoder":107            predictions = self.predictor(multi_scale_features, mask_features, mask, target_queries, target_vlp, task, extra)108        else:109            if self.transformer_in_feature == "transformer_encoder":110                assert (111                    transformer_encoder_features is not None112                ), "Please use the TransformerEncoderPixelDecoder."113                predictions = self.predictor(transformer_encoder_features, mask_features, mask)114            elif self.transformer_in_feature == "pixel_embedding":115                predictions = self.predictor(mask_features, mask_features, mask)116            else:117                predictions = self.predictor(features[self.transformer_in_feature], mask_features, mask)118        return predictions119 120 121@register_body122def get_xdecoder_head(cfg, input_shape, lang_encoder, extra):123    return XDecoderHead(cfg, input_shape, lang_encoder, extra)