MCPcopy Create free account
hub / github.com/HYUNJS/SGT / build_batch_data_loader

Function build_batch_data_loader

projects/Datasets/MOT/build.py:233–251  ·  view source on GitHub ↗
(dataset, sampler, total_batch_size, *, num_workers=0)

Source from the content-addressed store, hash-verified

231
232
233def build_batch_data_loader(dataset, sampler, total_batch_size, *, num_workers=0):
234 world_size = get_world_size()
235 assert (
236 total_batch_size > 0 and total_batch_size % world_size == 0
237 ), "Total batch size ({}) must be divisible by the number of gpus ({}).".format(
238 total_batch_size, world_size
239 )
240
241 batch_size = total_batch_size // world_size
242 batch_sampler = torch.utils.data.sampler.BatchSampler(
243 sampler, batch_size, drop_last=True
244 ) # drop_last so the batch always have the same size
245 return torch.utils.data.DataLoader(
246 dataset,
247 num_workers=num_workers,
248 batch_sampler=batch_sampler,
249 collate_fn=trivial_batch_collator,
250 worker_init_fn=worker_init_reset_seed,
251 )
252
253
254def trivial_batch_collator(batch):

Callers 2

build_mix_train_loaderFunction · 0.90
build_mot_train_loaderFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected