news 2026/9/23 14:18:15

DGL GraphBolt 快速入门:用数据管道(DataPipe)搭建 GNN 训练 Dataloader

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
DGL GraphBolt 快速入门:用数据管道(DataPipe)搭建 GNN 训练 Dataloader
  • 人工智能
  • 机器学习
  • 深度学习
  • 图计算

【免费下载链接】dgl

Python package built to ease deep learning on graph, on top of existing DL frameworks.

项目地址:https://gitcode.com/gh_mirrors/dg/dgl
点击查看免费下载

GraphBolt 是 DGL 中面向大规模图训练的数据加载解决方案,本文以 examples/graphbolt/quickstart 下的两个完整示例为主线,讲解如何用dgl.graphbolt以声明式数据管道的方式搭建 GNN 训练所需的 Dataloader:一条流水线同时完成小批量切分、设备搬运、邻居采样、负采样与特征抓取。读完本文,你将掌握ItemSampler → copy_to → sample_neighbor → sample_uniform_negative → fetch_feature → DataLoader的完整组装方法,并能独立复现 2 层 GCN 节点分类与 2 层 GraphSAGE 链接预测两种经典训练范式。

引言:GraphBolt 解决什么问题

GraphBolt 的设计目标非常聚焦——提供创建 Dataloader 以训练图神经网络所需的全部组件(原文:"Graphbolt provides all you need to create a dataloader to train a Graph Neural Networks")。它并不重新发明模型层,而是把 GNN 训练中最耗时的数据侧流程(小批量生成、邻居采样、特征抓取、负采样、数据搬运)抽象为一条可组合的数据管道(DataPipe),让用户以链式调用的方式声明式地组装完整的数据加载链路,而不是像传统写法那样手写采样循环、手动管理 batch 与特征索引。

本教程配套两个可运行的示例,均位于仓库 examples/graphbolt/quickstart 目录:

  • node_classification.py:在 Cora 数据集上训练一个 2 层图卷积网络(GCN)完成节点分类;
  • link_prediction.py:在 Cora 数据集上训练一个 2 层 GraphSAGE 完成链接预测,并用 AUROC 评估。

两个示例共用同一条建模思路:先加载数据集,再用graphbolt的 DataPipe 拼装 Dataloader,最后迭代 Dataloader 完成训练与评估。下面逐一拆解。

运行准备

示例依赖dgl.graphbolt(随 DGL 一并提供)与 PyTorch,另外需要torchmetrics(节点分类示例的准确率计算)和torcheval(链接预测示例的 AUROC 计算)。进入示例目录后直接运行即可:

# 节点分类 python examples/graphbolt/quickstart/node_classification.py # 链接预测 python examples/graphbolt/quickstart/link_prediction.py

两个脚本都会自动通过gb.BuiltinDataset("cora").load()下载并加载 Cora 数据集,无需手工准备数据文件。脚本会自动检测硬件:优先使用cuda:0,否则回退到 CPU(torch.device("cuda:0" if torch.cuda.is_available() else "cpu"))。

节点分类快速入门:2 层 GCN

第一步:加载数据集

import dgl.graphbolt as gb dataset = gb.BuiltinDataset("cora").load()

gb.BuiltinDatasetOnDiskDataset的子类(定义见 python/dgl/graphbolt/impl/ondisk_dataset.py),负责从 AWS S3 下载内置数据集并加载为OnDiskDataset。加载完成后,数据集对象提供三个关键成员:

  • dataset.graphFusedCSCSamplingGraph,用于采样的图结构;
  • dataset.feature:特征存储,可按节点/边类型与特征名查询特征;
  • dataset.tasks:任务列表,Cora 数据集包含两个任务——tasks[0]是节点分类,tasks[1]是链接预测。

dataset.tasks[0].metadata["num_classes"]可以拿到类别数,作为模型输出维度;输入维度则由特征形状推断:

