| 280 | return x |
| 281 | |
| 282 | class 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 | |
| 313 | class TransformerOCR(nn.Module): |