使用 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,这包括生成节点之间的最短路径矩阵以及入度/出度信息。

预处理可以通过两种方式处理:

  1. Static Preprocessing:使用 dataset.map(preprocess_item, batched=False) 在训练前处理数据。
  2. 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_sizegradient_accumulation_steps,以防止内存不足(OOM)错误。
  • Hardware:虽然可以在 CPU 上训练,但指南指出,在 Intel Core i7 上训练 20 个 epoch 大约需要一天,强烈建议在实际使用中采用 GPU 加速。

Sources