()
| 65 | |
| 66 | |
| 67 | def main(): |
| 68 | dataset = MyImageFolder(root_dir=config.DATASET) |
| 69 | loader = DataLoader( |
| 70 | dataset, |
| 71 | batch_size=config.BATCH_SIZE, |
| 72 | shuffle=True, |
| 73 | pin_memory=True, |
| 74 | num_workers=config.NUM_WORKERS, |
| 75 | ) |
| 76 | #加载模型 |
| 77 | gen = Generator(in_channels=3).to(config.DEVICE) |
| 78 | disc = Discriminator(in_channels=3).to(config.DEVICE) |
| 79 | #定义优化器 |
| 80 | opt_gen = optim.Adam(gen.parameters(), lr=config.LEARNING_RATE, betas=(0.9, 0.999)) |
| 81 | opt_disc = optim.Adam(disc.parameters(), lr=config.LEARNING_RATE, betas=(0.9, 0.999)) |
| 82 | #定义均分误差损失函数 |
| 83 | mse = nn.MSELoss() |
| 84 | #定义交叉熵损失函数 |
| 85 | bce = nn.BCEWithLogitsLoss() |
| 86 | vgg_loss = VGGLoss() |
| 87 | |
| 88 | if config.LOAD_MODEL: |
| 89 | load_checkpoint( |
| 90 | config.CHECKPOINT_GEN_PRE, |
| 91 | gen, |
| 92 | opt_gen, |
| 93 | config.LEARNING_RATE, |
| 94 | ) |
| 95 | load_checkpoint( |
| 96 | config.CHECKPOINT_DISC_PRE, |
| 97 | disc, |
| 98 | opt_disc, |
| 99 | config.LEARNING_RATE, |
| 100 | ) |
| 101 | |
| 102 | for epoch in range(config.NUM_EPOCHS): |
| 103 | train_fn(loader, disc, gen, opt_gen, opt_disc, mse, bce, vgg_loss,epoch) |
| 104 | |
| 105 | print('epoch: {}'.format(epoch)) |
| 106 | if config.SAVE_MODEL and epoch > 0: |
| 107 | save_checkpoint(gen, opt_gen, filename=config.CHECKPOINT_GEN) |
| 108 | save_checkpoint(disc, opt_disc, filename=config.CHECKPOINT_DISC) |
| 109 | |
| 110 | |
| 111 | if __name__ == "__main__": |
no test coverage detected