MCPcopy Create free account
hub / github.com/NVIDIA/FasterTransformer / build_loader

Function build_loader

examples/pytorch/vit/ViT-quantization/data.py:69–108  ·  view source on GitHub ↗
(config, args)

Source from the content-addressed store, hash-verified

67 return dataset_val, data_loader_val
68
69def build_loader(config, args):
70 config.defrost()
71 dataset_train, config.MODEL.NUM_CLASSES = build_dataset(is_train=True, config=config)
72 config.freeze()
73 # print(f"local rank {config.LOCAL_RANK} / global rank {dist.get_rank()} successfully build train dataset")
74 dataset_val, _ = build_dataset(is_train=False, config=config)
75 # print(f"local rank {config.LOCAL_RANK} / global rank {dist.get_rank()} successfully build val dataset")
76
77 num_tasks = dist.get_world_size()
78 global_rank = dist.get_rank()
79
80 sampler_train = torch.utils.data.DistributedSampler(
81 dataset_train, num_replicas=num_tasks, rank=global_rank, shuffle=True
82 )
83
84 if config.TEST.SEQUENTIAL:
85 sampler_val = torch.utils.data.SequentialSampler(dataset_val)
86 else:
87 sampler_val = torch.utils.data.distributed.DistributedSampler(
88 dataset_val, shuffle=False
89 )
90
91 data_loader_train = torch.utils.data.DataLoader(
92 dataset_train, sampler=sampler_train,
93 batch_size=args.calib_batchsz if args.calib else config.DATA.BATCH_SIZE,
94 num_workers=config.DATA.NUM_WORKERS,
95 pin_memory=config.DATA.PIN_MEMORY,
96 drop_last=True,
97 )
98
99 data_loader_val = torch.utils.data.DataLoader(
100 dataset_val, sampler=sampler_val,
101 batch_size=config.DATA.BATCH_SIZE,
102 shuffle=False,
103 num_workers=config.DATA.NUM_WORKERS,
104 pin_memory=config.DATA.PIN_MEMORY,
105 drop_last=True
106 )
107
108 return dataset_train, dataset_val, data_loader_train, data_loader_val
109
110
111def build_dataset(is_train, config):

Callers 5

calibFunction · 0.90
trainFunction · 0.90
calibFunction · 0.90
validate_trtFunction · 0.90
trainFunction · 0.90

Calls 1

build_datasetFunction · 0.85

Tested by

no test coverage detected