You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.25 07:32:50