| 36 | |
| 37 | |
| 38 | def 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 | |