MCPcopy Create free account
hub / github.com/DanielShalam/BPA / get_dataloader

Function get_dataloader

utils.py:51–72  ·  view source on GitHub ↗

Get dataloader with categorical sampler for few-shot classification.

(set_name: str, args: argparse, constant: bool = False)

Source from the content-addressed store, hash-verified

49
50
51def 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
75def get_optimizer_and_lr_scheduler(args, params):

Callers

nothing calls this directly

Calls 1

CategoriesSamplerClass · 0.90

Tested by

no test coverage detected