(self,
in_channels,
out_channels,
max_label_length=25,
embed_dim=512,
dec_num_heads=8,
dec_mlp_ratio=4,
dec_depth=6,
perm_num=6,
perm_forward=True,
perm_mirrored=True,
decode_ar=True,
refine_iters=1,
dropout=0.1,
is_pretrain=True,
ORP_path=None,
**kwargs: Any)
| 1153 | class DptrParseq(nn.Module): |
| 1154 | |
| 1155 | def __init__(self, |
| 1156 | in_channels, |
| 1157 | out_channels, |
| 1158 | max_label_length=25, |
| 1159 | embed_dim=512, |
| 1160 | dec_num_heads=8, |
| 1161 | dec_mlp_ratio=4, |
| 1162 | dec_depth=6, |
| 1163 | perm_num=6, |
| 1164 | perm_forward=True, |
| 1165 | perm_mirrored=True, |
| 1166 | decode_ar=True, |
| 1167 | refine_iters=1, |
| 1168 | dropout=0.1, |
| 1169 | is_pretrain=True, |
| 1170 | ORP_path=None, |
| 1171 | **kwargs: Any) -> None: |
| 1172 | super().__init__() |
| 1173 | self.pad_id = out_channels - 1 |
| 1174 | self.eos_id = 0 |
| 1175 | self.bos_id = out_channels - 2 |
| 1176 | self.max_label_length = max_label_length |
| 1177 | self.decode_ar = decode_ar |
| 1178 | self.refine_iters = refine_iters |
| 1179 | self.is_pretrain = is_pretrain |
| 1180 | if not is_pretrain: |
| 1181 | self.token_query = nn.Parameter(torch.Tensor(1, 26, embed_dim)) |
| 1182 | self.fmu = FMU(embed_dim, dec_num_heads, embed_dim * dec_mlp_ratio, |
| 1183 | dropout) |
| 1184 | |
| 1185 | decoder_layer = DecoderLayer(embed_dim, dec_num_heads, |
| 1186 | embed_dim * dec_mlp_ratio, dropout) |
| 1187 | self.decoder = Decoder(decoder_layer, |
| 1188 | num_layers=dec_depth, |
| 1189 | norm=nn.LayerNorm(embed_dim)) |
| 1190 | |
| 1191 | # Perm/attn mask stuff |
| 1192 | self.rng = np.random.default_rng() |
| 1193 | self.max_gen_perms = perm_num // 2 if perm_mirrored else perm_num |
| 1194 | self.perm_forward = perm_forward |
| 1195 | self.perm_mirrored = perm_mirrored |
| 1196 | |
| 1197 | # We don't predict <bos> nor <pad> |
| 1198 | self.head = nn.Linear(embed_dim, out_channels - 2) |
| 1199 | self.text_embed = TokenEmbedding(out_channels, embed_dim) |
| 1200 | |
| 1201 | # +1 for <eos> |
| 1202 | self.pos_queries = nn.Parameter( |
| 1203 | torch.Tensor(1, max_label_length + 1, embed_dim)) |
| 1204 | self.dropout = nn.Dropout(p=dropout) |
| 1205 | # Encoder has its own init. |
| 1206 | self.apply(self._init_weights) |
| 1207 | nn.init.trunc_normal_(self.pos_queries, std=0.02) |
| 1208 | |
| 1209 | if is_pretrain: |
| 1210 | self.clip_encoder, preprocess = load('ViT-B/16') |
| 1211 | for p in self.clip_encoder.parameters(): |
| 1212 | p.requires_grad = False |
nothing calls this directly
no test coverage detected