| 70 | |
| 71 | |
| 72 | class DeformableTransformerDecoderLayer(nn.Module): |
| 73 | def __init__( |
| 74 | self, |
| 75 | d_model=256, |
| 76 | d_ffn=1024, |
| 77 | dropout=0.1, |
| 78 | activation='relu', |
| 79 | n_levels=4, |
| 80 | n_heads=8, |
| 81 | n_points=4, |
| 82 | decoder_sa_type='ca', |
| 83 | module_seq=['sa', 'ca', 'ffn'], |
| 84 | ): |
| 85 | super().__init__() |
| 86 | # pdb.set_trace() |
| 87 | self.module_seq = module_seq |
| 88 | assert sorted(module_seq) == ['ca', 'ffn', 'sa'] |
| 89 | |
| 90 | # cross attention |
| 91 | self.cross_attn = MSDeformAttn(d_model, n_levels, n_heads, n_points) |
| 92 | self.dropout1 = nn.Dropout(dropout) |
| 93 | self.norm1 = nn.LayerNorm(d_model) |
| 94 | |
| 95 | # self attention |
| 96 | self.self_attn = nn.MultiheadAttention(d_model, |
| 97 | n_heads, |
| 98 | dropout=dropout) |
| 99 | self.dropout2 = nn.Dropout(dropout) |
| 100 | self.norm2 = nn.LayerNorm(d_model) |
| 101 | |
| 102 | # ffn |
| 103 | self.linear1 = nn.Linear(d_model, d_ffn) |
| 104 | self.activation = _get_activation_fn(activation, |
| 105 | d_model=d_ffn, |
| 106 | batch_dim=1) |
| 107 | self.dropout3 = nn.Dropout(dropout) |
| 108 | self.linear2 = nn.Linear(d_ffn, d_model) |
| 109 | self.dropout4 = nn.Dropout(dropout) |
| 110 | self.norm3 = nn.LayerNorm(d_model) |
| 111 | |
| 112 | self.key_aware_proj = None |
| 113 | self.decoder_sa_type = decoder_sa_type |
| 114 | assert decoder_sa_type in ['sa'] |
| 115 | |
| 116 | def rm_self_attn_modules(self): |
| 117 | self.self_attn = None |
| 118 | self.dropout2 = None |
| 119 | self.norm2 = None |
| 120 | |
| 121 | @staticmethod |
| 122 | def with_pos_embed(tensor, pos): |
| 123 | return tensor if pos is None else tensor + pos |
| 124 | |
| 125 | def forward_ffn(self, tgt): |
| 126 | tgt2 = self.linear2(self.dropout3(self.activation(self.linear1(tgt)))) |
| 127 | tgt = tgt + self.dropout4(tgt2) |
| 128 | tgt = self.norm3(tgt) |
| 129 | return tgt |