在PyTorch Lightning中处理多数据集/数据加载器的问题
解决PyTorch Lightning多数据集训练的两个核心问题
方案一:自定义多数据集迭代器(保留独立Batch,循环小数据集)
这个方案让小数据集在耗尽后自动重置迭代器,直到最大数据集的所有批次跑完,同时支持按数据集加权损失。
1. 定义自定义迭代器类
class MultiDatasetIterator: def __init__(self, loaders): self.loaders = loaders self.iterators = {key: iter(loader) for key, loader in loaders.items()} # 以最大数据集的批次数作为当前epoch的总迭代次数 self.total_batches = max(len(loader) for loader in loaders.values()) self.current_batch = 0 def __iter__(self): return self def __next__(self): if self.current_batch >= self.total_batches: raise StopIteration batch = {} for key in self.loaders.keys(): try: batch[key] = next(self.iterators[key]) except StopIteration: # 小数据集耗尽,重新生成迭代器继续采样 self.iterators[key] = iter(self.loaders[key]) batch[key] = next(self.iterators[key]) self.current_batch += 1 return batch def __len__(self): return self.total_batches
2. 修改train_dataloader方法
def train_dataloader(self): train_loaders = {} for key, value in self.train_dict.items(): train_loaders[key] = DataLoader(value, batch_size=self.batch_size, collate_fn=collate) # 返回自定义迭代器而非字典 return MultiDatasetIterator(train_loaders)
3. 实现损失加权的training_step
提前在模型__init__中定义好各数据集的权重字典(比如self.loss_weights = {"dataset_a": 0.4, "dataset_b": 0.6}),然后修改训练步骤:
def training_step(self, batch, batch_idx): total_batch_loss = 0 for key, value in batch.items(): anc, pos, neg = value emb_anc = F.normalize(self.forward(anc.x, anc.edge_index, anc.weights, anc.batch, training=True ), 2, dim=1) emb_pos = F.normalize(self.forward(pos.x, pos.edge_index, pos.weights, pos.batch, training=True ), 2, dim=1) emb_neg = F.normalize(self.forward(neg.x, neg.edge_index, neg.weights, neg.batch, training=True ), 2, dim=1) loss_dataset = LossFunc(emb_anc, emb_pos, emb_neg, anc.y, pos.y, neg.y) # 乘以对应数据集的权重 total_batch_loss += loss_dataset * self.loss_weights[key] self.log("Loss", total_batch_loss, prog_bar=True, on_epoch=True) return total_batch_loss
方案二:合并数据集并保留标识(灵活采样+加权)
通过给每个样本添加所属数据集的标签,合并成单个DataLoader,同时支持按数据集加权损失,还能通过采样器控制各数据集的采样频率。
1. 定义带标识的包装数据集
class LabeledDataset(Dataset): def __init__(self, dataset, dataset_key): self.dataset = dataset self.dataset_key = dataset_key def __len__(self): return len(self.dataset) def __getitem__(self, idx): # 返回样本+所属数据集的key return self.dataset[idx], self.dataset_key
2. 合并数据集并创建DataLoader
def train_dataloader(self): concat_datasets = [] for key, value in self.train_dict.items(): concat_datasets.append(LabeledDataset(value, key)) # 合并所有数据集 combined_dataset = ConcatDataset(concat_datasets) # 可选:用WeightedRandomSampler控制各数据集的采样比例(比如让小数据集被采样更多) # sample_weights = [] # for key, value in self.train_dict.items(): # sample_weights.extend([self.loss_weights[key]] * len(value)) # sampler = WeightedRandomSampler(sample_weights, len(combined_dataset)) return DataLoader(combined_dataset, batch_size=self.batch_size, collate_fn=self.custom_collate) # sampler=sampler # 可选启用加权采样
3. 自定义Collate函数分组样本
from collections import defaultdict def custom_collate(self, batch): # 按数据集key分组样本 grouped = defaultdict(list) for sample, key in batch: grouped[key].append(sample) # 对每组应用原有的collate_fn collated_batch = {} for key, samples in grouped.items(): collated_batch[key] = collate(samples) return collated_batch
4. 训练步骤同方案一
training_step的代码和方案一完全一致,直接按数据集key取权重计算即可。
内容的提问来源于stack exchange,提问作者James Arten
相关产品推荐
相关产品推荐

