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

Function main

SRGAN/train.py:67–108  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

65
66
67def 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
111if __name__ == "__main__":

Callers 1

train.pyFile · 0.70

Calls 7

MyImageFolderClass · 0.90
GeneratorClass · 0.90
DiscriminatorClass · 0.90
VGGLossClass · 0.90
load_checkpointFunction · 0.90
save_checkpointFunction · 0.90
train_fnFunction · 0.70

Tested by

no test coverage detected