in_size = dataset.feature.size("node", None, "feat")[0] out_size = dataset.tasks[0].metadata["num_classes"]

第二步:组装 Dataloader 数据管道

这是 GraphBolt 的核心。create_dataloader函数以链式调用把四类 DataPipe 串成一条流水线:

def create_dataloader(dataset, itemset, device): # 1. 从 itemset 中采样种子节点,切成 batch datapipe = gb.ItemSampler(itemset, batch_size=16) # 2. 将 mini-batch 搬运到指定设备,供后续采样与训练使用 datapipe = datapipe.copy_to(device) # 3. 为种子节点采样邻居(两层,fanout 分别为 4 和 2) datapipe = datapipe.sample_neighbor(dataset.graph, fanouts=[4, 2]) # 4. 抓取采样到的节点的特征 datapipe = datapipe.fetch_feature( dataset.feature, node_feature_keys=["feat"] ) # 5. 实例化为 DataLoader return gb.DataLoader(datapipe)

各阶段职责如下:

  1. gb.ItemSampler(itemset, batch_size=16)ItemSet定义了"要遍历的样本是什么"(这里是训练/验证/测试集的种子节点),ItemSampler负责将其切成指定大小的小批量。它支持shuffledrop_lastseed(可复现的随机打乱种子)等参数,详见 python/dgl/graphbolt/item_sampler.py。
  2. .copy_to(device):把 mini-batch 搬运到采样与训练所在的设备。注意:当copy_to放在管道前部时,DataLoadernum_workers必须为 0(DataLoader文档明确说明多进程下不支持 CUDA 使用,见 python/dgl/graphbolt/dataloader.py)。
  3. .sample_neighbor(dataset.graph, fanouts=[4, 2]):邻居采样,fanouts的长度即采样层数,[4, 2]表示第 1 层为每个节点采样 4 个邻居、第 2 层采样 2 个。注意 fanout 顺序是从最外层到最内层(源码注释:"The fanout order is from the outermost layer to innermost layer")。采样输出是紧凑化(compacted)后的子图,每个 batch 对应一层一个SampledSubgraph。采样还可配置replace(是否有放回)、prob_name(按边权重采样)、deduplicate(跨层种子去重)等参数,见 python/dgl/graphbolt/impl/neighbor_sampler.py。
  4. .fetch_feature(dataset.feature, node_feature_keys=["feat"]):把采样到的节点对应的feat特征取出来,装入 mini-batch。
  5. gb.DataLoader(datapipe):把整条管道包装为可迭代的数据加载器,支持num_workers(多进程)、persistent_workersmax_uva_threads等参数,见 python/dgl/graphbolt/dataloader.py。

第三步:定义并训练 2 层 GCN

模型使用dgl.nn.GraphConv堆叠两层,中间加 ReLU:

class GCN(nn.Module): def __init__(self, in_size, out_size, hidden_size=16): super().__init__() self.layers = nn.ModuleList() self.layers.append(dglnn.GraphConv(in_size, hidden_size)) self.layers.append(dglnn.GraphConv(hidden_size, out_size)) def forward(self, blocks, x): hidden_x = x for layer_idx, (layer, block) in enumerate(zip(self.layers, blocks)): hidden_x = layer(block, hidden_x) is_last_layer = layer_idx == len(self.layers) - 1 if not is_last_layer: hidden_x = F.relu(hidden_x) return hidden_x

注意forward的输入blocks:它是 mini-batch 中按层组织的采样子图列表(与fanouts=[4, 2]对应,共 2 个 block),每一层卷积作用在对应的 block 上——这正是 GraphBolt Dataloader 直接产出的训练就绪格式。

训练循环极其简洁,迭代 Dataloader 即可拿到模型所需的全部张量:

for step, data in enumerate(dataloader): x = data.node_features["feat"] # 采样节点的特征 y = data.labels # 种子节点的真实标签 y_hat = model(data.blocks, x) # 前向 loss = F.cross_entropy(y_hat, y) # 损失 optimizer.zero_grad() loss.backward() optimizer.step()

