案例 16:OGBG-MolHIV 图分类——图同构网络与图级读出#

学习目标与记号#

  节点分类为每个节点输出标签;图分类则需要把节点数量不同的整张图汇总成固定长度的表示。本案例以分子图二分类为背景,介绍图同构网络(graph isomorphism network,GIN)的邻居求和聚合,以及图级求和、均值和最大值读出。本案例的学习目标包括:

  1. 区分节点特征、边、图编号和图标签,核对批量中各个量的维度;

  2. 从多重集合聚合角度解释 GIN: \(\boldsymbol{h}_v^{[l]}=\operatorname{MLP}^{[l]}((1+\epsilon^{[l]})\boldsymbol{h}_v^{[l-1]}+\sum_{u\in\mathcal{N}(v)}\boldsymbol{h}_u^{[l-1]})\)

  3. 实现不依赖 PyG 的断开图批量、GIN 层和置换不变图读出;

  4. 用“忽略边的节点袋模型”和训练集阳性率作为基线;

  5. 保留官方分子骨架(scaffold)划分,用受试者工作特征曲线下面积(area under the receiver operating characteristic curve,ROC-AUC)评价不平衡二分类,并辨认结构泄漏。

  这是算法案例,不构成药物、安全或临床判断。HIV 标签来自特定实验与处理流程,不能外推到人体疗效。

数据来源、许可与运行模式#

  默认使用快速模式:在官方训练集、验证集和测试集的分子骨架划分内固定随机选取一部分图,不根据标签挑选样本。将代码变量 FAST_MODE 设为 False 后使用全部图,建议准备 3--6 GB 内存并使用图形处理器(graphics processing unit,GPU)。数据缓存在 AI_COURSE_DATA_DIR 指定的目录中;未设置该变量时使用用户缓存目录。

  分子骨架划分尽量让不同集合具有不同的分子骨架,比随机拆分更能检验模型面对新化学结构时的表现;快速模式的结果不是 OGB 排行榜成绩。

from pathlib import Path
import copy
import gzip
import hashlib
import os
import random
import urllib.request
import zipfile

import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
import torch
from sklearn.metrics import average_precision_score, roc_auc_score

SEED = 42
FAST_MODE = os.getenv("AI_COURSE_FAST_MODE", os.getenv("FAST_MODE", "1")) != "0"
random.seed(SEED)
np.random.seed(SEED)
torch.manual_seed(SEED)

cache_root = Path(
    os.environ.get(
        "AI_COURSE_DATA_DIR",
        Path(os.environ.get("XDG_CACHE_HOME", Path.home() / ".cache")) / "ai-course-cases",
    )
).expanduser()
cache_root.mkdir(parents=True, exist_ok=True)
device = torch.device("cuda" if (not FAST_MODE and torch.cuda.is_available()) else "cpu")
print({"快速模式": FAST_MODE, "计算设备": str(device), "缓存目录": str(cache_root)})
{'快速模式': True, '计算设备': 'cpu', '缓存目录': '/private/tmp/ai-course-case-data'}

第 1 步:从 OGB 压缩包读取图数据#

  OGB 的逗号分隔值(comma-separated values,CSV)版本把所有图的节点特征和边按顺序放在同一组文件中,并用 num-node-listnum-edge-list 记录每张图占用的行数。读取后必须用累计和恢复各张图的边界,不能把两张不同分子图中的节点误连。

  • node-feat:\((\sum_g n_g,d_x)\)每行一个原子的离散属性编码;

  • edge:\((\sum_g m_g,2)\)每张图内部的局部端点;

  • graph-label:\((G,1)\)

  • scaffold split:图编号,而不是节点编号。

  下载使用临时文件和原子替换,并显示 SHA-256 以记录数据快照。

DATA_URL = "https://snap.stanford.edu/ogb/data/graphproppred/csv_mol_download/hiv.zip"
archive = cache_root / "ogbg_molhiv.zip"

def download_cached(url, target):
    if not target.exists():
        temporary = target.with_suffix(target.suffix + ".part")
        request = urllib.request.Request(url, headers={"User-Agent": "ai-course-case/1.0"})
        with urllib.request.urlopen(request, timeout=120) as response, temporary.open("wb") as out:
            while chunk := response.read(1024 * 1024):
                out.write(chunk)
        temporary.replace(target)
    digest = hashlib.sha256(target.read_bytes()).hexdigest()
    size_mib = round(target.stat().st_size / 2**20, 1)
    print(f"文件:{target.name};大小:{size_mib} 兆字节;SHA-256:{digest}")
    return target

def find_member(names, suffix):
    matches = [name for name in names if name.endswith(suffix)]
    if len(matches) != 1:
        raise FileNotFoundError(f"需要唯一的 {suffix},实际找到 {matches[:5]}")
    return matches[0]

def read_gzip_csv(zf, suffix, dtype):
    with zf.open(find_member(zf.namelist(), suffix)) as raw, gzip.GzipFile(fileobj=raw) as uncompressed:
        return pd.read_csv(uncompressed, header=None).to_numpy(dtype=dtype)

download_cached(DATA_URL, archive)
with zipfile.ZipFile(archive) as zf:
    node_feat_all = read_gzip_csv(zf, "/raw/node-feat.csv.gz", np.int64)
    edge_rows_all = read_gzip_csv(zf, "/raw/edge.csv.gz", np.int64)
    node_counts = read_gzip_csv(zf, "/raw/num-node-list.csv.gz", np.int64).reshape(-1)
    edge_counts = read_gzip_csv(zf, "/raw/num-edge-list.csv.gz", np.int64).reshape(-1)
    graph_labels = read_gzip_csv(zf, "/raw/graph-label.csv.gz", np.float32).reshape(-1)
    split_global = {
        "train": read_gzip_csv(zf, "/split/scaffold/train.csv.gz", np.int64).reshape(-1),
        "valid": read_gzip_csv(zf, "/split/scaffold/valid.csv.gz", np.int64).reshape(-1),
        "test": read_gzip_csv(zf, "/split/scaffold/test.csv.gz", np.int64).reshape(-1),
    }

if edge_rows_all.shape[1] != 2 and edge_rows_all.shape[0] == 2:
    edge_rows_all = edge_rows_all.T
assert node_counts.sum() == len(node_feat_all)
assert edge_counts.sum() == len(edge_rows_all)
assert len(node_counts) == len(graph_labels)
print("图数、节点特征、边:", len(node_counts), node_feat_all.shape, edge_rows_all.shape)
official_split_sizes = {"训练集": len(split_global["train"]),
                        "验证集": len(split_global["valid"]),
                        "测试集": len(split_global["test"])}
print("官方划分:", official_split_sizes)
文件:ogbg_molhiv.zip;大小:2.0 兆字节;SHA-256:47d747664b9e1653de5aac99bf26c015d88a3daa474f5a32f3fe5c014111375b
图数、节点特征、边: 41127 (1049163, 9) (1129688, 2)
官方划分: {'训练集': 32901, '验证集': 4113, '测试集': 4113}

第 2 步:检查数据与任务难点#

  MolHIV 的阳性样本较少,即使把所有样本都预测为阴性,也可能得到较高的准确率。因此,主要指标采用 ROC-AUC,并补充平均精确率(average precision,AP),后者更关注阳性样本的排序。还要观察每张图的节点数和边数;过大的图会占用更多内存,也会放大求和读出结果的数值范围。

  我们按官方集合分别报告阳性率,但不据此重排或筛选测试样本。测试分布只能用于描述,不能用来选择阈值、读出方式或训练轮数。

node_starts = np.r_[0, np.cumsum(node_counts[:-1])]
edge_starts = np.r_[0, np.cumsum(edge_counts[:-1])]
audit = []
split_names = {"train": "训练集", "valid": "验证集", "test": "测试集"}
for name, indices in split_global.items():
    audit.append({
        "数据集合": split_names[name],
        "图数量": len(indices),
        "阳性率": float(np.nanmean(graph_labels[indices])),
        "节点数中位数": float(np.median(node_counts[indices])),
        "节点数第95百分位": float(np.quantile(node_counts[indices], 0.95)),
        "有向边数中位数": float(np.median(edge_counts[indices])),
    })
display(pd.DataFrame(audit))
assert set(np.unique(graph_labels[~np.isnan(graph_labels)])).issubset({0.0, 1.0})
数据集合 图数量 阳性率 节点数中位数 节点数第95百分位 有向边数中位数
0 训练集 32901 0.037446 23.0 45.0 25.0
1 验证集 4113 0.019694 25.0 50.0 27.0
2 测试集 4113 0.031607 22.0 48.0 25.0

第 3 步:保留官方分子骨架划分并创建断开图批量#

  快速模式只在每个官方集合内部抽样;没有重新随机划分,也不使用标签做分层。随后把选中的图串成一个“大而不连通”的批:

  • node_features:\((N_{\text{sub}},d_x)\)

  • edge_index:\((2,M_{\text{sub}})\)端点加上各图节点偏移;

  • graph_id:\((N_{\text{sub}},)\)说明每个节点属于哪张图;

  • y_graph:\((G_{\text{sub}},)\)

  不同图之间没有边,所以一次前向传播仍等价于逐图传播;我们也避免 BatchNorm,以免让测试图统计量影响训练图。

rng = np.random.default_rng(SEED)
caps = {"train": 1500, "valid": 400, "test": 400}
if FAST_MODE:
    selected = {
        name: np.sort(rng.choice(ids, size=min(caps[name], len(ids)), replace=False))
        for name, ids in split_global.items()
    }
else:
    selected = {name: ids.copy() for name, ids in split_global.items()}

selected_global = np.concatenate([selected["train"], selected["valid"], selected["test"]])
global_to_local_graph = {int(g): i for i, g in enumerate(selected_global)}
feature_parts, edge_parts, graph_parts = [], [], []
new_node_offset = 0

for local_graph, global_graph in enumerate(selected_global):
    n = int(node_counts[global_graph])
    m = int(edge_counts[global_graph])
    ns, es = int(node_starts[global_graph]), int(edge_starts[global_graph])
    features = node_feat_all[ns:ns+n]
    edges = edge_rows_all[es:es+m].copy()
    if len(edges) and edges.max() >= n:
        if edges.min() >= ns and edges.max() < ns + n:
            edges -= ns
        else:
            raise ValueError(f"图 {global_graph} 的边端点超出节点范围")
    feature_parts.append(features)
    edge_parts.append(edges.T + new_node_offset)
    graph_parts.append(np.full(n, local_graph, dtype=np.int64))
    new_node_offset += n

node_features = torch.as_tensor(np.concatenate(feature_parts), dtype=torch.long, device=device)
edge_index = torch.as_tensor(np.concatenate(edge_parts, axis=1), dtype=torch.long, device=device)
graph_id = torch.as_tensor(np.concatenate(graph_parts), dtype=torch.long, device=device)
y_graph = torch.as_tensor(graph_labels[selected_global], dtype=torch.float32, device=device)
split_local = {
    name: torch.as_tensor([global_to_local_graph[int(g)] for g in ids], dtype=torch.long, device=device)
    for name, ids in selected.items()
}
print({"图数量": len(selected_global), "节点数": len(node_features), "边数": edge_index.shape[1]})
print("各个量的维度:", node_features.shape, edge_index.shape, graph_id.shape, y_graph.shape)
{'图数量': 2300, '节点数': 59590, '边数': 64492}
各个量的维度: torch.Size([59590, 9]) torch.Size([2, 64492]) torch.Size([59590]) torch.Size([2300])

第 4 步:离散原子属性编码与图读出#

  原子属性列是类别编码,不应把“类别编号相差 2”解释为物理距离 2。我们为每列建立嵌入表,并把各字段嵌入相加,得到 \(\boldsymbol{H}^{[0]}\in\mathbb{R}^{N\times d_h}\)嵌入表大小只依据公开特征编码范围,不读取图标签。

  图读出必须对节点排列不变。求和、均值、最大值读出都满足

\[ \operatorname{READOUT}(\{\boldsymbol{h}_{\pi(1)},\ldots,\boldsymbol{h}_{\pi(n)}\}) =\operatorname{READOUT}(\{\boldsymbol{h}_1,\ldots,\boldsymbol{h}_n\}). \]

  求和保留规模信息,均值消除节点数尺度,最大值只保留每维最强响应;哪种更好是需要验证集回答的建模问题。

cardinalities = (node_features.max(dim=0).values + 1).tolist()

class AtomEncoder(torch.nn.Module):
    def __init__(self, cardinalities, hidden_dim):
        super().__init__()
        self.embeddings = torch.nn.ModuleList(
            [torch.nn.Embedding(int(size), hidden_dim) for size in cardinalities]
        )
        for embedding in self.embeddings:
            torch.nn.init.xavier_uniform_(embedding.weight)

    def forward(self, x):
        parts = [embedding(x[:, j]) for j, embedding in enumerate(self.embeddings)]
        return torch.stack(parts, dim=0).sum(dim=0)

def graph_readout(h, graph_id, num_graphs, mode):
    out = torch.zeros((num_graphs, h.shape[1]), device=h.device, dtype=h.dtype)
    if mode in {"sum", "mean"}:
        out.index_add_(0, graph_id, h)
        if mode == "mean":
            counts = torch.bincount(graph_id, minlength=num_graphs).clamp_min(1)
            out = out / counts[:, None]
    elif mode == "max":
        out.fill_(-torch.inf)
        index = graph_id[:, None].expand_as(h)
        out.scatter_reduce_(0, index, h, reduce="amax", include_self=True)
        out = torch.nan_to_num(out, neginf=0.0)
    else:
        raise ValueError(mode)
    return out

toy_h = torch.arange(12, dtype=torch.float32, device=device).reshape(4, 3)
toy_gid = torch.tensor([0, 0, 1, 1], device=device)
perm = torch.tensor([1, 0, 3, 2], device=device)
for mode in ["sum", "mean", "max"]:
    assert torch.allclose(
        graph_readout(toy_h, toy_gid, 2, mode),
        graph_readout(toy_h[perm], toy_gid[perm], 2, mode),
    )
print("读出置换不变检查通过;字段基数:", cardinalities)
读出置换不变检查通过;字段基数: [80, 1, 10, 9, 4, 1, 6, 2, 2]

第 5 步:两个无结构基线#

  第一条基线是训练集阳性率,它不给图排序能力,ROC-AUC 理论上是 0.5;它提醒我们不能只用准确率掩盖类别不平衡。

  第二条是节点袋模型:编码原子后直接读出,不使用 edge_index。它保留“有哪些原子及其数量”,却不知道原子如何连接。GIN 相对它的增益才更接近结构贡献,而不是仅仅相对一个过弱常数。

train_ids = split_local["train"]
train_rate = float(y_graph[train_ids].mean().item())
neg = float((y_graph[train_ids] == 0).sum().item())
pos = float((y_graph[train_ids] == 1).sum().item())
pos_weight = torch.tensor(neg / max(pos, 1.0), device=device)
print({"训练集阳性率": train_rate, "阳性损失权重": float(pos_weight.item())})

class BagOfAtoms(torch.nn.Module):
    def __init__(self, cardinalities, hidden_dim, readout="sum"):
        super().__init__()
        self.encoder = AtomEncoder(cardinalities, hidden_dim)
        self.readout = readout
        self.head = torch.nn.Sequential(
            torch.nn.Linear(hidden_dim, hidden_dim),
            torch.nn.ReLU(),
            torch.nn.Dropout(0.25),
            torch.nn.Linear(hidden_dim, 1),
        )

    def forward(self, x, edges, graph_id, num_graphs):
        h = self.encoder(x)
        return self.head(graph_readout(h, graph_id, num_graphs, self.readout)).squeeze(1)
{'训练集阳性率': 0.04266666620969772, '阳性损失权重': 22.4375}

第 6 步:图同构网络的多重集合聚合#

  普通均值聚合会丢失重复次数:集合 \(\{a,b\}\)\(\{a,a,b,b\}\) 均值相同。GIN 使用求和,并让多层感知机(MLP)学习组合,在有限条件下能更有力地区分邻居多重集合。自身表示乘 \(1+\epsilon\)使中心节点与邻居和不被强制等权。

  代码中的 src、dst 表示消息从 src 发往 dst;index_add 对相同 dst 求和。原始分子边通常用两个方向存储,因此每个原子能接收相邻原子的消息。为聚焦 GIN,本简化模型未编码键类型,后面会把它列为局限。

class GINLayer(torch.nn.Module):
    def __init__(self, hidden_dim):
        super().__init__()
        self.epsilon = torch.nn.Parameter(torch.zeros(()))
        self.mlp = torch.nn.Sequential(
            torch.nn.Linear(hidden_dim, 2 * hidden_dim),
            torch.nn.ReLU(),
            torch.nn.Linear(2 * hidden_dim, hidden_dim),
        )

    def forward(self, h, edge_index):
        src, dst = edge_index
        neighbor_sum = torch.zeros_like(h)
        neighbor_sum.index_add_(0, dst, h[src])
        return torch.relu(self.mlp((1.0 + self.epsilon) * h + neighbor_sum))

class GINClassifier(torch.nn.Module):
    def __init__(self, cardinalities, hidden_dim, layers=3, readout="sum"):
        super().__init__()
        self.encoder = AtomEncoder(cardinalities, hidden_dim)
        self.layers = torch.nn.ModuleList([GINLayer(hidden_dim) for _ in range(layers)])
        self.readout = readout
        self.head = torch.nn.Sequential(
            torch.nn.Linear(hidden_dim, hidden_dim),
            torch.nn.ReLU(),
            torch.nn.Dropout(0.25),
            torch.nn.Linear(hidden_dim, 1),
        )

    def forward(self, x, edges, graph_id, num_graphs):
        h = self.encoder(x)
        for layer in self.layers:
            h = layer(h, edges)
        graph_h = graph_readout(h, graph_id, num_graphs, self.readout)
        return self.head(graph_h).squeeze(1)

shape_model = GINClassifier(cardinalities, 24, layers=2).to(device)
with torch.no_grad():
    shape_logits = shape_model(node_features, edge_index, graph_id, len(selected_global))
assert shape_logits.shape == (len(selected_global),)
print("从节点特征到图预测结果的维度:", tuple(node_features.shape), "→", tuple(shape_logits.shape))
从节点特征到图预测结果的维度: (59590, 9) → (2300,)

第 7 步:训练时只让训练图产生梯度#

  二元交叉熵直接根据输出层的线性运算结果计算,避免先计算 sigmoid 函数再取对数带来的数值问题。阳性权重 pos_weight 只根据训练标签计算,使数量较少的阳性样本在损失函数中获得更大权重。验证集的 ROC-AUC 用于确定早停时刻;在确定模型与训练方案前,不查看测试指标。

  所有图虽在一个断开批中前向传播,但没有跨图边,且无 BatchNorm,所以测试图不会改变训练图表示。

def split_metrics(logits, indices):
    truth = y_graph[indices].detach().cpu().numpy()
    score = torch.sigmoid(logits[indices]).detach().cpu().numpy()
    if len(np.unique(truth)) < 2:
        raise RuntimeError("当前抽样集合只有一个类别,无法计算受试者工作特征曲线下面积;请增大快速模式样本上限")
    return {
        "roc_auc": roc_auc_score(truth, score),
        "average_precision": average_precision_score(truth, score),
    }

def fit(model, epochs, lr=2e-3):
    model = model.to(device)
    optimizer = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=1e-5)
    criterion = torch.nn.BCEWithLogitsLoss(pos_weight=pos_weight)
    best_state, best_auc, stale = None, -1.0, 0
    records = []
    for epoch in range(epochs):
        model.train()
        optimizer.zero_grad()
        logits = model(node_features, edge_index, graph_id, len(selected_global))
        loss = criterion(logits[train_ids], y_graph[train_ids])
        loss.backward()
        torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)
        optimizer.step()

        model.eval()
        with torch.no_grad():
            logits = model(node_features, edge_index, graph_id, len(selected_global))
            valid = split_metrics(logits, split_local["valid"])
        records.append((epoch, float(loss.item()), valid["roc_auc"], valid["average_precision"]))
        if valid["roc_auc"] > best_auc + 1e-4:
            best_auc = valid["roc_auc"]
            best_state = copy.deepcopy(model.state_dict())
            stale = 0
        else:
            stale += 1
        if stale >= 8:
            break
    model.load_state_dict(best_state)
    return model, pd.DataFrame(records, columns=["训练轮次", "训练损失", "验证ROC-AUC", "验证平均精确率"])

epochs = 22 if FAST_MODE else 100
hidden = 48 if FAST_MODE else 128

第 8 步:比较边结构与读出方式#

  在其他条件相同的情况下,每次只改变一个因素进行比较:

  1. 节点袋模型与 GIN:比较是否使用边;

  2. GIN 的求和、均值和最大值读出:比较如何把节点集合变成图向量。

  其他训练协议保持相同。严格研究应为多个随机种子报告均值与标准差;这里单种子用于控制课堂运行时间。

factories = {
    "节点袋(求和读出、无边)": lambda: BagOfAtoms(cardinalities, hidden, "sum"),
    "图同构网络(求和读出)": lambda: GINClassifier(cardinalities, hidden, layers=3, readout="sum"),
    "图同构网络(均值读出)": lambda: GINClassifier(cardinalities, hidden, layers=3, readout="mean"),
    "图同构网络(最大值读出)": lambda: GINClassifier(cardinalities, hidden, layers=3, readout="max"),
}
trained, histories = {}, {}
for name, factory in factories.items():
    torch.manual_seed(SEED)
    trained[name], histories[name] = fit(factory(), epochs)
    print(name, "最佳验证ROC-AUC", histories[name]["验证ROC-AUC"].max())

fig, axes = plt.subplots(1, 2, figsize=(10, 3.5))
for name, history in histories.items():
    axes[0].plot(history["训练轮次"], history["训练损失"], label=name)
    axes[1].plot(history["训练轮次"], history["验证ROC-AUC"], label=name)
axes[0].set(title="训练加权二元交叉熵", xlabel="训练轮次")
axes[1].set(title="验证ROC-AUC", xlabel="训练轮次", ylim=(0, 1))
axes[1].legend(fontsize=8)
plt.tight_layout()
节点袋(求和读出、无边) 最佳验证ROC-AUC 0.74448419797257
图同构网络(求和读出) 最佳验证ROC-AUC 0.6575233551977738
图同构网络(均值读出) 最佳验证ROC-AUC 0.6789902603856093
图同构网络(最大值读出) 最佳验证ROC-AUC 0.40946133969389786
../../_images/9b1fe1bbd848e2a137c0f75ce18807fefdeda9610e0cb670d022a4aa5f70b077.png

第 9 步:确定模型与训练方案后评价测试集#

  ROC-AUC 衡量随机选取一个阳性样本和一个阴性样本时,模型把阳性样本排在阴性样本之前的概率,它不依赖某一个分类阈值;AP 对阳性样本较少的情况更敏感,其随机水平接近阳性样本所占比例,而不是 0.5。二者都不能说明模型输出的概率已经校准。

  常数基线不能产生排序,因此 AUC 为 0.5,AP 约为测试阳性率。若 GIN 优于节点袋模型,才有证据表明连接结构带来额外信息;若不同读出方法的结果差异明显,说明图大小与极值信息的处理值得关注。

test_ids = split_local["test"]
test_truth = y_graph[test_ids].cpu().numpy()
constant_score = np.full(len(test_truth), train_rate)
rows = [{
    "模型": "训练阳性率常数",
    "测试ROC-AUC": roc_auc_score(test_truth, constant_score),
    "测试平均精确率": average_precision_score(test_truth, constant_score),
    "参数量": 0,
}]
for name, model in trained.items():
    model.eval()
    with torch.no_grad():
        logits = model(node_features, edge_index, graph_id, len(selected_global))
    metrics = split_metrics(logits, test_ids)
    rows.append({
        "模型": name,
        "测试ROC-AUC": metrics["roc_auc"],
        "测试平均精确率": metrics["average_precision"],
        "参数量": sum(p.numel() for p in model.parameters()),
    })
results = pd.DataFrame(rows).sort_values("测试ROC-AUC", ascending=False)
display(results.style.format({"测试ROC-AUC": "{:.3f}", "测试平均精确率": "{:.3f}"}))
  模型 测试ROC-AUC 测试平均精确率 参数量
1 节点袋(求和读出、无边) 0.826 0.151 7921
2 图同构网络(求和读出) 0.754 0.087 36004
3 图同构网络(均值读出) 0.753 0.154 36004
0 训练阳性率常数 0.500 0.025 0
4 图同构网络(最大值读出) 0.396 0.021 36004

第 10 步:读出对图大小的敏感性#

  不训练模型也能看见读出先验。对同一节点表示整体复制一次:

  • sum 向量扩大为 2 倍,明确保留节点数;

  • mean 不变;

  • max 不变。

  因此 sum 可能利用“分子大小”这个信号,也可能被大图尺度主导。下面用数值断言把概念变成可测试性质。

h = torch.tensor([[1., 2.], [3., 1.]], device=device)
g = torch.zeros(2, dtype=torch.long, device=device)
h_twice = torch.cat([h, h])
g_twice = torch.zeros(4, dtype=torch.long, device=device)
for mode in ["sum", "mean", "max"]:
    once = graph_readout(h, g, 1, mode)
    twice = graph_readout(h_twice, g_twice, 1, mode)
    mode_names = {"sum": "求和读出", "mean": "均值读出", "max": "最大值读出"}
    print(mode_names[mode], "原图", once.cpu().numpy(), "复制节点后", twice.cpu().numpy())
assert torch.allclose(graph_readout(h_twice, g_twice, 1, "sum"),
                      2 * graph_readout(h, g, 1, "sum"))
assert torch.allclose(graph_readout(h_twice, g_twice, 1, "mean"),
                      graph_readout(h, g, 1, "mean"))
求和读出 原图 [[4. 3.]] 复制节点后 [[8. 6.]]
均值读出 原图 [[2.  1.5]] 复制节点后 [[2.  1.5]]
最大值读出 原图 [[3. 2.]] 复制节点后 [[3. 2.]]

结论、局限与常见错误#

  本案例用 GIN 与节点袋基线分离“原子组成”和“连接结构”,并通过 scaffold 划分减少骨架近邻跨集合造成的泄漏。局限包括:

  1. 快速模式是固定抽样,估计方差较大,不能作排行榜比较;

  2. 简化 GIN 忽略键类型和边属性,而化学键显然重要;

  3. 没有使用虚拟节点、专业分子特征、超参数搜索或多种子集成;

  4. ROC-AUC 只评价排序,不评价校准和实际决策代价;

  5. 数据标签来自特定筛选实验,不能直接解释成临床结论。

  常见错误:随机拆分相似骨架、把节点分类与图分类混淆、在节点维直接输出图标签、用测试集 AUC 选择读出方法、把类别编码当连续数,以及漏加图偏移导致不同分子的节点被错误连接。

综合练习#

  1. 在 GIN 中加入边特征编码,令消息同时依赖原子和键类型。

  2. 比较 2、3、5 层 GIN,报告验证 AUC、训练时间和节点表示相似度。

  3. 实现 jumping knowledge:拼接各层图读出,再与只用最后一层比较。

  4. 用训练集选择概率阈值,并在验证集报告精确率和召回率;说明为什么仍不能用测试集选择阈值。

  5. 关闭快速模式,接入 OGB 官方评价程序,至少使用 3 个随机种子运行实验,并报告结果的均值、标准差、硬件环境和运行时间。