| 21 | from torch.utils.data import DataLoader,Dataset |
| 22 | |
| 23 | def train_fn(disc_X,disc_Y,gen_G,gen_F,loader,opt_disc,opt_gen,L1,mse,d_scale,g_scale,epoch,cudaIsAvailable = False): |
| 24 | loop = tqdm(loader,leave=True) |
| 25 | for idx ,(vango,photo) in enumerate(loop): |
| 26 | vango = vango.to(config.DEVICE) |
| 27 | photo = photo.to(config.DEVICE) |
| 28 | |
| 29 | #train discriminator |
| 30 | # if cudaIsAvailable==True: |
| 31 | # with torch.cuda.amp.autocast(): |
| 32 | #X -> Y |
| 33 | fake_photo = gen_G(vango) |
| 34 | D_X_real = disc_X(photo) |
| 35 | D_X_fake = disc_X(fake_photo.detach()) |
| 36 | D_X_real_loss = mse(D_X_real,torch.ones_like(D_X_real)) |
| 37 | D_X_fake_loss = mse(D_X_fake, torch.zeros_like(D_X_fake)) |
| 38 | D_X_loss = D_X_fake_loss + D_X_real_loss |
| 39 | |
| 40 | #Y -> X |
| 41 | fake_vango = gen_F(photo) |
| 42 | D_Y_real = disc_Y(vango) |
| 43 | D_Y_fake = disc_Y(fake_vango.detach()) |
| 44 | D_Y_real_loss = mse(D_Y_real, torch.ones_like(D_Y_real)) |
| 45 | D_Y_fake_loss = mse(D_Y_fake, torch.zeros_like(D_Y_fake)) |
| 46 | D_Y_loss = D_Y_fake_loss + D_Y_real_loss |
| 47 | |
| 48 | D_loss = D_X_loss + D_Y_loss |
| 49 | |
| 50 | if cudaIsAvailable: |
| 51 | opt_disc.zero_grad() |
| 52 | d_scale.scale(D_loss).backward() |
| 53 | d_scale.step(opt_disc) |
| 54 | d_scale.update() |
| 55 | else: |
| 56 | opt_disc.zero_grad() |
| 57 | D_loss.backward() |
| 58 | opt_disc.step() |
| 59 | |
| 60 | #train Generator H and Z |
| 61 | #with torch.cuda.amp.autocast(): |
| 62 | #adversarial loss for both generator |
| 63 | D_X_fake = disc_X(fake_photo) |
| 64 | D_Y_fake = disc_Y(fake_vango) |
| 65 | loss_G_Y = mse(D_X_fake,torch.ones_like(D_X_fake)) |
| 66 | loss_G_X = mse(D_X_fake,torch.ones_like(D_Y_fake)) |
| 67 | |
| 68 | #cycle loss |
| 69 | cycle_vango = gen_F(fake_photo) |
| 70 | cycle_photo = gen_G(fake_vango) |
| 71 | cycle_vango_loss = L1(vango,cycle_vango) |
| 72 | cycle_photo_loss = L1(photo,cycle_photo) |
| 73 | |
| 74 | # identity loss |
| 75 | identity_vango = gen_F(vango) |
| 76 | identity_photo = gen_G(photo) |
| 77 | identity_vango_loss = L1(vango,identity_vango) |
| 78 | identity_photo_loss = L1(photo,identity_photo) |
| 79 | |
| 80 | G_loss = ( |