overwrite the original forward_decoder, now input latent are fused
(self, latent, mask, ids_restore)
| 357 | return x, mask, ids_restore |
| 358 | |
| 359 | def forward_decoder(self, latent, mask, ids_restore): |
| 360 | """ |
| 361 | overwrite the original forward_decoder, now input latent are fused |
| 362 | """ |
| 363 | # print("to decoder, latent", latent.size()) |
| 364 | latent = self.decompressor(latent) |
| 365 | # # embed tokens |
| 366 | latent = self.decoder_embed(latent) |
| 367 | |
| 368 | restored_latent = self.unmasking_handle(latent, mask, ids_restore) |
| 369 | restored_latent = restored_latent.reshape(restored_latent.size(0), self.patch_h*self.patch_w, -1) |
| 370 | # x = torch.cat([latent_cls, restored_latent], dim=1) |
| 371 | |
| 372 | # add pos embed |
| 373 | x = restored_latent + self.decoder_pos_embed[:, 1:, :] |
| 374 | |
| 375 | # apply Transformer blocks |
| 376 | for blk in self.decoder_blocks: |
| 377 | x = blk(x) |
| 378 | x = self.decoder_norm(x) |
| 379 | |
| 380 | # VERSION mlp ----- |
| 381 | x_occ = self.decoder_pred_occ(x) |
| 382 | x_free = self.decoder_pred_free(x) |
| 383 | x_occ = self.unpatchify(x_occ) |
| 384 | x_free = self.unpatchify(x_free) |
| 385 | # remove cls token |
| 386 | # x = x[:, 1:, :] |
| 387 | # print("pred size", x.size()) |
| 388 | # -------------------- |
| 389 | # x_pred = x_occ |
| 390 | # print("x_pred", x_pred.size()) |
| 391 | x_pred = torch.stack((x_free, x_occ), dim=1) # [B, class, C, H, W] |
| 392 | return x_pred |
| 393 | |
| 394 | def forward_loss(self, teacher, pred): |
| 395 | """ |
no test coverage detected