使用 Transformers 进行图分类
Hugging Face 提供了一份关于使用 Transformers 库实现图分类的技术指南。库中目前可用于此任务的主要模型是微软的 Graphormer,它允许用户将基于 transformer 的架构应用于图结构数据,以完成二分类、多分类或回归等任务。
图数据格式化与加载
Hugging Face Hub 上的图数据集主要以 jsonl 格式存储,作为图的列表。每个图以字典形式表示,包含以下必需和可选字段:
edge_index:一个列表,包含两个平行的整数列表,表示边中节点的索引(例如,[[1, 1, 3], [2, 3, 1]]表示边 1–2、1–3 和 3–1)。num_nodes:一个整数,指示图中节点的总数,确保孤立节点也被计入。y:预测的目标值,可以是整数(用于多分类)、浮点数(用于回归),或二进制标签列表(用于多任务分类)。node_feat(可选):一个整数列表的列表,包含每个节点的特征,按节点索引顺序排列。edge_attr(可选):一个整数列表的列表,包含每条边的属性,遵循edge_index的顺序。
用户可以使用 datasets 库直接从 Hub 加载这些数据集,例如来自 Open Graph Benchmark (OGB) 的 ogbg-molhiv 数据集。
Graphormer 的预处理
图 transformer 框架需要特定的预处理,以生成有助于学习过程的特征。对于 Graphormer,这包括生成节点之间的最短路径矩阵以及入度/出度信息。
预处理可以通过两种方式处理:
- Static Preprocessing:使用
dataset.map(preprocess_item, batched=False)在训练前处理数据。 - On-the-fly Processing:在
GraphormerDataCollator中设置on_the_fly_processing=True,在训练期间处理数据,这对大图更节省内存。
模型实现与微调
图分类使用 GraphormerForGraphClassification 类实现。工作流程包括加载预训练检查点并将分类头适配到特定的下游任务。
加载与配置
通过 from_pretrained 加载模型时,用户可以指定 num_classes 以匹配其特定任务(例如,二分类时 num_classes=2)。将 ignore_mismatched_sizes=True 设置为 true,库会创建自定义分类头,替换预训练检查点中的原始解码头。
训练工作流
训练过程通过 Trainer API 管理。关键技术考虑包括:
- Data Collator:必须将
GraphormerDataCollator传递给Trainer,以将单个图转换为批次。 - Memory Management:由于模型体积,需要调整
per_device_train_batch_size和gradient_accumulation_steps,以防止内存不足(OOM)错误。 - Hardware:虽然可以在 CPU 上训练,但指南指出,在 Intel Core i7 上训练 20 个 epoch 大约需要一天,强烈建议在实际使用中采用 GPU 加速。