MCPcopy Create free account
hub / github.com/KeepTryingTo/Pytorch-GAN / train_fn

Function train_fn

Code/train.py:23–103  ·  view source on GitHub ↗
(disc_X,disc_Y,gen_G,gen_F,loader,opt_disc,opt_gen,L1,mse,d_scale,g_scale,epoch,cudaIsAvailable = False)

Source from the content-addressed store, hash-verified

21from torch.utils.data import DataLoader,Dataset
22
23def 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 = (

Callers 1

main_Function · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected