(args, datapath, listfile, nviews, mode="train",force_test=False)
| 9 | |
| 10 | |
| 11 | def get_loader(args, datapath, listfile, nviews, mode="train",force_test=False): |
| 12 | |
| 13 | |
| 14 | if args.dataset_name == "dtu_yao": |
| 15 | dataset = DtuDataset(datapath, listfile, mode, nviews, args.img_size, args.numdepth, args.interval_scale) |
| 16 | elif args.dataset_name == "general_eval": |
| 17 | dataset = EvalDataset(datapath, listfile, mode, nviews, args.numdepth, args.interval_scale, args.inverse_depth, |
| 18 | max_h=args.max_h, max_w=args.max_w, fix_res=args.fix_res) |
| 19 | elif args.dataset_name == "blendedmvs": |
| 20 | dataset = BlendedMVSDataset(datapath, listfile, mode, nviews, args.numdepth, args.interval_scale) |
| 21 | else: |
| 22 | raise NotImplementedError("Don't support dataset: {}".format(args.dataset_name)) |
| 23 | |
| 24 | if args.distributed: |
| 25 | sampler = torch.utils.data.DistributedSampler(dataset, num_replicas=dist.get_world_size(), rank=dist.get_rank()) |
| 26 | else: |
| 27 | sampler = RandomSampler(dataset) if (mode == "train") else SequentialSampler(dataset) |
| 28 | |
| 29 | data_loader = data.DataLoader(dataset, args.batch_size, sampler=sampler, num_workers=4, drop_last=(mode == "train"), pin_memory=True) |
| 30 | |
| 31 | return data_loader, sampler |
no test coverage detected