news 2026/9/6 16:29:08

PyG 加载 QM9 数据集:从下载到跑通训练

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyG 加载 QM9 数据集:从下载到跑通训练

PyG 加载 QM9 数据集:从下载到跑通训练

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

第一次用 PyTorch Geometric(PyG,一个基于 PyTorch 的图神经网络库)做分子属性预测时,QM9 几乎是绕不开的入门数据集——它装着 13 万多个有机小分子和 19 个量子化学回归目标。可实际操作里,下载依赖、目标字段切列、量纲归一化这三步经常卡住人。这篇文章的思路是:先用最短的代码把数据跑起来,再回过头解释每一步在做什么,遇到报错时知道往哪儿查。

十分钟把 QM9 加载进内存

先跑起来再说。下面这段代码在任意能装torch-geometric的目录里直接执行,首次运行会自动把原始文件下载到data/QM9/raw/,再处理成data/QM9/processed/data_v3.pt

from torch_geometric.datasets import QM9 dataset = QM9(root='data/QM9') print(len(dataset)) # 130831 个分子 print(dataset.num_features) # 11 维原子特征 print(dataset.num_targets) # 19 个回归目标 print(dataset[0].z.shape) # 第一个分子的原子序数

第二次运行同一段代码时,PyG 检测到processed/目录里已经有data_v3.pt,就跳过下载与处理直接读盘,通常几秒内返回。这个"缓存到磁盘"的行为来自QM9的父类InMemoryDataset,具体逻辑在 torch_geometric/datasets/qm9.py 里可以看到。

如果你只是想验证环境能不能跑通,到这里就够了。下面我们把"数据里到底装了什么"拆开看。

QM9 里每个分子长什么样

dataset[i]返回的是一个Data对象,把它想象成一行"分子体检报告":

  • z:每个原子的原子序数,整数张量;
  • pos:每个原子的三维坐标(笛卡尔坐标),QM9 只保留了能量最低的构象;
  • x:11 维原子特征,由 5 类原子 one-hot(H/C/N/O/F)加上原子序数、芳香性、sp/sp²/sp³ 杂化、连氢数拼起来;
  • y:19 个回归目标,从偶极矩、极化率一直到各种能量;
  • edge_index/edge_attr:化学键连接关系和键型 one-hot。

19 个目标对应的物理量和单位,源码里用一张表格列得很清楚,直接打开 torch_geometric/datasets/qm9.py 翻到QM9的 docstring 就能看到完整清单。做能量类目标(U_0UHG以及对应的 atomization 版本)时,通常要先减去原子参考能量再做归一化,QM9.atomref(target)就是干这个的——它按原子序数查一张内置表,把每个原子的孤立原子能量扣掉,让网络只需学"成键带来的能量变化",收敛会快很多。

下载与处理:RDKit 决定走哪条路

这里有个容易踩的坑:PyG 的QM9内部会尝试import rdkit,成功与否决定了它下载哪份原始文件、走哪条处理分支。

  • 有 RDKit:下载qm9.zip(SDF 结构 + 目标 CSV + 一个 skip 列表),然后自己解析 SDF 生成z / pos / x / edge_index / smiles。好处是能拿到 SMILES 字符串,后续做 SMILES 相关模型时方便。
  • 无 RDKit:改下qm9_v3.zip,里面是官方预处理好的qm9_v3.pt。好处是不用装化学库,坏处是数据里没有smiles字段,其他字段与手动处理版一致(源码process()里也做了对应分支)。

所以,如果你后面要写涉及 SMILES 的代码,建议先把 RDKit 装上:

conda install -c conda-forge rdkit # 或者:pip install rdkit-pypi

装完再触发一次处理(QM9(root=..., force_reload=True)会强制重跑download()process())。如果你只想快速看效果、暂时不想装 RDKit,PyG 会自动切到预处理好版本,功能基本不受影响,只是data.smilesNone

第一次跑的时候,如果下载中断或者目录被手动删了一半,可能会看到类似这样的报错:

FileNotFoundError: [Errno 2] No such file or directory: 'data/QM9/processed/data_v3.pt'

通常是raw/下的原始文件缺了一份(比如uncharacterized.txt没下载完),PyG 判断 raw 文件不齐就跳过了process()。处理方法是把data/QM9/整个目录清掉,再重新执行QM9(root='data/QM9'),让它从头走一遍下载与处理。

训练时最容易卡住的三处

数据集加载本身只是第一步,真正训练时几个字段格式问题反复出现。

第一处:y有 19 列,但你的模型只预测 1 个。直接拿data.y喂给SchNet之类的回归头,形状对不上。最简单的做法是在数据集构造时就用transform把目标列切出来,参考 examples/qm9_nn_conv.py 里的MyTransform

import copy class KeepOneTarget: def __init__(self, target: int): self.t = target def __call__(self, data): data = copy.copy(data) data.y = data.y[:, self.t:self.t+1] return data

把它塞进QM9(root=..., transform=KeepOneTarget(0)),之后data.y就只有一列了。

第二处:不同目标量级差很多。偶极矩、极化率、HOMO 能量、自由能各自量纲和数值范围都不一样,直接 MSE 训练时大数值目标会主导梯度。惯例是先做标准化(减均值除标准差),训练完评估 MAE 时再乘回去。官方示例也是这么做的,dataset.data.y = (dataset.data.y - mean) / std之后,std要单独存下来给评估用。

