diff --git a/GanAttack.py b/GanAttack.py index 4335c6b..31a8bd6 100644 --- a/GanAttack.py +++ b/GanAttack.py @@ -76,7 +76,7 @@ def main(cfg: DictConfig) -> None: criterion = nn.CrossEntropyLoss() max_loss=nn.MarginRankingLoss(0.1) clip_loss=CLIPLoss().to(device) -# vgg_loss=VggLoss().to(device) + # vgg_loss=VggLoss().to(device) # summary(model, input_size = (3, 256, 256), batch_size = 5) # set_requires_grad(model.mlp.parameters()) for p in (model.mlp.parameters()): @@ -97,12 +97,14 @@ def main(cfg: DictConfig) -> None: # _, _, _, clean_refine_images, clean_latent_codes, _=inverter(inputs,img_path) optimizer.zero_grad() generated_img,adv_latent_codes=model(inputs) - loss_vgg=vgg_loss(inputs,generated_img) -# loss_l1=F.l1_loss(codes,adv_latent_codes) + # loss_vgg=vgg_loss(inputs,generated_img) + loss_l1=F.l1_loss(codes,adv_latent_codes) loss_clip=clip_loss(generated_img,prompt) - outputs = classifier(generated_img) - preds=criterion(outputs,labels) - loss_classifier=max_loss(torch.ones_like(preds),preds,-torch.ones_like(preds)) + adv_outputs = classifier(generated_img) + clean_outputs=classifier(inputs) + adv_preds=criterion(adv_outputs,labels) + clean_preds=criterion(clean_outputs,labels) + loss_classifier=max_loss(clean_preds,adv_preds,torch.ones_like(preds)) # _, preds = torch.max(classifier(generated_img), 1) # loss_classifier=max_loss(torch.ones_like(criterion(outputs, labels)),criterion(outputs, labels),criterion(outputs, labels))