CoolFace
Apppublic

naver/PUMP

sourceHugging Faceupdated 4y agoView on Hugging Face
1likes
train.py122 linesDownload Raw Back to root
1# Copyright 2022-present NAVER Corp.2# CC BY-NC-SA 4.03# Available only for non-commercial use4 5from pdb import set_trace as bb6import os7import torch8import torch.optim as optim9import torchvision.transforms as tvf10 11from tools import common, trainer12from datasets import *13from core.conv_mixer import ConvMixer14from core.losses import *15 16 17def parse_args():18    import argparse19    parser = argparse.ArgumentParser("Script to train PUMP")20 21    parser.add_argument("--pretrained", type=str, default="", help='pretrained model path')22    parser.add_argument("--save-path", type=str, required=True, help='directory to save model')23 24    parser.add_argument("--epochs", type=int, default=50, help='number of training epochs')25    parser.add_argument("--batch-size", "--bs", type=int, default=16, help="batch size")26    parser.add_argument("--learning-rate", "--lr", type=str, default=1e-4)27    parser.add_argument("--weight-decay", "--wd", type=float, default=5e-4)28    29    parser.add_argument("--threads", type=int, default=8, help='number of worker threads')30    parser.add_argument("--device", default='cuda')31    32    args = parser.parse_args()33    return args34 35 36def main( args ):37    device = args.device38    common.mkdir_for(args.save_path)39 40    # Create data loader41    db = BalancedCatImagePairs(42            3125, SyntheticImagePairs(RandomWebImages(0,52),distort='RandomTilting(0.5)'),43            4875, SyntheticImagePairs(SfM120k_Images(),distort='RandomTilting(0.5)'),44            8000, SfM120k_Pairs())45 46    db = FastPairLoader(db,47            crop=256, transform='RandomRotation(20), RandomScale(256,1536,ar=1.3,can_upscale=True), PixelNoise(25)', 48            p_swap=0.5, p_flip=0.5, scale_jitter=0.5)49 50    print("Training image database =", db)51    data_loader = torch.utils.data.DataLoader(db, batch_size=args.batch_size, shuffle=True, 52            num_workers=args.threads, collate_fn=collate_ordered, pin_memory=False, drop_last=True, 53            worker_init_fn=WorkerWithRngInit())54 55    # create network56    net = ConvMixer(output_dim=128, hidden_dim=512, depth=7, patch_size=4, kernel_size=9)57    print(f"\n>> Creating {type(net).__name__} net ( Model size: {common.model_size(net)/1e6:.1f}M parameters )")58 59    # create losses60    loss = MultiLoss(alpha=0.3, 61            loss_sup = PixelAPLoss(nq=20, inner_bw=True, sampler=NghSampler(ngh=7)), 62            loss_unsup = DeepMatchingLoss(eps=0.03))63    64    # create optimizer65    optimizer = optim.Adam( [p for p in net.parameters() if p.requires_grad], 66                            lr=args.learning_rate, weight_decay=args.weight_decay)67 68    train = MyTrainer(net, loss, optimizer).to(device)69 70    # initialization71    final_model_path = osp.join(args.save_path,'model.pt')72    last_model_path = osp.join(args.save_path,'model.pt.last')73    if osp.exists( final_model_path ):74        print('Already trained, nothing to do!')75        return76    elif args.pretrained: 77        train.load( args.pretrained )78    elif osp.exists( last_model_path ):79        train.load( last_model_path )80 81    train = train.to(args.device)82    if ',' in os.environ.get('CUDA_VISIBLE_DEVICES',''):83        train.distribute()84 85    # Training loop #86    while train.epoch < args.epochs:87        # shuffle dataset (select new pairs)88        data_loader.dataset.set_epoch(train.epoch)89 90        train(data_loader)91 92        train.save(last_model_path)93 94    # save final model95    torch.save(train.model.state_dict(), open(final_model_path,'wb'))96 97 98totensor = tvf.Compose([99    common.ToTensor(), 100    tvf.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])101    ])102 103class MyTrainer (trainer.Trainer):104    """ This class implements the network training.105        Below is the function I need to overload to explain how to do the backprop.106    """107    def forward_backward(self, inputs):108        assert torch.is_grad_enabled() and self.net.training109 110        (img1, img2), labels = inputs111        output1 = self.net(totensor(img1))112        output2 = self.net(totensor(img2))113 114        loss, details = trainer.get_loss(self.loss(output1, output2, img1=img1, img2=img2, **labels))115        trainer.backward(loss)116        return details117 118 119 120if __name__ == '__main__':121    main(parse_args())122