naver/PUMP
1
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 