()
| 46 | |
| 47 | |
| 48 | def main(): |
| 49 | # log-file setting |
| 50 | global args |
| 51 | args = parser.parse_args() |
| 52 | log_file_name = args.save_name + time.strftime('%Y-%m-%d-%H-%M-%S', time.localtime(time.time())) + '.log' |
| 53 | global logger |
| 54 | logger = create_logger(os.path.join(args.exp, log_file_name)) |
| 55 | logger.info("============ Initialized logger ============") |
| 56 | logger.info("\n".join("%s: %s" % (k, str(v)) |
| 57 | for k, v in sorted(dict(vars(args)).items()))) |
| 58 | logger.info("The experiment will be stored in %s\n" % args.exp) |
| 59 | logger.info("") |
| 60 | |
| 61 | # fix random seeds |
| 62 | torch.manual_seed(args.seed) |
| 63 | torch.cuda.manual_seed_all(args.seed) |
| 64 | np.random.seed(args.seed) |
| 65 | |
| 66 | # CNN |
| 67 | if args.verbose: |
| 68 | logger.info('Architecture: {}'.format(args.arch)) |
| 69 | # extra mlp head & random gaussian bluring augmentation |
| 70 | model = models.__dict__[args.arch](out=args.nmb_cluster, extra_mlp=True, random_gblur=True) |
| 71 | model = torch.nn.DataParallel(model) |
| 72 | model.cuda() |
| 73 | cudnn.benchmark = True |
| 74 | |
| 75 | # create optimizer |
| 76 | optimizer = torch.optim.SGD( |
| 77 | filter(lambda x: x.requires_grad, model.parameters()), |
| 78 | lr=args.lr, |
| 79 | momentum=args.momentum, |
| 80 | weight_decay=10**args.wd, |
| 81 | ) |
| 82 | lr_scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.epochs, \ |
| 83 | eta_min=0, last_epoch=-1) |
| 84 | |
| 85 | # define loss function |
| 86 | criterion = nn.CrossEntropyLoss() |
| 87 | |
| 88 | # optionally resume from a checkpoint |
| 89 | start_epoch = 0 |
| 90 | if args.resume: |
| 91 | if os.path.isfile(args.resume): |
| 92 | logger.info("=> loading checkpoint '{}'".format(args.resume)) |
| 93 | checkpoint = torch.load(args.resume) |
| 94 | start_epoch = checkpoint['epoch'] |
| 95 | model.load_state_dict(checkpoint['state_dict']) |
| 96 | optimizer.load_state_dict(checkpoint['optimizer']) |
| 97 | logger.info("=> loaded checkpoint '{}' (epoch {})" |
| 98 | .format(args.resume, checkpoint['epoch'])) |
| 99 | else: |
| 100 | logger.info("=> no checkpoint found at '{}'".format(args.resume)) |
| 101 | |
| 102 | end = time.time() |
| 103 | # preprocessing of data |
| 104 | normalize = transforms.Normalize(mean=[0.485, 0.456, 0.406], |
| 105 | std=[0.229, 0.224, 0.225]) |
no test coverage detected