(args)
| 85 | return trainset, trainloader, testloader |
| 86 | |
| 87 | def 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 |
no outgoing calls
no test coverage detected