MCPcopy Create free account
hub / github.com/csxmli2016/MARCONetPlusPlus / Decoder

Class Decoder

networks/transocr_arch.py:282–310  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

280 return x
281
282class Decoder(nn.Module):
283
284 def __init__(self):
285 super(Decoder, self).__init__()
286
287 self.mask_multihead = MultiHeadedAttention(h=4, d_model=1024, dropout=0.1)
288 self.mul_layernorm1 = LayerNorm(features=1024)
289
290 self.multihead = MultiHeadedAttention(h=4, d_model=1024, dropout=0.1, compress_attention=False)
291 self.mul_layernorm2 = LayerNorm(features=1024)
292
293 self.pff = PositionwiseFeedForward(1024, 2048)
294 self.mul_layernorm3 = LayerNorm(features=1024)
295
296 def forward(self, text, conv_feature):
297 text_max_length = text.shape[1]
298 mask = subsequent_mask(text_max_length).cuda()
299
300 result = text
301 result = self.mul_layernorm1(result + self.mask_multihead(result, result, result, mask=mask)[0])
302
303 b, c, h, w = conv_feature.shape
304 conv_feature = conv_feature.view(b, c, h * w).permute(0, 2, 1).contiguous()
305 word_image_align, attention_map = self.multihead(result, conv_feature, conv_feature, mask=None)
306 result = self.mul_layernorm2(result + word_image_align)
307
308 result = self.mul_layernorm3(result + self.pff(result))
309
310 return result, attention_map
311
312
313class TransformerOCR(nn.Module):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected