Team Ai
Apppublic

David310/Detect_AI-generated_Image

sourceHugging Faceupdated 2y agoView on Hugging Face
4likes
train.py86 linesDownload Raw Back to root
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