MCPcopy Create free account
hub / github.com/Topdu/OpenOCR / build_dataloader

Function build_dataloader

tools/data/__init__.py:39–136  ·  view source on GitHub ↗
(config, mode, logger, seed=None, epoch=1, task='rec')

Source from the content-addressed store, hash-verified

37
38
39def build_dataloader(config, mode, logger, seed=None, epoch=1, task='rec'):
40 config = copy.deepcopy(config)
41 mode = mode.capitalize()
42
43 # 获取 dataset 配置
44 dataset_config = config[mode]['dataset']
45 module_name = dataset_config['name']
46
47 # 动态导入 dataset 类
48 if module_name not in DATASET_MODULES:
49 raise ValueError(
50 f'Unsupported dataset: {module_name}. Supported datasets: {list(DATASET_MODULES.keys())}'
51 )
52
53 dataset_module = importlib.import_module(DATASET_MODULES[module_name])
54 dataset_class = getattr(dataset_module, module_name)
55 dataset = dataset_class(config, mode, logger, seed, epoch=epoch, task=task)
56
57 # DataLoader 配置
58 loader_config = config[mode]['loader']
59 num_workers = loader_config['num_workers']
60 pin_memory = loader_config.get('pin_memory', False)
61 if module_name == 'CMERWebDataSet':
62 logger.info(f"Building WebLoader for {module_name} (IterableDataset mode)...")
63 import webdataset as wds
64 persistent = num_workers > 0
65 data_loader = wds.WebLoader(
66 dataset,
67 batch_size=None, # 必须为 None,因为 dataset yield 的已经是 batch
68 shuffle=False, # 外部不打乱,内部处理
69 num_workers=num_workers,
70 pin_memory=True,
71 prefetch_factor=4,
72 persistent_workers=persistent,
73 )
74 total_iter_steps = config['Global'].get('total_iter_steps', 1000000)
75 data_loader = data_loader.with_length(total_iter_steps)
76 return data_loader
77 else:
78 batch_size = loader_config['batch_size_per_card']
79 drop_last = loader_config['drop_last']
80 shuffle = loader_config['shuffle']
81 sampler = None
82 batch_sampler = None
83 if 'sampler' in config[mode]:
84 sampler_config = config[mode]['sampler']
85 sampler_name = sampler_config.pop('name')
86
87 if sampler_name not in SAMPLER_MODULES:
88 raise ValueError(
89 f'Unsupported sampler: {sampler_name}. Supported samplers: {list(SAMPLER_MODULES.keys())}'
90 )
91
92 sampler_module = importlib.import_module(SAMPLER_MODULES[sampler_name])
93 sampler_class = getattr(sampler_module, sampler_name)
94 batch_sampler = sampler_class(dataset, **sampler_config)
95 elif config['Global']['distributed'] and mode == 'Train':
96 sampler = DistributedSampler(dataset=dataset, shuffle=shuffle)

Callers 6

mainFunction · 0.90
mainFunction · 0.90
mainFunction · 0.90
mainFunction · 0.90
__init__Method · 0.90
trainMethod · 0.90

Calls 1

getMethod · 0.80

Tested by

no test coverage detected