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

Function train_fn

fc-CGANCode/train.py:38–86  ·  view source on GitHub ↗
(generator,discriminator,optimizer_G,optimizer_D,adversarial_loss,dataloader,epoch)

Source from the content-addressed store, hash-verified

36
37
38def train_fn(generator,discriminator,optimizer_G,optimizer_D,adversarial_loss,dataloader,epoch):
39 loop = tqdm(dataloader,leave=True)
40 loop.set_description(desc="Training: ")
41 for step,(imgs,labels) in enumerate(loop):
42 valid = torch.ones(config.BATCH_SIZE)
43 fake = torch.zeros(config.BATCH_SIZE)
44
45 # print('labels.shape: {}'.format(np.shape(labels)))
46 # print('labels: {}'.format(labels))
47
48 real_imgs = imgs.to(config.DEVICE)
49 real_y = torch.zeros(config.BATCH_SIZE,config.NUM_CLASSES)
50 #dim = 1,index = config.BATCH_SIZE,src = 1
51 #https://mbd.baidu.com/ma/s/HT3QuRvI
52 real_y = real_y.scatter_(1, labels.view(config.BATCH_SIZE, 1), 1)
53
54 noise = torch.randn(size = (config.BATCH_SIZE,config.LATENT_DIM)).to(config.DEVICE)
55 gen_labels = (torch.rand(config.BATCH_SIZE,1)*config.NUM_CLASSES).type(torch.LongTensor)
56
57 # print('gen_labels.shape: {}'.format(gen_labels.shape)) => [BATCH_SIZE,1]
58 # print('gen_labels: {}'.format(gen_labels))[[3.6908],[8.2607],[7.1017],......]
59
60 gen_y = torch.zeros(config.BATCH_SIZE,config.NUM_CLASSES)
61 # dim = 1,index = config.BATCH_SIZE,src = 1
62 gen_y = gen_y.scatter_(1, gen_labels.view(config.BATCH_SIZE, 1), 1)
63
64 #compute the discriminator's loss
65 optimizer_D.zero_grad()
66 d_real_loss = adversarial_loss(np.squeeze(discriminator(real_imgs,real_y)),valid)
67 # gen_imgs = generator(noise,gen_y)
68 gen_imgs = generator(noise,real_y)
69 # d_fake_loss = adversarial_loss(np.squeeze(discriminator(gen_imgs.detach(),gen_y)),fake)
70 d_fake_loss = adversarial_loss(np.squeeze(discriminator(gen_imgs.detach(), real_y)), fake)
71 d_loss = d_real_loss + d_fake_loss
72 d_loss.backward()
73 optimizer_D.step()
74
75 #compute the generator's loss
76 optimizer_G.zero_grad()
77 # g_loss = adversarial_loss(np.squeeze(discriminator(gen_imgs,gen_y)),valid)
78 g_loss = adversarial_loss(np.squeeze(discriminator(gen_imgs, real_y)), valid)
79 g_loss.backward()
80 optimizer_G.step()
81
82 loop.set_postfix(
83 epoch = epoch,
84 d_loss = d_loss.item(),
85 g_loss = g_loss.item()
86 )
87
88
89

Callers 1

mainFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected