| 37 | |
| 38 | |
| 39 | def 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) |