MCPcopy Create free account
hub / github.com/DIVE128/DMVSNet / get_loader

Function get_loader

datasets/__init__.py:11–31  ·  view source on GitHub ↗
(args, datapath, listfile, nviews, mode="train",force_test=False)

Source from the content-addressed store, hash-verified

9
10
11def 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

Callers 2

__init__Method · 0.90
testMethod · 0.90

Calls 1

BlendedMVSDatasetClass · 0.85

Tested by

no test coverage detected