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

Method __init__

openrec/modeling/decoders/cdistnet_decoder.py:196–253  ·  view source on GitHub ↗
(self,
                 in_channels,
                 out_channels,
                 n_head=None,
                 num_encoder_blocks=3,
                 num_decoder_blocks=3,
                 beam_size=0,
                 max_len=25,
                 residual_dropout_rate=0.1,
                 add_conv=False,
                 **kwargs)

Source from the content-addressed store, hash-verified

194class CDistNetDecoder(nn.Module):
195
196 def __init__(self,
197 in_channels,
198 out_channels,
199 n_head=None,
200 num_encoder_blocks=3,
201 num_decoder_blocks=3,
202 beam_size=0,
203 max_len=25,
204 residual_dropout_rate=0.1,
205 add_conv=False,
206 **kwargs):
207 super(CDistNetDecoder, self).__init__()
208 dst_vocab_size = out_channels
209 self.ignore_index = dst_vocab_size - 1
210 self.bos = dst_vocab_size - 2
211 self.eos = 0
212 self.beam_size = beam_size
213 self.max_len = max_len
214 self.add_conv = add_conv
215 d_model = in_channels
216 dim_feedforward = d_model * 4
217 n_head = n_head if n_head is not None else d_model // 32
218
219 if add_conv:
220 self.convbnrelu = ConvBnRelu(
221 in_channels=in_channels,
222 out_channels=in_channels,
223 kernel_size=(1, 3),
224 stride=(1, 2),
225 )
226 if num_encoder_blocks > 0:
227 self.positional_encoding = PositionalEncoding(
228 dropout=0.1,
229 dim=d_model,
230 )
231 self.trans_encoder = Transformer_Encoder(
232 n_layers=num_encoder_blocks,
233 n_head=n_head,
234 d_model=d_model,
235 d_inner=dim_feedforward,
236 )
237 else:
238 self.trans_encoder = None
239 self.semantic_branch = SEM_Pre(
240 d_model=d_model,
241 dst_vocab_size=dst_vocab_size,
242 residual_dropout_rate=residual_dropout_rate,
243 )
244 self.positional_branch = POS_Pre(d_model=d_model)
245
246 self.mdcdp = MDCDP(d_model, n_head, dim_feedforward // 2,
247 num_decoder_blocks)
248 self._reset_parameters()
249
250 self.tgt_word_prj = nn.Linear(
251 d_model, dst_vocab_size - 2,
252 bias=False) # We don't predict <bos> nor <pad>
253 self.tgt_word_prj.weight.data.normal_(mean=0.0, std=d_model**-0.5)

Callers

nothing calls this directly

Calls 8

_reset_parametersMethod · 0.95
PositionalEncodingClass · 0.90
Transformer_EncoderClass · 0.90
ConvBnReluClass · 0.85
SEM_PreClass · 0.85
POS_PreClass · 0.85
MDCDPClass · 0.85
__init__Method · 0.45

Tested by

no test coverage detected