()
| 113 | |
| 114 | |
| 115 | def main(): |
| 116 | args = get_args() |
| 117 | utils.set_seed(seed=args.seed) |
| 118 | print("\n".join("%s: %s" % (k, str(v)) for k, v in sorted(dict(vars(args)).items()))) |
| 119 | output_dir = utils.get_output_dir(args=args) |
| 120 | |
| 121 | # define datasets and loaders |
| 122 | args.set_episodes = dict(train=args.train_episodes, val=args.eval_episodes, test=args.test_episodes) |
| 123 | if not args.eval: |
| 124 | train_dataloader = utils.get_dataloader(set_name='train', args=args, constant=False) |
| 125 | val_dataloader = utils.get_dataloader(set_name='val', args=args, constant=True) |
| 126 | else: |
| 127 | val_dataloader = utils.get_dataloader(set_name='test', args=args, constant=False) |
| 128 | train_dataloader = None |
| 129 | |
| 130 | # define model and load pretrained weights if available |
| 131 | model = utils.get_model(args.backbone, args) |
| 132 | model = model.to(device) |
| 133 | utils.load_weights(model, args.pretrained_path) |
| 134 | |
| 135 | # BPA and few-shot classification method (e.g. proto, pt-map...) |
| 136 | bpa = None |
| 137 | if 'bpa' in args.method.lower(): |
| 138 | bpa = BPA( |
| 139 | distance_metric=args.distance_metric, |
| 140 | ot_reg=args.ot_reg, |
| 141 | mask_diag=args.mask_diag, |
| 142 | sinkhorn_iterations=args.sink_iters, |
| 143 | max_scale=args.max_scale |
| 144 | ) |
| 145 | fewshot_method = utils.get_method(args=args, bpa=bpa) |
| 146 | |
| 147 | # few-shot labels |
| 148 | train_labels = utils.get_fs_labels(args.method, args.train_way, args.num_query, args.num_shot) |
| 149 | val_labels = utils.get_fs_labels(args.method, args.val_way, args.num_query, args.num_shot) |
| 150 | |
| 151 | # initialized wandb |
| 152 | if args.wandb: |
| 153 | utils.init_wandb(exp_name=output_dir.split('/')[-1] if output_dir[-1] != '/' else output_dir.split('/')[-2], |
| 154 | args=args) |
| 155 | |
| 156 | # define loss |
| 157 | criterion = utils.get_criterion_by_method(method=args.method) |
| 158 | |
| 159 | # Test-set evaluation |
| 160 | if args.eval: |
| 161 | print(f"Evaluate model for {args.test_episodes} episodes... ") |
| 162 | loss, acc = eval_one_epoch(model, val_dataloader, fewshot_method, criterion, val_labels, 0, args, set_name='test') |
| 163 | print("Final evaluation results:\nAccuracy={:.4f}, Loss={:.4f}".format(acc, loss)) |
| 164 | exit(1) |
| 165 | |
| 166 | # define optimizer and scheduler |
| 167 | optimizer, lr_scheduler = utils.get_optimizer_and_lr_scheduler(args=args, params=model.parameters()) |
| 168 | |
| 169 | # evaluate model before training |
| 170 | if args.eval_first: |
| 171 | print("Evaluate model before training... ") |
| 172 | eval_one_epoch(model, val_dataloader, fewshot_method, criterion, val_labels, -1, args, set_name='val') |
no test coverage detected