(self, imgs, class_label,
gen_image=False, bsz=None, num_iter=None, choice_temperature=None,
sampled_rep=None, rdm_steps=250, eta=1.0, cfg=0.0, class_label_gen=None)
| 449 | return loss |
| 450 | |
| 451 | def forward(self, imgs, class_label, |
| 452 | gen_image=False, bsz=None, num_iter=None, choice_temperature=None, |
| 453 | sampled_rep=None, rdm_steps=250, eta=1.0, cfg=0.0, class_label_gen=None): |
| 454 | if gen_image: |
| 455 | return self.gen_image(bsz, num_iter, choice_temperature, sampled_rep, rdm_steps, eta, cfg, class_label_gen) |
| 456 | |
| 457 | self.pretrained_encoder.eval() |
| 458 | with torch.no_grad(): |
| 459 | mean = torch.Tensor([0.485, 0.456, 0.406]).cuda().unsqueeze(0).unsqueeze(-1).unsqueeze(-1) |
| 460 | std = torch.Tensor([0.229, 0.224, 0.225]).cuda().unsqueeze(0).unsqueeze(-1).unsqueeze(-1) |
| 461 | x_normalized = (imgs - mean) / std |
| 462 | x_normalized = torch.nn.functional.interpolate(x_normalized, 224, mode='bicubic') |
| 463 | rep = self.pretrained_encoder.forward_features(x_normalized) |
| 464 | if self.pretrained_enc_withproj: |
| 465 | rep = self.pretrained_encoder.head(rep) |
| 466 | rep_std = torch.std(rep, dim=1, keepdim=True) |
| 467 | rep_mean = torch.mean(rep, dim=1, keepdim=True) |
| 468 | rep = (rep - rep_mean) / rep_std |
| 469 | |
| 470 | latent, gt_indices, token_drop_mask, token_all_mask = self.forward_encoder(imgs, rep, class_label) |
| 471 | logits = self.forward_decoder(latent, token_drop_mask, token_all_mask) |
| 472 | loss = self.forward_loss(gt_indices, logits, token_all_mask) |
| 473 | return loss, imgs, token_all_mask |
| 474 | |
| 475 | def gen_image(self, bsz, num_iter=12, choice_temperature=4.5, sampled_rep=None, rdm_steps=250, eta=1.0, |
| 476 | cfg=0.0, class_label=None): |
nothing calls this directly
no test coverage detected