每个 epoch 结束后,分别用task.validation_settask.test_set构建验证/测试 Dataloader 评估准确率(torchmetrics.functional.accuracytask="multiclass")。默认训练 10 个 epoch,学习率1e-2,输出格式为Epoch {epoch:03d} | Loss ... | Val Acc ... | Test Acc ...

链接预测快速入门:2 层 GraphSAGE

数据管道:训练与测试的差异

链接预测的数据管道与节点分类类似,但多了负采样种子边排除两个环节,且训练/测试时管道不同:

def create_dataloader(dataset, device, is_train=True): task = dataset.tasks[1] # 链接预测任务 itemset = task.train_set if is_train else task.test_set datapipe = gb.ItemSampler(itemset, batch_size=256) datapipe = datapipe.copy_to(device) if is_train: # 训练:为每条种子边采样 1 条负边 datapipe = datapipe.sample_uniform_negative( dataset.graph, negative_ratio=1 ) # 训练:两层邻居采样 datapipe = datapipe.sample_neighbor(dataset.graph, fanouts=[4, 2]) # 训练:从子图中剔除种子边,防止标签泄露 datapipe = datapipe.transform(gb.exclude_seed_edges) else: # 测试:全量邻居(fanout=-1 表示采样所有邻居) datapipe = datapipe.sample_neighbor(dataset.graph, fanouts=[-1, -1]) datapipe = datapipe.fetch_feature( dataset.feature, node_feature_keys=["feat"] ) return gb.DataLoader(datapipe)

三个训练专属环节的作用:

  1. sample_uniform_negative(dataset.graph, negative_ratio=1):对每条种子(正)边均匀采样negative_ratio条负边。UniformNegativeSampler的具体实现见 python/dgl/graphbolt/impl/uniform_negative_sampler.py。
  2. sample_neighbor(dataset.graph, fanouts=[4, 2]):与节点分类一致的两层邻居采样。链接预测中采样器会先从正负节点对中收集去重后的节点作为种子再采样(NeighborSampler文档对此有专门说明)。
  3. .transform(gb.exclude_seed_edges):把种子边从采样子图中剔除。exclude_seed_edges定义于 python/dgl/graphbolt/external_utils.py,它防止模型在消息传递阶段"看见"用于训练的正样本边,避免信息泄露。

测试时则用fanouts=[-1, -1]全量邻居采样,不做负采样、不排除种子边,以保证评估的完整性。

模型与损失

模型是两层SAGEConv(..., "mean")编码器加一个两层 MLP 打分器(predictor),对源/目标节点嵌入做逐元素乘积后输出单个 logit:

class GraphSAGE(nn.Module): def __init__(self, in_size, hidden_size=16): super().__init__() self.layers = nn.ModuleList() self.layers.append(SAGEConv(in_size, hidden_size, "mean")) self.layers.append(SAGEConv(hidden_size, hidden_size, "mean")) self.predictor = nn.Sequential( nn.Linear(hidden_size, hidden_size), nn.ReLU(), nn.Linear(hidden_size, 1), )

训练循环中通过data.compacted_seeds.T拿到紧凑化后的正负节点对索引,进而取出两端嵌入:

compacted_seeds = data.compacted_seeds.T y = model(data.blocks, x) logits = model.predictor( y[compacted_seeds[0].long()] * y[compacted_seeds[1].long()] ).squeeze() loss = F.binary_cross_entropy_with_logits(logits, labels.float())

评估

测试阶段用torcheval.metrics.BinaryAUROC计算 AUROC:把测试 Dataloader 产出的所有 logit 与标签拼接后一次性更新指标并compute()

注意:示例文件头部有免责声明——该示例没有从原始图中剔除测试边,可能存在数据泄露,作者明确表示"We are ignoring this issue for this example because we are focused on demonstrating usability",即仅用于演示 GraphBolt 的易用性,正式实验需自行处理边划分。

深入理解:数据管道背后的核心对象

MiniBatch:贯穿全流程的统一数据结构

MiniBatch(定义见 python/dgl/graphbolt/minibatch.py)是数据加载过程中所有阶段的输入输出统一结构,上面示例中反复使用的字段包括:

字段含义
seeds种子项:一维张量表示种子节点,二维张量每行表示一条边/超链接等
labels与种子对应的标签(节点分类时是节点类别,链接预测时是边标签)
sampled_subgraphs每层采样一个SampledSubgraph,按层组织
input_nodes最外层(所有采样层)涉及的全部输入节点
node_features抓取到的节点特征字典,键为特征名(异构图下为(节点类型, 特征名)元组)
compacted_seeds紧凑化后的种子,对应采样子图中的新索引
blockssampled_subgraphs转换得到的 DGL block 列表,可直接喂给 GNN 层

节点分类中只用到node_featureslabelsblocks;链接预测额外用到compacted_seeds做节点对索引。

ItemSet 与 ItemSampler

gb.ItemSet(见 python/dgl/graphbolt/itemset.py)是样本集合的轻量包装,支持四种形态:

  • 单个整数:等价于torch.arange生成的节点序列;
  • 单个张量:按第一维索引的节点集合;
  • 张量元组(同形状):如(节点, 标签)配对;
  • 张量元组(不同形状):如(边, 标签)配对,逐项对应切分。

每个 ItemSet 可以指定names(如"seeds""labels"),这些名字与MiniBatch的属性名对齐,从而让ItemSampler自动把切好的 batch 组装成带语义的MiniBatchItemSampler本身基于torch.utils.data.IterDataPipe实现,因此可以与 PyTorch 官方的任意迭代型 DataPipe 继续拼接(见 python/dgl/graphbolt/item_sampler.py 的说明)。

DataLoader 的多进程策略

gb.DataLoader(见 python/dgl/graphbolt/dataloader.py)会在num_workers > 0时自动改造数据管道:在ItemSampler后插入sharding_filter()实现各 worker 均匀分配 mini-batch,并在FeatureFetcher处"切断"管道——特征抓取之前的阶段(采样等)在子进程中执行,特征抓取之后的阶段回到主进程。这个设计避免了特征张量在进程间反复拷贝,是规模化训练时的关键性能点。

内置数据集

gb.BuiltinDatasetcora外还内置ogbn-magogbl-citation2ogbn-arxivogbn-papers100Mogbn-productsogb-lsc-mag240migb-hom系列、igb-het系列等(完整清单见 python/dgl/graphbolt/impl/ondisk_dataset.py),覆盖同质/异质、节点分类/链接预测等常见场景,且大部分大规模数据集在预处理时加入了反向边并去重。

GPU 训练:特征与图的固定内存(Pinned Memory)

两个示例都包含一段 GPU 专属优化代码:

if device == torch.device("cuda:0"): dataset.graph.pin_memory_() dataset.feature.pin_memory_()

将图结构与特征固定到内存(pinned memory)后,GPU 可以通过 UVA(Unified Virtual Addressing)直接访问它们,从而让采样与特征抓取在 GPU 端完成、无需逐批拷贝。配合NeighborSampleroverlap_fetch(用独立 CUDA 流重叠图抓取与其他运算)与num_gpu_cached_edges(GPU 缓存高频访问的顶点邻域,降低 PCIe 带宽压力)参数,可进一步压榨 GPU 数据管线性能。如果使用 CPU 训练,这两行会自然跳过。

小结:一条管道吃透两类任务

对比两个示例可以发现 GraphBolt 的统一抽象:节点分类与链接预测的差异,最终只体现在数据管道的两个节点上——链接预测多接一个sample_uniform_negative(负采样)和一个exclude_seed_edges(防泄露),模型侧用compacted_seeds取节点对。其余环节(ItemSampler 切批、copy_to 搬运、sample_neighbor 采样、fetch_feature 抓特征、DataLoader 迭代)完全一致。

