MCPcopy Create free account
hub / github.com/Topdu/OpenOCR / ARDecoder

Class ARDecoder

openrec/modeling/decoders/ote_decoder.py:31–153  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

29
30
31class ARDecoder(nn.Module):
32
33 def __init__(
34 self,
35 in_channels,
36 out_channels,
37 nhead=None,
38 num_decoder_layers=6,
39 max_len=25,
40 attention_dropout_rate=0.0,
41 residual_dropout_rate=0.1,
42 scale_embedding=True,
43 ):
44 super(ARDecoder, self).__init__()
45 self.out_channels = out_channels
46 self.ignore_index = out_channels - 1
47 self.bos = out_channels - 2
48 self.eos = 0
49 self.max_len = max_len
50 d_model = in_channels
51 dim_feedforward = d_model * 4
52 nhead = nhead if nhead is not None else d_model // 32
53 self.embedding = Embeddings(
54 d_model=d_model,
55 vocab=self.out_channels,
56 padding_idx=0,
57 scale_embedding=scale_embedding,
58 )
59 self.pos_embed = nn.Parameter(torch.zeros([1, max_len + 1, d_model],
60 dtype=torch.float32),
61 requires_grad=True)
62 trunc_normal_(self.pos_embed, std=0.02)
63 self.decoder = nn.ModuleList([
64 TransformerBlock(
65 d_model,
66 nhead,
67 dim_feedforward,
68 attention_dropout_rate,
69 residual_dropout_rate,
70 with_self_attn=True,
71 with_cross_attn=False,
72 ) for i in range(num_decoder_layers)
73 ])
74
75 self.tgt_word_prj = nn.Linear(d_model,
76 self.out_channels - 2,
77 bias=False)
78 self.apply(self._init_weights)
79
80 def _init_weights(self, m):
81 if isinstance(m, nn.Linear):
82 nn.init.xavier_normal_(m.weight)
83 if m.bias is not None:
84 nn.init.zeros_(m.bias)
85
86 def forward_train(self, src, tgt):
87 tgt = tgt[:, :-1]
88

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected