MCPcopy Create free account
hub / github.com/Closed11/Unsupervised-Image-Classification / main

Function main

main.py:48–192  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

46
47
48def 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])

Callers 1

main.pyFile · 0.70

Calls 7

create_loggerFunction · 0.90
color_distortionFunction · 0.90
DatasetGivenLabelsClass · 0.90
UnifLabelSamplerClass · 0.90
compute_labelsFunction · 0.85
formatMethod · 0.80
trainFunction · 0.70

Tested by

no test coverage detected