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

Method __init__

openrec/modeling/decoders/parseq_decoder.py:170–215  ·  view source on GitHub ↗
(self,
                 in_channels,
                 out_channels,
                 max_label_length=25,
                 embed_dim=384,
                 dec_num_heads=12,
                 dec_mlp_ratio=4,
                 dec_depth=1,
                 perm_num=6,
                 perm_forward=True,
                 perm_mirrored=True,
                 decode_ar=True,
                 refine_iters=1,
                 dropout=0.1,
                 **kwargs: Any)

Source from the content-addressed store, hash-verified

168class PARSeqDecoder(nn.Module):
169
170 def __init__(self,
171 in_channels,
172 out_channels,
173 max_label_length=25,
174 embed_dim=384,
175 dec_num_heads=12,
176 dec_mlp_ratio=4,
177 dec_depth=1,
178 perm_num=6,
179 perm_forward=True,
180 perm_mirrored=True,
181 decode_ar=True,
182 refine_iters=1,
183 dropout=0.1,
184 **kwargs: Any) -> None:
185 super().__init__()
186 self.pad_id = out_channels - 1
187 self.eos_id = 0
188 self.bos_id = out_channels - 2
189 self.max_label_length = max_label_length
190 self.decode_ar = decode_ar
191 self.refine_iters = refine_iters
192
193 decoder_layer = DecoderLayer(embed_dim, dec_num_heads,
194 embed_dim * dec_mlp_ratio, dropout)
195 self.decoder = Decoder(decoder_layer,
196 num_layers=dec_depth,
197 norm=nn.LayerNorm(embed_dim))
198
199 # Perm/attn mask stuff
200 self.rng = np.random.default_rng()
201 self.max_gen_perms = perm_num // 2 if perm_mirrored else perm_num
202 self.perm_forward = perm_forward
203 self.perm_mirrored = perm_mirrored
204
205 # We don't predict <bos> nor <pad>
206 self.head = nn.Linear(embed_dim, out_channels - 2)
207 self.text_embed = TokenEmbedding(out_channels, embed_dim)
208
209 # +1 for <eos>
210 self.pos_queries = nn.Parameter(
211 torch.Tensor(1, max_label_length + 1, embed_dim))
212 self.dropout = nn.Dropout(p=dropout)
213 # Encoder has its own init.
214 self.apply(self._init_weights)
215 nn.init.trunc_normal_(self.pos_queries, std=0.02)
216
217 def _init_weights(self, module: nn.Module):
218 """Initialize the weights using the typical initialization schemes used

Callers

nothing calls this directly

Calls 5

DecoderLayerClass · 0.70
DecoderClass · 0.70
TokenEmbeddingClass · 0.70
__init__Method · 0.45
applyMethod · 0.45

Tested by

no test coverage detected