PyTorch Geometric 加载 QM9 数据集:从环境自检到跑通一次分子属性预测
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
本文带你用 PyTorch Geometric 完成 QM9 数据集加载的完整链路:先做环境与依赖自检,再一次性把数据落盘,然后核对数据、接入 SchNet 跑通一个训练 epoch,最后给出批处理提速与排错清单。
QM9 是 GNN 分子建模的标准测试集:约 13 万个含 C、H、N、O、F 的有机小分子,每个分子带 3D 原子坐标与 19 种量子化学性质(偶极矩、HOMO/LUMO 能级、内能、焓等),平均 18 个原子、37 条键。它也是 PyG 中InMemoryDataset(内存数据集基类:首次处理后把全部图存成一个.pt缓存,之后再打开直接载入内存)的典型使用者,官方实现在 torch_geometric/datasets/qm9.py。
QM9 环境自检:RDKit 依赖与降级加载
开始之前先确认两件事,避免「PyG QM9 加载失败」类问题在训练中途才暴露。
- 装好 PyG 后,先在终端验证导入:执行
python -c "import torch_geometric",再单独执行python -c "import rdkit"。 - RDKit 是化学信息学工具包,这里负责解析 QM9 的原始 SDF 文件(一种存放分子结构与坐标的文本格式),把分子转成原子、键、坐标等图字段。
QM9 的处理流程对 RDKit 有依赖,缺失时会静默降级:
download()尝试import rdkit失败时,改从processed_url下载官方预处理包qm9_v3.pt,跳过 SDF 解析;process()检测到未安装时,在 stderr 打印一条提示后直接载入预处理数据;- 功能影响:拿不到 SMILES 字符串与
name等分子结构字段,其余 19 列目标值不受影响,正常训练可以继续。
安装建议:优先conda install -c conda-forge rdkit,其次pip install rdkit-pypi(偶有依赖冲突,按需处理)。
💡 注意一个坑:如果你曾在无 RDKit 状态下加载过 QM9,processed/data_v3.pt缓存里存的就是降级版本。装上 RDKit 后不会自动重新处理,需要删掉该数据集目录,或构造时传force_reload=True触发重处理。
QM9 一次加载:路径配置与 FileNotFoundError 排查
QM9 构造时只需给一个root目录,下载与处理全自动完成:
from torch_geometric.datasets import QM9 dataset = QM9(root='./data/QM9') # 首次运行自动下载并处理 print(len(dataset), dataset.data.y.shape)目录约定:raw/存放原始文件(装了 RDKit 是gdb9.sdf等,没装是qm9_v3.pt),processed/生成data_v3.pt缓存。构造完成后数据已整体载入内存,再次实例化时直接读缓存,重开成本从分钟级降到秒级;需要重新处理时传force_reload=True即可。
路径配置建议与排错提示:
- 不要照抄示例里基于
__file__的相对拼接路径:Jupyter 中__file__不可用,osp.dirname(osp.realpath(__file__))会拿到非预期位置。改用显式的绝对路径,或相对当前工作目录的./data/QM9。 - 报
FileNotFoundError: .../raw/gdb9.sdf:多为 root 写权限不足、目录被提前删了一半,或断网导致下载中断。检查目录可写、补全下载,或手动把文件放进raw/。 - 报 RDKit 相关
ImportError:不是崩溃,而是走了上述降级路径,装上 RDKit 再按提示重建缓存即可。
数据核对:19 列目标与 y 属性索引错位排查
加载后先核对规模与形状,再决定用哪一列训练:
import torch print(len(dataset)) # 130831 print(dataset.data.y.shape) # torch.Size([130831, 19]) # DimeNet 预训练权重对应原子化能量列,需要先重排 y: idx = torch.tensor([0, 1, 2, 3, 4, 5, 6, 12, 13, 14, 15, 11]) dataset.data.y = dataset.data.y[:, idx]19 列依次为:偶极矩 μ、极化率 α、HOMO 能量、LUMO 能量、能隙 Δε、电子空间展布、ZPVE、内能 U₀/U、焓 H、自由能 G、热容 cv,以及四组原子化能量(U₀ᴬᵗᵒᵐ、Uᴬᵗᵒᵐ、Hᴬᵗᵒᵐ、Gᴬᵗᵒᵐ)和三个转动常数。处理时 PyG 已把 Hartree、kcal/mol 等原始单位统一换算成 eV,你拿到的就是可训练的量纲。
⚠️ 索引错位问题:DimeNet 官方预训练权重是按原子化能量列训练的,所以 examples/qm9_pretrained_dimenet.py 先把 y 的第 7–10 列(U₀、U、H、G)替换成第 12–15 列(对应原子化能量),再把 cv 挪到第 11 位。若不重排就传target索引,预测值会跟错列,评估 MAE 明显异常甚至越界报错。
另一条捷径是工厂方法:SchNet.from_qm9_pretrained(path, dataset, target)会自行下载权重、切分 train/val/test,参考 examples/qm9_pretrained_schnet.py。注意调用它时传原始未重排的 dataset,索引逻辑由方法内部处理,先重排再传target反而会错位。
接入训练:QM9 分子属性预测最小主线
下面以偶极矩(第 0 列)为目标的 SchNet 训练为主线,完整版见 examples/qm9_nn_conv.py:
import torch.nn.functional as F import torch_geometric.transforms as T from torch_geometric.datasets import QM9 from torch_geometric.loader import DataLoader from torch_geometric.nn import SchNet dataset = QM9('./data/QM9', transform=T.Distance(norm=False)).shuffle() data = dataset[0] data.y = data.y[:, 0] # 只保留目标列,其余同理 train, val, test = dataset[:100000], dataset[100000:110000], dataset[110000:] loader = DataLoader(train, batch_size=64, shuffle=True) model = SchNet(hidden_channels=128, num_filters=128, num_interactions=6, num_gaussians=50).to('cuda') opt = torch.optim.Adam(model.parameters(), lr=1e-3) for epoch in range(1, 3): for data in loader: data = data.to('cuda') opt.zero_grad() out = model(data.z, data.pos, data.batch) loss = F.mse_loss(out.view(-1), data.y) loss.backward(); opt.step()要点说明:
T.Distance在线计算原子对距离特征,SchNet 的 forward 只吃z(原子序数)、pos(坐标)与batch(批内分子归属);dataset.shuffle()后按位置切片做划分,是官方示例的标准做法;- 追求精度可对目标做均值/标准差归一化,评估时再乘回去;QM9 还提供
dataset.atomref(target)原子参考能量,用于把预测校正到原子化能量基线,进阶再试。
DataLoader 批处理提速与 pre_transform 预处理缓存
- 动态批处理:分子大小有差异时,固定
batch_size的批次内节点数会波动。改用DynamicBatchSampler(dataset, max_num_nodes=1000)并传给DataLoader(dataset, batch_sampler=sampler),按节点数上限组批,单批显存更平稳,训练吞吐通常更高。 - pre_transform 落盘缓存:
transform每次访问数据都会重新执行;pre_transform只在首次处理时执行一次,结果直接写进processed/data_v3.pt。凡是不依赖随机性、每次都要做的特征加工(如只保留某一目标列),都应放进pre_transform,后续加载零成本。 - 多进程加载可设
num_workers=4;13 万分子全量加载约占数 GB 内存,机器紧张时先切片到子集验证流程。
排错清单与延伸练习
把前文分散的排错提示汇总成自检清单:
- 报
FileNotFoundError且路径含raw/gdb9.sdf:确认 RDKit 状态与目录权限,必要时手动补文件; - stderr 出现 "Using a pre-processed version of the dataset":属正常降级,但装好 RDKit 后需重建缓存才能拿到 SMILES;
- 评估时 target 索引越界或 MAE 异常:检查是否做了 y 列重排、
target是否落在 [0, 11]; - 重开 notebook 数据「不见了」:确认
processed/data_v3.pt是否存在,存在则加载应只需秒级; - 内存不足:用切片子集先跑通,再上全量。
延伸练习:
- 把 SchNet 换成 DimeNetPlusPlus,走一遍
from_qm9_pretrained+ y 列重排的完整流程; - 用
dataset.atomref(12)对原子化能量目标做参考能量校正,比较校正前后 MAE; - 把数据接入
LightningDataModule(见 torch_geometric/data/lightning/),或直接换用 PCQM4M 数据集体会预训练迁移效果。
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考