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

Method __init__

openrec/modeling/decoders/mdiff_decoder.py:28–104  ·  view source on GitHub ↗
(self,
                 in_channels,
                 out_channels,
                 nhead=None,
                 num_decoder_layers=6,
                 max_len=25,
                 attention_dropout_rate=0.0,
                 residual_dropout_rate=0.1,
                 scale_embedding=True,
                 parallel_decoding=False,
                 autoregressive_decoding=False,
                 sampler_step=5,
                 low_confidence_decoding=False,
                 random_mask_decoding=False,
                 semi_autoregressive_decoding=False,
                 cloze_mask_decoding=False,
                 rec_loss_weight=1.0,
                 reflect_loss_weight=1.0,
                 sample_k=0,
                 temperature=1.0)

Source from the content-addressed store, hash-verified

26 """
27
28 def __init__(self,
29 in_channels,
30 out_channels,
31 nhead=None,
32 num_decoder_layers=6,
33 max_len=25,
34 attention_dropout_rate=0.0,
35 residual_dropout_rate=0.1,
36 scale_embedding=True,
37 parallel_decoding=False,
38 autoregressive_decoding=False,
39 sampler_step=5,
40 low_confidence_decoding=False,
41 random_mask_decoding=False,
42 semi_autoregressive_decoding=False,
43 cloze_mask_decoding=False,
44 rec_loss_weight=1.0,
45 reflect_loss_weight=1.0,
46 sample_k=0,
47 temperature=1.0):
48 super(MDiffDecoder, self).__init__()
49 self.out_channels = out_channels
50 self.ignore_index = out_channels - 1
51 self.mask_token_id = out_channels - 2
52 self.eos = 0
53 self.max_len = max_len
54 d_model = in_channels
55 dim_feedforward = d_model * 4
56 self.pd = parallel_decoding
57 self.ar = autoregressive_decoding
58 self.sampler_step = sampler_step
59 self.lc = low_confidence_decoding
60 self.rm = random_mask_decoding
61 self.semiar = semi_autoregressive_decoding
62 self.cm = cloze_mask_decoding
63 self.rec_loss_weight = rec_loss_weight
64 self.reflect_loss_weight = reflect_loss_weight
65 self.temperature = temperature
66 self.sample_k = sample_k
67 nhead = nhead if nhead is not None else d_model // 32
68 self.embedding = Embeddings(
69 d_model=d_model,
70 vocab=self.out_channels,
71 padding_idx=0,
72 scale_embedding=scale_embedding,
73 )
74 self.pos_embed = nn.Parameter(torch.zeros(
75 [1, self.max_len + 1, d_model], dtype=torch.float32),
76 requires_grad=True)
77 nn.init.trunc_normal_(self.pos_embed, std=0.02)
78 self.positional_encoding = PositionalEncoding(
79 dropout=residual_dropout_rate, dim=d_model)
80
81 self.decoder = nn.ModuleList([
82 TransformerBlock(
83 d_model,
84 nhead,
85 dim_feedforward,

Callers 1

__init__Method · 0.45

Calls 4

EmbeddingsClass · 0.90
PositionalEncodingClass · 0.90
TransformerBlockClass · 0.70
applyMethod · 0.45

Tested by

no test coverage detected