change loss
This commit is contained in:
parent
63def0a779
commit
cc87c09e86
14
GanAttack.py
14
GanAttack.py
|
|
@ -76,7 +76,7 @@ def main(cfg: DictConfig) -> None:
|
||||||
criterion = nn.CrossEntropyLoss()
|
criterion = nn.CrossEntropyLoss()
|
||||||
max_loss=nn.MarginRankingLoss(0.1)
|
max_loss=nn.MarginRankingLoss(0.1)
|
||||||
clip_loss=CLIPLoss().to(device)
|
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)
|
# summary(model, input_size = (3, 256, 256), batch_size = 5)
|
||||||
# set_requires_grad(model.mlp.parameters())
|
# set_requires_grad(model.mlp.parameters())
|
||||||
for p in (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)
|
# _, _, _, clean_refine_images, clean_latent_codes, _=inverter(inputs,img_path)
|
||||||
optimizer.zero_grad()
|
optimizer.zero_grad()
|
||||||
generated_img,adv_latent_codes=model(inputs)
|
generated_img,adv_latent_codes=model(inputs)
|
||||||
loss_vgg=vgg_loss(inputs,generated_img)
|
# loss_vgg=vgg_loss(inputs,generated_img)
|
||||||
# loss_l1=F.l1_loss(codes,adv_latent_codes)
|
loss_l1=F.l1_loss(codes,adv_latent_codes)
|
||||||
loss_clip=clip_loss(generated_img,prompt)
|
loss_clip=clip_loss(generated_img,prompt)
|
||||||
outputs = classifier(generated_img)
|
adv_outputs = classifier(generated_img)
|
||||||
preds=criterion(outputs,labels)
|
clean_outputs=classifier(inputs)
|
||||||
loss_classifier=max_loss(torch.ones_like(preds),preds,-torch.ones_like(preds))
|
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)
|
# _, preds = torch.max(classifier(generated_img), 1)
|
||||||
# loss_classifier=max_loss(torch.ones_like(criterion(outputs, labels)),criterion(outputs, labels),criterion(outputs, labels))
|
# loss_classifier=max_loss(torch.ones_like(criterion(outputs, labels)),criterion(outputs, labels),criterion(outputs, labels))
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue