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

PyTorch Geometric DataLoader两次迭代返回结果异常求助

PyTorch Geometric DataLoader 首次遍历批次异常问题

训练网络时发现PyTorch Geometric的DataLoader存在以下异常:

  • 首次遍历数据集时,前几个batch的batch索引仅为0,后续才恢复正常;第二次遍历所有batch均包含多个索引
  • 存在批次合并、数据被覆盖的情况,导致模型前向传播崩溃且训练数据丢失

复现代码

from Pointcloud.Modules.FileDataset import FileDataset
from Pointcloud.Modules import Config as config

dm = FileDataset(config.DATA_DIR, split_name=config.SPLIT_NAME, split=config.SPLIT)
train_dl = dm.train_dataloader(config.BATCH_SIZE, config.NUM_WORKERS)

for i in range(2):
    sum = 0
    for j, minibatch in enumerate(train_dl):
        if j < 8: # 仅展示前8个batch以复现问题
            print(j, minibatch.batch.unique(return_counts=True))
        sum += 1
    print(f"Iteration {i}\nNumber of iterations: {sum}\nNumber of batches: {len(train_dl.dataset) / config.BATCH_SIZE}")

输出结果

train bs: 25008
val bs: 8335
test bs: 8339
0 (tensor([0], device='cuda:0'), tensor([756], device='cuda:0'))
1 (tensor([0], device='cuda:0'), tensor([756], device='cuda:0'))
2 (tensor([0], device='cuda:0'), tensor([774], device='cuda:0'))
3 (tensor([0], device='cuda:0'), tensor([787], device='cuda:0'))
4 (tensor([ 0,  1,  2,  3,  4,  5,  6,  7,  8,  9, 10, 11, 12, 13, 14, 15],
       device='cuda:0'), tensor([39, 60, 44, 54, 52, 50, 51, 48, 36, 54, 61, 51, 51, 56, 55, 48],
       device='cuda:0'))
5 (tensor([ 0,  1,  2,  3,  4,  5,  6,  7,  8,  9, 10, 11, 12, 13, 14, 15],
       device='cuda:0'), tensor([47, 54, 56, 50, 51, 26, 75, 48, 43, 52, 55, 53, 35, 38, 51, 43],
       device='cuda:0'))
6 (tensor([ 0,  1,  2,  3,  4,  5,  6,  7,  8,  9, 10, 11, 12, 13, 14, 15],
       device='cuda:0'), tensor([47, 62, 53, 50, 54, 53, 52, 48, 57, 50, 53, 43, 53, 47, 59, 56],
       device='cuda:0'))
7 (tensor([ 0,  1,  2,  3,  4,  5,  6,  7,  8,  9, 10, 11, 12, 13, 14, 15],
       device='cuda:0'), tensor([52, 39, 43, 49, 52, 40, 49, 52, 52, 72, 57, 36, 28, 53, 50, 52],
       device='cuda:0'))
Iteration 0
Number of iterations: 1563
Number of batches: 1563.0
0 (tensor([ 0,  1,  2,  3,  4,  5,  6,  7,  8,  9, 10, 11, 12, 13, 14, 15],
       device='cuda:0'), tensor([43, 46, 50, 52, 53, 49, 63, 50, 36, 49, 53, 51, 40, 51, 56, 46],
       device='cuda:0'))
1 (tensor([ 0,  1,  2,  3,  4,  5,  6,  7,  8,  9, 10, 11, 12, 13, 14, 15],
       device='cuda:0'), tensor([50, 64, 44, 55, 54, 53, 53, 50, 49, 44, 52, 51, 47, 38, 54, 45],
       device='cuda:0'))
2 (tensor([ 0,  1,  2,  3,  4,  5,  6,  7,  8,  9, 10, 11, 12, 13, 14, 15],
       device='cuda:0'), tensor([44, 52, 49, 53, 49, 48, 52, 51, 44, 51, 42, 49, 53, 61, 46, 50],
       device='cuda:0'))
