Get dataloader with categorical sampler for few-shot classification.
(set_name: str, args: argparse, constant: bool = False)
| 49 | |
| 50 | |
| 51 | def get_dataloader(set_name: str, args: argparse, constant: bool = False): |
| 52 | """ |
| 53 | Get dataloader with categorical sampler for few-shot classification. |
| 54 | """ |
| 55 | num_episodes = args.set_episodes[set_name] |
| 56 | num_way = args.train_way if set_name == 'train' else args.val_way |
| 57 | |
| 58 | # define dataset sampler and data loader |
| 59 | data_set = DATASETS[args.dataset.lower()]( |
| 60 | args.data_path, set_name, args.backbone, |
| 61 | augment=set_name == 'train' and args.augment |
| 62 | ) |
| 63 | args.img_size = data_set.image_size |
| 64 | |
| 65 | data_sampler = CategoriesSampler( |
| 66 | set_name, data_set.label, num_episodes, const_loader=constant, |
| 67 | num_way=num_way, num_shot=args.num_shot, num_query=args.num_query, |
| 68 | replace=set_name == 'train', |
| 69 | ) |
| 70 | return DataLoader( |
| 71 | data_set, batch_sampler=data_sampler, num_workers=args.num_workers, pin_memory=not constant |
| 72 | ) |
| 73 | |
| 74 | |
| 75 | def get_optimizer_and_lr_scheduler(args, params): |
nothing calls this directly
no test coverage detected