PyTorch Geometric 加载 QM9 数据集:从环境自检到跑通一次分子属性预测
2026/9/5 21:54:03 网站建设 项目流程

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 有依赖,缺失时会静默降级:

  1. download()尝试import rdkit失败时,改从processed_url下载官方预处理包qm9_v3.pt,跳过 SDF 解析;
  2. process()检测到未安装时,在 stderr 打印一条提示后直接载入预处理数据;
  3. 功能影响:拿不到 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是否存在,存在则加载应只需秒级;
  • 内存不足:用切片子集先跑通,再上全量。

延伸练习:

  1. 把 SchNet 换成 DimeNetPlusPlus,走一遍from_qm9_pretrained+ y 列重排的完整流程;
  2. dataset.atomref(12)对原子化能量目标做参考能量校正,比较校正前后 MAE;
  3. 把数据接入LightningDataModule(见 torch_geometric/data/lightning/),或直接换用 PCQM4M 数据集体会预训练迁移效果。

【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询