MCPcopy Create free account
hub / github.com/LAMDA-CL/CVPR22-Fact / get_base_dataloader

Function get_base_dataloader

dataloader/data_utils.py:87–117  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

85 return trainset, trainloader, testloader
86
87def get_base_dataloader(args):
88 txt_path = "data/index_list/" + args.dataset + "/session_" + str(0 + 1) + '.txt'
89 class_index = np.arange(args.base_class)
90 if args.dataset == 'cifar100':
91
92 trainset = args.Dataset.CIFAR100(root=args.dataroot, train=True, download=True,
93 index=class_index, base_sess=True)
94 testset = args.Dataset.CIFAR100(root=args.dataroot, train=False, download=False,
95 index=class_index, base_sess=True)
96
97 if args.dataset == 'cub200':
98 trainset = args.Dataset.CUB200(root=args.dataroot, train=True,
99 index=class_index, base_sess=True)
100 testset = args.Dataset.CUB200(root=args.dataroot, train=False, index=class_index)
101
102 if args.dataset == 'mini_imagenet':
103 trainset = args.Dataset.MiniImageNet(root=args.dataroot, train=True,
104 index=class_index, base_sess=True)
105 testset = args.Dataset.MiniImageNet(root=args.dataroot, train=False, index=class_index)
106
107 if args.dataset == 'imagenet100' or args.dataset == 'imagenet1000':
108 trainset = args.Dataset.ImageNet(root=args.dataroot, train=True,
109 index=class_index, base_sess=True)
110 testset = args.Dataset.ImageNet(root=args.dataroot, train=False, index=class_index)
111
112 trainloader = torch.utils.data.DataLoader(dataset=trainset, batch_size=args.batch_size_base, shuffle=True,
113 num_workers=8, pin_memory=True)
114 testloader = torch.utils.data.DataLoader(
115 dataset=testset, batch_size=args.test_batch_size, shuffle=False, num_workers=8, pin_memory=True)
116
117 return trainset, trainloader, testloader
118
119
120

Callers 3

get_dataloaderFunction · 0.85
get_dataloaderMethod · 0.85
get_dataloaderMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected