0Pricing
Learn AI with Python · 课时

链接预测与图分类

边预测任务、负采样、图级池化,以及用于图分类的 GINConv。

链接预测与图分类 是 CoddyKit 上的免费 Learn AI with Python 课时。 这是第 4 节课,共 4 节。 你可以在下方免费阅读本课时的完整内容 — 然后在浏览器中使用内置代码编辑器和全天候 AI 导师进行实践。 这是 Learn AI with Python 学习路径的一部分,你的进度在网页和 CoddyKit 应用中同步。 Learn AI with Python 课程共包含 4 节课。

两个新的图任务

除了节点分类之外,GNN 还可以处理:

  • 链接预测:两个节点之间是否会存在一条边?(好友推荐、药物相互作用)
  • 图分类:为整个图分配一个标签(这个分子有毒吗?)

链接预测设置

在链接预测中,我们首先使用 GNN 计算节点嵌入表示,然后对候选节点对进行评分。高分表示模型认为这两个节点之间应该存在一条边。

为边评分

一种常见的边评分方法是两个节点嵌入表示的点积:score = dot(h_u, h_v)。相似的嵌入表示会产生较高的点积,因此可以预测这两个节点很可能相连。

h = gnn(data.x, data.edge_index)        # node embeddings
score = (h[u] * h[v]).sum(dim=-1)        # dot product per pair

负采样

图只列出实际存在的边(正样本)。要训练分类器,我们还需要不存在的边。负采样会随机选取没有连接的节点对作为负样本,从而平衡训练集。

from torch_geometric.utils import negative_sampling

neg_edge_index = negative_sampling(
    edge_index=data.edge_index,
    num_nodes=data.num_nodes,
    num_neg_samples=data.edge_index.size(1),
)

带逻辑值的二元交叉熵损失

链接预测是二元任务(有边或无边)。我们为正样本对和负样本对打上 1 和 0 的标签,并使用 BCEWithLogitsLoss进行训练。该损失以数值稳定的方式将 S 形函数与二元交叉熵结合起来。

import torch

pos = (h[pos_u] * h[pos_v]).sum(-1)
neg = (h[neg_u] * h[neg_v]).sum(-1)
scores = torch.cat([pos, neg])
labels = torch.cat([torch.ones_like(pos), torch.zeros_like(neg)])
loss = torch.nn.functional.binary_cross_entropy_with_logits(scores, labels)

切换到图分类

对于图分类,我们需要每个图对应一个向量,而不是每个节点对应一个向量。在 GNN 层生成节点嵌入表示后,我们将其池化为一个图级表示。

global_mean_pool

global_mean_pool会对一个图中的所有节点嵌入表示取平均,从而生成一个固定大小的向量,与图的大小无关。当多个图组成一个批次时,batch索引会告诉它每个节点属于哪个图。

from torch_geometric.nn import global_mean_pool

h = gnn(x, edge_index)               # [num_nodes, dim]
hg = global_mean_pool(h, batch)       # [num_graphs, dim]
logits = classifier(hg)

池化为何重要

池化使模型不受节点顺序和图大小的影响:两个同构图会产生相同的池化向量。平均池化很简单;求和池化和最大池化是另外两种选择,它们具有不同的敏感性。

GINConv

GINConv(图同构网络)是一种表达能力更强的卷积层。它使用 MLP 和求和聚合,专门用于最大化消息传递在图级任务中的判别能力。

from torch_geometric.nn import GINConv
import torch

mlp = torch.nn.Sequential(
    torch.nn.Linear(in_dim, hid),
    torch.nn.ReLU(),
    torch.nn.Linear(hid, hid),
)
conv = GINConv(mlp)

与 Weisfeiler-Leman 的联系

GIN 的设计目标是达到 Weisfeiler-Leman(WL)测试的能力。Weisfeiler-Leman 是一种用于区分非同构图的经典算法。许多更简单的 GNN 无法区分某些图,而 GIN 可以做到这一点,其能力上限与 WL 测试相当,因此非常适合图分类。

选择合适的工具

请根据任务匹配架构:

  • 链接预测:GNN 嵌入 + 点积评分 + 负采样 + BCE 损失
  • 图分类:使用 GINConv 等表达能力强的卷积层 + 全局池化 + 分类器

快速检查

请测试您对知识的掌握情况。

回顾

您已经学习了链接预测和图分类:

  • 边得分 = dot(h_u, h_v),使用负采样和BCEWithLogitsLoss进行训练
  • global_mean_pool将节点嵌入转换为图级向量
  • GINConv具有很强的表达能力,可以达到Weisfeiler-Leman 测试的能力

常见问题解答

「链接预测与图分类」课时是免费的吗?

是的 — 「链接预测与图分类」的完整文本可在网页上免费阅读。要进行交互式练习(内置代码编辑器和全天候 AI 导师)并解锁 Learn AI with Python 课程的其余内容,请升级到 CoddyKit PRO。 Learn AI with Python 课程共包含 4 节课。

「链接预测与图分类」这节课中我会学到什么?

边预测任务、负采样、图级池化,以及用于图分类的 GINConv。 你通过在浏览器中直接运行的动手代码来练习 Learn AI with Python,全天候 AI 导师会在你学习这节课的过程中回答你的问题。

学习 Learn AI with Python 需要有经验吗?

无需任何先前经验。CoddyKit 上的 Learn AI with Python 课程适合初学者到高级学习者,你可以从这里开始或从头开始,按照自己的节奏学习。 这是第 4 节课,共 4 节。

「链接预测与图分类」课时需要多长时间?

大多数 CoddyKit 课程大约需要 5–10 分钟。每节课都很精短且互动,所以你能稳步进步,并在网页和应用中从离开的地方继续。

我能在这节 Learn AI with Python 课中编写并运行代码吗?

能。每节 Learn AI with Python 课都包含内置代码编辑器,你可以在浏览器中直接编写并运行真实代码,并获得即时 AI 反馈 — 无需本地设置。

此课程中的所有课时

  1. 机器学习的图论基础
  2. 图卷积网络(GCN)
  3. 使用 GNN 进行节点分类
  4. 链接预测与图分类
← 返回 Learn AI with Python