这正是 GraphBolt 的设计哲学:把 GNN 训练中高频复用的数据侧逻辑沉淀为可组合的标准件,让开发者从"手写采样循环 + 特征索引 + batch 拼接"中解放出来,把精力集中在模型与实验本身。若需深入阅读,可继续探索:

  • 邻居采样实现:python/dgl/graphbolt/impl/neighbor_sampler.py
  • 负采样实现:python/dgl/graphbolt/impl/uniform_negative_sampler.py
  • 特征抓取:python/dgl/graphbolt/feature_fetcher.py
  • 数据加载器:python/dgl/graphbolt/dataloader.py
  • 人工智能
  • 机器学习
  • 深度学习
  • 图计算

【免费下载链接】dgl

Python package built to ease deep learning on graph, on top of existing DL frameworks.

项目地址:https://gitcode.com/gh_mirrors/dg/dgl
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/23 14:18:15

2026最新怎么样哄女朋友代码性能优化实战指南

2026最新怎么样哄女朋友代码性能优化实战指南 面试被问原理答不上来,是不是让你瞬间大脑空白?别慌,2026最新的实战案例里,连“怎么样哄女朋友”这种生活化场景都能变成代码优化的绝佳载体。 性能瓶颈:为什么你的“哄法”这么慢…

作者头像 李华
网站建设 2026/9/23 14:18:10

一文搞懂如何进行视频剪辑:面试突击与避坑指南

一文搞懂如何进行视频剪辑:面试突击与避坑指南 版本升级后 API 全变了?别慌。很多开发者一提到 如何进行视频剪辑 ,脑子里全是 ffmpeg 命令行或者 After Effects 的操作界面,但在编程面试中,这往往考察的是对媒体处理流水线、流式处理 API…

作者头像 李华
网站建设 2026/9/23 14:17:51

关于教育孩子的文章:手写实现3种架构避开新手坑

关于教育孩子的文章:手写实现3种架构避开新手坑 别急着背八股文,你现在的困境很典型:语法书翻烂了,变量、循环、类都会写,但一让你搭个像样的项目,脑子瞬间空白。这种“会语法不会工程”的断崖式落差,是绝大多数初学者甚至转行者的死穴。 很多教程只教你怎么用 print…

作者头像 李华
网站建设 2026/9/23 14:17:48

实战项目避坑:价格表设计3个死穴一次讲透

实战项目避坑:价格表设计3个死穴一次讲透 昨天凌晨三点,我还在帮一个做装修报价系统的哥们修 Bug。他盯着屏幕问我:“为什么加了个折扣字段,整个数据库索引全挂了?” 别笑,这场景太常见了。很多刚接触后端开发的朋友,在搭建 实战项目 时,一上来就照着网上那些“高大上”的范式去建表。结果呢?…

作者头像 李华
网站建设 2026/9/23 14:17:37

3个技巧一文搞懂魔力宝贝论坛后端源码逻辑

3个技巧一文搞懂魔力宝贝论坛后端源码逻辑 满屏的 NullPointerException 和 StackOverflowError 堆栈,看着就头大?很多开发者接手“魔力宝贝论坛”这类经典社区项目时,第一反应就是懵:这代码到底哪错了?别慌,今天咱们不整虚的,直接扒开这个项目的核心源码, 一文搞懂…

作者头像 李华
网站建设 2026/9/23 14:17:26

3分钟吃透correspond源码:报错不再懵的速查手册

3分钟吃透correspond源码:报错不再懵的速查手册 盯着满屏的 TypeError 和 ReferenceError ,StackTrace 里全是陌生的文件名和行号,你是不是也懵了?别慌,今天这篇 correspond 源码速查手册,专治各种“报错看不懂”的疑难杂症。…

作者头像 李华