3 (tensor([ 0,  1,  2,  3,  4,  5,  6,  7,  8,  9, 10, 11, 12, 13, 14, 15],
       device='cuda:0'), tensor([59, 50, 40, 47, 33, 54, 59, 63, 53, 46, 44, 40, 49, 52, 53, 36],
       device='cuda:0'))
4 (tensor([ 0,  1,  2,  3,  4,  5,  6,  7,  8,  9, 10, 11, 12, 13, 14, 15],
       device='cuda:0'), tensor([54, 53, 49, 55, 54, 46, 28, 43, 45, 46, 51, 60, 49, 48, 54, 48],
       device='cuda:0'))
5 (tensor([ 0,  1,  2,  3,  4,  5,  6,  7,  8,  9, 10, 11, 12, 13, 14, 15],
       device='cuda:0'), tensor([54, 44, 59, 50, 45, 50, 35, 40, 54, 47, 39, 52, 52, 39, 50, 55],
       device='cuda:0'))
6 (tensor([ 0,  1,  2,  3,  4,  5,  6,  7,  8,  9, 10, 11, 12, 13, 14, 15],
       device='cuda:0'), tensor([49, 51, 42, 40, 51, 46, 48, 52, 45, 52, 38, 43, 50, 44, 49, 41],
       device='cuda:0'))
7 (tensor([ 0,  1,  2,  3,  4,  5,  6,  7,  8,  9, 10, 11, 12, 13, 14, 15],
       device='cuda:0'), tensor([49, 54, 61, 52, 48, 50, 36, 47, 49, 56, 41, 43, 52, 52, 54, 39],
       device='cuda:0'))
Iteration 1
Number of iterations: 1563
Number of batches: 1563.0

数据集示例

数据集为PyTorch Geometric Data对象列表,前5个元素如下:

[Data(x=[59, 8], edge_index=[2, 357], y=[1, 3]),
 Data(x=[50, 8], edge_index=[2, 292], y=[1, 3]),
 Data(x=[52, 8], edge_index=[2, 324], y=[1, 3]),
 Data(x=[56, 8], edge_index=[2, 362], y=[1, 3]),
 Data(x=[44, 8], edge_index=[2, 256], y=[1, 3])]

DataLoader创建函数

def train_dataloader(self, batch_size, num_workers):
    return tg_loader_DataLoader(
        dataset=self.train_ds,
        batch_size=batch_size,
        shuffle=True,
        num_workers=num_workers,
        persistent_workers=True,
        drop_last=True
    )

问题排查与解决

1. 核心原因

启用persistent_workers=True时,DataLoader的子进程会在epoch间保持存活,但如果FileDataset存在线程不安全的状态共享,或子进程中数据集未正确初始化,会导致PyG的batch索引生成逻辑失效,出现首次遍历的异常batch。此外部分旧版本PyG的DataLoader在该模式下存在bug,也会引发此类问题。

2. 修复方案

  • 临时禁用persistent_workers:这是最快的验证方式,修改后观察问题是否消失:
    def train_dataloader(self, batch_size, num_workers):
        return tg_loader_DataLoader(
            dataset=self.train_ds,
            batch_size=batch_size,
            shuffle=True,
            num_workers=num_workers,
            persistent_workers=False,  # 修改此处
            drop_last=True
        )
    
  • 检查数据集线程安全性:确保FileDataset的__getitem__无共享可变状态(如全局缓存、未初始化类变量),必要时将数据加载逻辑延迟到__getitem__中,避免子进程复用父进程状态。
  • 升级PyTorch Geometric版本:更新到最新稳定版(如2.3.0+),修复旧版本的已知bug。

3. 验证步骤

  1. 禁用persistent_workers后重新运行代码,确认首次epoch的前8个batch均包含多个索引。
  2. 若问题解决,再逐步排查数据集的线程安全问题,在确保安全的前提下重新启用persistent_workers。

内容的提问来源于stack exchange,提问作者Ruben Band

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 21:52:33