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

如何在Torch Geometric Dataset中实现多进程加载与数据预取

PyTorch Geometric 数据集并行处理与预取实现方案

你不需要修改自定义Dataset类的get()/__getitem__单样本加载逻辑,基于PyTorch与PyG原生的DataLoader能力,即可对齐tf.data.Dataset的两类核心特性,无需自行实现复杂的多进程或预取逻辑。

1. 多进程并行处理(对应num_parallel_calls能力)

这部分能力直接通过PyG DataLoader的num_workers参数实现,和自定义PyG数据集完全兼容:

  • 保持原有Dataset.get()方法的单样本加载、单样本处理逻辑不变,不要在方法内部加入全局锁、跨进程共享可变对象,避免触发多进程死锁
  • 图数据的转换逻辑(节点特征归一化、边采样、子图提取、数据增强等)可以直接写在get()方法或者transform回调中,多进程模式下每个worker进程会独立执行处理逻辑,不会阻塞主训练进程
  • 如果是规模超过内存的大图数据集,不要使用InMemoryDataset,改用按需读盘的普通Dataset类,避免多进程启动时全量数据复制导致内存溢出

基础配置示例:

from torch_geometric.loader import DataLoader

# MyDataset为你自定义实现的PyG数据集类
dataset = MyDataset(root="./your_dataset_path")

loader = DataLoader(
    dataset,
    batch_size=32,
    shuffle=True,
    num_workers=8,  # 等价于tf.data的num_parallel_calls,可根据CPU核心数调整
    pin_memory=True,  # 锁页内存配置,配合预取加速CPU到GPU的数据传输
    persistent_workers=True  # 训练全程保留worker进程,避免每个epoch重复启停进程的额外开销
)

2. 数据预取(对应.prefetch()能力)

PyG DataLoader内置了CPU侧预取能力,配合轻量包装类即可实现GPU训练、CPU预处理并行的流水线效果,完全覆盖数据加载最佳实践的要求。

基础CPU预取配置

直接通过DataLoader内置参数即可开启:

  • prefetch_factor:控制每个worker进程提前缓存的batch数量,一般设置为2~4即可,过大会占用过多不必要的内存,比如设置为2时,8个worker会提前准备16个batch在内存队列中
  • 配合pin_memory=True配置,预取的数据会存入锁页内存,拷贝到GPU时不需要额外CPU寻址开销,传输速度提升明显

只需要在之前的DataLoader配置中加入对应参数即可:

loader = DataLoader(
    dataset,
    batch_size=32,
    shuffle=True,
    num_workers=8,
    pin_memory=True,
    persistent_workers=True,
    prefetch_factor=2  # 开启CPU侧预取
)

进阶GPU异步预取

如果需要实现和tf.data.prefetch()完全一致的、GPU计算与数据传输并行的效果,可以用轻量包装类实现CUDA流异步预取,不需要修改原有数据集或训练循环逻辑:

import torch
from torch_geometric.loader import DataLoader

class PrefetchLoader:
    def __init__(self, base_loader, target_device):
        self.loader = base_loader
        self.device = target_device
        # 开辟独立CUDA流负责数据传输,和训练计算流并行
        self.copy_stream = torch.cuda.Stream(device=target_device)
        self.next_batch = None

    def _preload_batch(self):
        try:
            self.next_batch = next(self.loader_iter)
        except StopIteration:
            self.next_batch = None
            return
        # 在独立传输流中异步将数据拷贝到GPU
        with torch.cuda.stream(self.copy_stream):
            self.next_batch = self.next_batch.to(self.device, non_blocking=True)

    def __iter__(self):
        self.loader_iter = iter(self.loader)
        self._preload_batch()
        while self.next_batch is not None:
            # 等待当前批次数据拷贝完成
            torch.cuda.current_stream().wait_stream(self.copy_stream)
            current_batch = self.next_batch
            # 处理当前批次时,异步预取下一批数据
            self._preload_batch()
            yield current_batch
    
    def __len__(self):
        return len(self.loader)

使用方式和普通Loader完全一致:

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# 先配置好多进程、CPU预取的基础Loader
base_loader = DataLoader(
    dataset,
    batch_size=32,
    shuffle=True,
    num_workers=8,
    pin_memory=True,
    persistent_workers=True,
    prefetch_factor=2
)
# 包装为GPU异步预取Loader
train_loader = PrefetchLoader(base_loader, device)

# 训练循环无需修改,直接迭代即可
for epoch in range(total_epochs):
    for batch in train_loader:
        # batch已经提前加载到GPU,直接执行前向计算即可
        pred = model(batch)
        # 反向传播、参数更新逻辑保持不变
        loss = criterion(pred, batch.y)
        loss.backward()
        optimizer.step()
        optimizer.zero_grad()

调优提示:尽量保证单样本处理耗时均匀,不要在get()方法中加入无法并行的全局操作,正常配置下多进程+预取的组合可以将数据加载等待的开销降到接近0,和原生tf.data流水线性能持平。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 10:27:23