David310/Detect_AI-generated_Image
4
1import os2import time3from tensorboardX import SummaryWriter4 5from validate import validate6from data import create_dataloader7from earlystop import EarlyStopping8from networks.trainer import Trainer9from options.train_options import TrainOptions10 11 12"""Currently assumes jpg_prob, blur_prob 0 or 1"""13def get_val_opt():14 val_opt = TrainOptions().parse(print_options=False)15 val_opt.isTrain = False16 val_opt.no_resize = False17 val_opt.no_crop = False18 val_opt.serial_batches = True19 val_opt.data_label = 'val'20 val_opt.jpg_method = ['pil']21 if len(val_opt.blur_sig) == 2:22 b_sig = val_opt.blur_sig23 val_opt.blur_sig = [(b_sig[0] + b_sig[1]) / 2]24 if len(val_opt.jpg_qual) != 1:25 j_qual = val_opt.jpg_qual26 val_opt.jpg_qual = [int((j_qual[0] + j_qual[-1]) / 2)]27 28 return val_opt29 30 31 32if __name__ == '__main__':33 opt = TrainOptions().parse()34 val_opt = get_val_opt()35 36 model = Trainer(opt)37 38 data_loader = create_dataloader(opt)39 val_loader = create_dataloader(val_opt)40 41 train_writer = SummaryWriter(os.path.join(opt.checkpoints_dir, opt.name, "train"))42 val_writer = SummaryWriter(os.path.join(opt.checkpoints_dir, opt.name, "val"))43 44 early_stopping = EarlyStopping(patience=opt.earlystop_epoch, delta=-0.001, verbose=True)45 start_time = time.time()46 print ("Length of data loader: %d" %(len(data_loader)))47 for epoch in range(opt.niter):48 49 for i, data in enumerate(data_loader):50 model.total_steps += 151 52 model.set_input(data)53 model.optimize_parameters()54 55 if model.total_steps % opt.loss_freq == 0:56 print("Train loss: {} at step: {}".format(model.loss, model.total_steps))57 train_writer.add_scalar('loss', model.loss, model.total_steps)58 print("Iter time: ", ((time.time()-start_time)/model.total_steps) )59 60 if model.total_steps in [10,30,50,100,1000,5000,10000] and False: # save models at these iters 61 model.save_networks('model_iters_%s.pth' % model.total_steps)62 63 if epoch % opt.save_epoch_freq == 0:64 print('saving the model at the end of epoch %d' % (epoch))65 model.save_networks( 'model_epoch_best.pth' )66 model.save_networks( 'model_epoch_%s.pth' % epoch )67 68 # Validation69 model.eval()70 ap, r_acc, f_acc, acc = validate(model.model, val_loader)71 val_writer.add_scalar('accuracy', acc, model.total_steps)72 val_writer.add_scalar('ap', ap, model.total_steps)73 print("(Val @ epoch {}) acc: {}; ap: {}".format(epoch, acc, ap))74 75 early_stopping(acc, model)76 if early_stopping.early_stop:77 cont_train = model.adjust_learning_rate()78 if cont_train:79 print("Learning rate dropped by 10, continue training...")80 early_stopping = EarlyStopping(patience=opt.earlystop_epoch, delta=-0.002, verbose=True)81 else:82 print("Early stopping.")83 break84 model.train()85 86 