MCPcopy Create free account
hub / github.com/DragonisCV/RAM / create_train_val_dataloader

Function create_train_val_dataloader

ram/train.py:29–66  ·  view source on GitHub ↗
(opt, logger)

Source from the content-addressed store, hash-verified

27
28
29def create_train_val_dataloader(opt, logger):
30 # create train and val dataloaders
31 train_loader, val_loaders = None, []
32 for phase, dataset_opt in opt['datasets'].items():
33 dataset_opt['gt_size'] = opt['gt_size']
34 if phase == 'train':
35 dataset_enlarge_ratio = dataset_opt.get('dataset_enlarge_ratio', 1)
36 train_set = build_dataset(dataset_opt)
37 train_sampler = EnlargedSampler(train_set, opt['world_size'], opt['rank'], dataset_enlarge_ratio)
38 train_loader = build_dataloader(
39 train_set,
40 dataset_opt,
41 num_gpu=opt['num_gpu'],
42 dist=opt['dist'],
43 sampler= train_sampler,
44 seed=opt['manual_seed'])
45
46 num_iter_per_epoch = math.ceil(
47 len(train_set) * dataset_enlarge_ratio / (dataset_opt['batch_size_per_gpu'] * opt['world_size']))
48 total_iters = int(opt['train']['total_iter'])
49 total_epochs = math.ceil(total_iters / (num_iter_per_epoch))
50 logger.info('Training statistics:'
51 f'\n\tNumber of train images: {len(train_set)}'
52 f'\n\tDataset enlarge ratio: {dataset_enlarge_ratio}'
53 f'\n\tBatch size per gpu: {dataset_opt["batch_size_per_gpu"]}'
54 f'\n\tWorld size (gpu number): {opt["world_size"]}'
55 f'\n\tRequire iter number per epoch: {num_iter_per_epoch}'
56 f'\n\tTotal epochs: {total_epochs}; iters: {total_iters}.')
57 elif phase.split('_')[0] == 'val':
58 val_set = build_dataset(dataset_opt)
59 val_loader = build_dataloader(
60 val_set, dataset_opt, num_gpu=opt['num_gpu'], dist=opt['dist'], sampler=None, seed=opt['manual_seed'])
61 logger.info(f'Number of val images/folders in {dataset_opt["name"]}: {len(val_set)}')
62 val_loaders.append(val_loader)
63 else:
64 raise ValueError(f'Dataset phase {phase} is not recognized.')
65
66 return train_loader, train_sampler, val_loaders, total_epochs, total_iters
67
68
69def load_resume_state(opt):

Callers 1

train_pipelineFunction · 0.85

Calls 4

build_datasetFunction · 0.90
EnlargedSamplerClass · 0.90
build_dataloaderFunction · 0.90
getMethod · 0.45

Tested by

no test coverage detected