(overwrite_args=None)
| 390 | |
| 391 | |
| 392 | def 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 |
no test coverage detected