第三处:DimeNet 系列模型的预训练示例里有一行列重排。如果你在 examples/qm9_pretrained_dimenet.py 里看到dataset.data.y = dataset.data.y[:, idx]idx = [0,1,2,3,4,5,6,12,13,14,15,11],别困惑——那是因为 DimeNet 论文里把U_0/U/H/G(列 7/8/9/10)换成了对应的 atomization 版本(列 12/13/14/15),δe(列 4)又用e_LUMO - e_HOMO代替。跑 SchNet 或 NNConv 时不需要这一步,只有加载 DimeNet 预训练权重时才要对齐列顺序。

一个端到端的分子属性预测工作流

把上面几步串起来,就是一个最小可跑的 QM9 训练脚本。为了短,这里用SchNet(一个基于连续滤波器的分子 GNN,输入是原子序数 + 坐标,输出标量属性),完整代码可以对照 examples/qm9_nn_conv.py 扩展。

准备数据:切目标列并做归一化,然后按官方习惯做随机划分(最后 1 万做测试,前 1 万做验证,其余训练):

import torch from torch_geometric.datasets import QM9 from torch_geometric.loader import DataLoader dataset = QM9(root='data/QM9', transform=KeepOneTarget(0)).shuffle() y = dataset.data.y mean, std = y.mean(dim=0, keepdim=True), y.std(dim=0, keepdim=True) dataset.data.y = (y - mean) / std train_ds = dataset[10000:] val_ds = dataset[:10000] train_loader = DataLoader(train_ds, batch_size=64, shuffle=True) val_loader = DataLoader(val_ds, batch_size=64)

模型和优化器:SchNet只吃z / pos / batch三个字段,坐标是分子图的关键输入,别忘了data.pos必须在数据集里(QM9 默认就带):

import torch.nn.functional as F from torch_geometric.nn import SchNet model = SchNet(hidden_channels=128, num_filters=128, num_interactions=6, num_gaussians=50) opt = torch.optim.Adam(model.parameters(), lr=1e-3)

训练循环,评估时用std把归一化后的 MAE 还原回原始量纲:

for epoch in range(1, 21): model.train(); total = 0.0 for batch in train_loader: opt.zero_grad() pred = model(batch.z, batch.pos, batch.batch) loss = F.mse_loss(pred.view(-1), batch.y.view(-1)) loss.backward() opt.step() total += loss.item() * batch.num_graphs print(f'epoch {epoch:02d} train loss {total/len(train_ds):.4f}')

跑几个 epoch 后 loss 应该明显下降。如果训练速度不够快,可以把num_workers加到 2~4;如果显存紧张,先把hidden_channels/num_filters从 128 砍到 64。

源码与示例入口

  • 数据集实现:torch_geometric/datasets/qm9.py(download()/process()/atomref()都在这个文件里)
  • 分子 GNN 最小示例(NNConv + GRU + Set2Set):examples/qm9_nn_conv.py
  • 用 SchNet 预训练权重跑 12 个目标的 MAE 评测:examples/qm9_pretrained_schnet.py
  • 用 DimeNet / DimeNet++ 预训练权重评测(含列重排示例):examples/qm9_pretrained_dimenet.py
  • 更大规模的分子图数据(PCQM4M,分布式训练参考):examples/multi_gpu/pcqm4m_ogb.py
  • 同类分子数据集 ZINC(15 万分子,1 个二分类目标):torch_geometric/datasets/zinc.py

QM9 的关键就三件事:先跑起来拿到Data对象,再切出你要的那个目标列并归一化,最后按分子图模型的输入习惯(z / pos / batchx / edge_index / batch)把字段喂对。把这三步走顺,换到 ZINC 或 PCQM4M 上只是换个数据集类名的事。

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

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/6 16:27:24

自动测试系统与网络化仪器:从总线演进到系统落地全解析

简介:这份《自动测试系统与网络化仪器》PPT课件聚焦电子测量与自动测试领域,适合测控、仪器科学与技术专业学生及从事测试系统开发的工程师。内容系统梳理了ATS的组成架构,从控制器、程控仪器到总线接口、测试软件与被测对象,并着…

作者头像 李华
网站建设 2026/9/6 16:20:42

Starship 安装指南:5 分钟在全平台装好你的跨 Shell 提示符

Starship 安装指南:5 分钟在全平台装好你的跨 Shell 提示符 【免费下载链接】starship ☄🌌️ The minimal, blazing-fast, and infinitely customizable prompt for any shell! 项目地址: https://gitcode.com/GitHub_Trending/st/starship Star…

作者头像 李华
网站建设 2026/9/6 16:20:23

华为业务变革框架与战略级项目管理实战指南

简介:这是一份华为业务变革管理框架(BTMS)V2.0与战略级项目管理的完整讲义,共111页PPT,适合企业变革管理者、PMO成员、战略规划人员及组织管理研究者学习参考。内容以BTMS框架为主线,系统梳理年度规划流程、…

作者头像 李华
网站建设 2026/9/6 16:18:00

ComfyUI 性能优化与显存配置:从 45 秒到 15 秒的实战指南

ComfyUI 性能优化与显存配置:从 45 秒到 15 秒的实战指南 【免费下载链接】ComfyUI The most powerful and modular diffusion model GUI, api and backend with a graph/nodes interface. 项目地址: https://gitcode.com/GitHub_Trending/co/ComfyUI ComfyU…

作者头像 李华