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. 验证步骤
- 禁用
persistent_workers后重新运行代码,确认首次epoch的前8个batch均包含多个索引。 - 若问题解决,再逐步排查数据集的线程安全问题,在确保安全的前提下重新启用
persistent_workers。
内容的提问来源于stack exchange,提问作者Ruben Band
相关产品推荐
相关产品推荐

