| 16 | |
| 17 | |
| 18 | class TransformerDecoder(nn.Module): |
| 19 | |
| 20 | def __init__(self, decoder_layer, num_layers, norm=None, |
| 21 | return_intermediate=False, |
| 22 | d_model=256, query_dim=4, |
| 23 | modulate_hw_attn=True, |
| 24 | num_feature_levels=1, |
| 25 | deformable_decoder=True, |
| 26 | decoder_query_perturber=None, |
| 27 | dec_layer_number=None, # number of queries each layer in decoder |
| 28 | rm_dec_query_scale=True, |
| 29 | dec_layer_share=False, |
| 30 | dec_layer_dropout_prob=None, |
| 31 | ): |
| 32 | super().__init__() |
| 33 | if num_layers > 0: |
| 34 | self.layers = _get_clones(decoder_layer, num_layers, layer_share=dec_layer_share) |
| 35 | else: |
| 36 | self.layers = [] |
| 37 | self.num_layers = num_layers |
| 38 | self.norm = norm |
| 39 | self.return_intermediate = return_intermediate |
| 40 | assert return_intermediate, "support return_intermediate only" |
| 41 | self.query_dim = query_dim |
| 42 | assert query_dim in [2, 4], "query_dim should be 2/4 but {}".format(query_dim) |
| 43 | self.num_feature_levels = num_feature_levels |
| 44 | |
| 45 | self.ref_point_head = MLP(query_dim // 2 * d_model, d_model, d_model, 2) |
| 46 | if not deformable_decoder: |
| 47 | self.query_pos_sine_scale = MLP(d_model, d_model, d_model, 2) |
| 48 | else: |
| 49 | self.query_pos_sine_scale = None |
| 50 | |
| 51 | if rm_dec_query_scale: |
| 52 | self.query_scale = None |
| 53 | else: |
| 54 | raise NotImplementedError |
| 55 | self.query_scale = MLP(d_model, d_model, d_model, 2) |
| 56 | self.bbox_embed = None |
| 57 | self.class_embed = None |
| 58 | |
| 59 | self.d_model = d_model |
| 60 | self.modulate_hw_attn = modulate_hw_attn |
| 61 | self.deformable_decoder = deformable_decoder |
| 62 | |
| 63 | if not deformable_decoder and modulate_hw_attn: |
| 64 | self.ref_anchor_head = MLP(d_model, d_model, 2, 2) |
| 65 | else: |
| 66 | self.ref_anchor_head = None |
| 67 | |
| 68 | self.decoder_query_perturber = decoder_query_perturber |
| 69 | self.box_pred_damping = None |
| 70 | |
| 71 | self.dec_layer_number = dec_layer_number |
| 72 | if dec_layer_number is not None: |
| 73 | assert isinstance(dec_layer_number, list) |
| 74 | assert len(dec_layer_number) == num_layers |
| 75 | # assert dec_layer_number[0] == |