MCPcopy Create free account
hub / github.com/Mattdl/ContinualPrototypeEvolution / main

Function main

main.py:392–499  ·  view source on GitHub ↗
(overwrite_args=None)

Source from the content-addressed store, hash-verified

390
391
392def main(overwrite_args=None):
393 args = parser.parse_args()
394 if overwrite_args is not None:
395 for k, v in overwrite_args.items(): # Debugging
396 setattr(args, k, v)
397
398 args.dyn_mem = True if args.dyn_mem == 'yes' else False
399 args.cuda = True if args.cuda == 'yes' else False
400 args.finetune = True if args.finetune == 'yes' else False
401 args.normalize = True if args.normalize == 'yes' else False
402 args.shared_head = True if args.shared_head == 'yes' else False
403 args.iid = True if args.iid == 'yes' else False
404
405 # unique identifier
406 uid = uuid.uuid4().hex if args.uid is None else args.uid
407 now = str(datetime.datetime.now().date()) + "_" + ':'.join(str(datetime.datetime.now().time()).split(':')[:-1])
408 runname = 'T={}_id={}'.format(now, uid) if not args.resume else args.resume
409
410 # Paths
411 setupname = [args.exp_name, args.model, args.data_file.split('.')[0]]
412 parentdir = os.path.join(args.save_path, '_'.join(setupname))
413
414 print("Init args={}".format(args))
415 stat_files = []
416 seeds = [args.seed] if args.seed is not None else list(range(args.n_seeds))
417 for seed in seeds:
418 # initialize seeds
419 print("STARTING SEED {}/{}".format(seed, args.n_seeds - 1))
420 torch.backends.cudnn.deterministic = False
421 torch.backends.cudnn.enabled = False
422 torch.manual_seed(seed)
423 np.random.seed(seed)
424 random.seed(seed)
425 if args.cuda:
426 torch.cuda.manual_seed_all(seed)
427
428 # load data
429 x_tr, x_te, n_inputs, n_classes, n_tasks = load_datasets(args)
430 args.is_cifar = ('cifar10' in args.data_file)
431 args.is_mnist = ('mnist' in args.data_file)
432 assert not (args.is_cifar and args.is_mnist)
433
434 args.input_shape = x_tr[0][1][0].shape
435 if args.input_shape[-1] == 3072: # CIFAR
436 assert args.is_cifar
437 args.CHW = (3, 32, 32)
438 elif args.input_shape[-1] == 784: # MNIST
439 assert args.is_mnist
440 args.CHW = (1, 28, 28)
441 else:
442 raise NotImplementedError()
443
444 args.n_classes = n_classes
445 n_outputs = args.n_classes if args.n_outputs is None else args.n_outputs # Embedding or Softmax
446
447 # set up continuum
448 continuum = Continuum(x_tr, args)
449

Callers 1

main.pyFile · 0.85

Calls 9

createdirsFunction · 0.90
confusion_matrixFunction · 0.90
load_datasetsFunction · 0.85
ContinuumClass · 0.85
ResultTrackerClass · 0.85
get_modelFunction · 0.85
life_experienceFunction · 0.85
stat_summarizeFunction · 0.85
get_allMethod · 0.80

Tested by

no test coverage detected