(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)
| 168 | class 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 |
nothing calls this directly
no test coverage detected