MCPcopy Create free account
hub / github.com/UX-Decoder/Semantic-SAM / TransformerDecoder

Class TransformerDecoder

semantic_sam/body/transformer_blocks.py:105–151  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

103
104
105class TransformerDecoder(nn.Module):
106 def __init__(self, decoder_layer, num_layers, norm=None, return_intermediate=False):
107 super().__init__()
108 self.layers = _get_clones(decoder_layer, num_layers)
109 self.num_layers = num_layers
110 self.norm = norm
111 self.return_intermediate = return_intermediate
112
113 def forward(
114 self,
115 tgt,
116 memory,
117 tgt_mask: Optional[Tensor] = None,
118 memory_mask: Optional[Tensor] = None,
119 tgt_key_padding_mask: Optional[Tensor] = None,
120 memory_key_padding_mask: Optional[Tensor] = None,
121 pos: Optional[Tensor] = None,
122 query_pos: Optional[Tensor] = None,
123 ):
124 output = tgt
125
126 intermediate = []
127
128 for layer in self.layers:
129 output = layer(
130 output,
131 memory,
132 tgt_mask=tgt_mask,
133 memory_mask=memory_mask,
134 tgt_key_padding_mask=tgt_key_padding_mask,
135 memory_key_padding_mask=memory_key_padding_mask,
136 pos=pos,
137 query_pos=query_pos,
138 )
139 if self.return_intermediate:
140 intermediate.append(self.norm(output))
141
142 if self.norm is not None:
143 output = self.norm(output)
144 if self.return_intermediate:
145 intermediate.pop()
146 intermediate.append(output)
147
148 if self.return_intermediate:
149 return torch.stack(intermediate)
150
151 return output.unsqueeze(0)
152
153
154class TransformerEncoderLayer(nn.Module):

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected