最近在整理做分子几何表征项目的思路,想把这个过程中的关键取舍和经验捋一遍。这个主题叫“不止生成:迈向更好的分子几何表征”,说白了就是:我们不只是想让模型生成一个看起来合理的分子三维构象,更重要的是让这些几何信息真正变成“表征”,喂给下游性质预测、分子生成、虚拟筛选这些任务时,能稳定地带来增益。如果你也在做AI制药、材料数据挖掘、构象分析这类事情,那这篇东西应该能对得上你的胃口。
很多人上手就奔向“生成”这个目标:给定一个二维骨架,去预测三维坐标。但项目推进到中后期你会发现,生成质量只是地基,真正的分水岭在于“几何表征”好不好。同样的坐标数据,放进不同的编码结构,用不同的预训练目标去学,最后的任务性能可能差一大截。所以这篇文章我想从需求定义讲到方案选型,再讲实操细节、评估方法和避坑经验,争取让你读完能直接复制一套完整的流程去跑你自己的分子体系。
1. 先搞清楚:分子几何表征到底要解决什么问题
1.1 从2D到3D:为什么几何信息这么关键
分子在二维层面只是一张拓扑图,原子连成快照,但现实中分子是三维的。同一个分子骨架可以扭转出无数种构象,有些构象能量高,有些能量低,而生物活性和物理性质往往取决于低能量构象群,甚至取决于分子在靶点口袋里的“变形能力”。这就是几何信息不能被忽略的直接原因,一个简单的例子是手性分子,互为镜像但空间结构完全不同,化学性质却可以天差地别。如果没有三维坐标,模型根本没法区分这种差异。
我之前在处理一个药物小分子的性质预测任务时做过对比:只用拓扑图特征,模型的精度卡在一个平台期上不去;后来引入3D构象和原子间距离张量,误差直接降了一截。这就是几何表征的价值所在——它补全了拓扑图缺失的空间约束信息,让模型有机会理解分子的实际形状、电子云排布、位阻效应。
但这里必须警惕一个陷阱:有坐标不等于有好表征。如果你只是把坐标堆成一个向量塞进MLP,模型学到的其实是绝对坐标的某种记忆,转动一下分子,输出就变了。这不符合物理规律,一个分子的能量和性质不该因你在电脑里把分子旋转90度而改变。所以几何表征的核心问题可以归纳成两句话:怎么编码几何信息,才能让模型尊重物理对称性;怎么学习几何信息,才能让下游任务真正受益,而不只是头疼医头。
1.2 几何表征的技术难点与核心约束
分子几何表征有三个绕不开的硬约束,我建议在做方案选型之前先背下来。
第一个是等变性(equivariance)。想象你手里拿着一张椅子,你把它原地转动一下,椅子的外形、你该坐的位置都没变,但每条腿的坐标全变了。分子的情况更苛刻:合法操作不只是旋转,还有平移。一个模型的预测目标如果是坐标、位移、速度向量这类量,那它必须随输入一起旋转,这叫SE(3)等变性。如果预测目标是标量性质,比如能量、溶解度,那模型必须做到旋转平移后输出不变,这叫不变性。一个不具备等变性的几何表征模型,在物理任务上总会有系统性偏差。
第二个是连续空间的高维性。分子几何存在于一个连续的高维空间,每个原子有三个空间自由度,一个50个原子的分子就是150维空间。想在这个空间里做采样、做优化、做生成,朴素方法会遇到维度灾难。这也是为什么纯靠坐标回归的模型往往抖得厉害,因为高维连续空间里没有足够的数据密度来支撑稳定学习。
第三个是构象分布的多样性。分子在室温下不是固定在一个构象里的,它有一个Boltzmann分布,不同构象之间只有微小的能量差异。生成任务里要求模型“找出一个合理构象”并不难,难的是让模型理解整个构象集合的分布特征,并把这种分布编码进表征里。好多项目在“生成”阶段玩得飞起,一到下游任务就拉胯,问题往往就出在这里:模型没见过足够多样的几何变化,学到的表征是“断断续续”的,换一个构象就认不出来了。
理解这三点之后,你就能看懂为什么主流方案会往等变图神经网络和去噪预训练方向走,这不是炫技,是物理约束决定的。
2. 方案拆解:从构象生成到几何表征学习的完整链路
2.1 第一环:构象样本的质量决定表征天花板
很多人觉得既然要学几何表征,那输入坐标随便用RDKit生成一个就行。这个想法,坦白说,浪费了整个管线的一半潜力。构象样本的质量决定了模型能看到多少几何变化,也决定了表征的天花板在哪儿。
我一般把构象获取分成三个层次。第一个层次是快速力场法,比如RDKit里面的ETKDG、MMFF优化,速度和成本都很低,一个分子几毫秒就能出一个构象。它的问题是只能覆盖低能量区域的一些局部极小值,很多高能量或长程折叠的构象根本采不到。第二个层次是半经验或密度泛函级别的结构优化,精度上去了,但计算成本也上去了,适合中小规模数据集,不适合动不动几十万分子的场景。第三个层次是生成式方法,比如基于扩散模型的GeoDiff、Torsional Diffusion这类方案,它们直接从数据里学习扭转角分布,出来的构象多样性和合理性都明显更好,但训练和推理成本都不低。
我的建议是:如果你做的是大规模预训练,第一层就可以,但要做多样化的筛选;如果你做的是小规模高精度表征,尽量用带力场精修的第二层;如果你有新分子骨架的专项任务,第三层的生成模型能带来额外优势。我在实际项目里用的是“RDKit粗采样 + 能量排序 + 去重”的组合管线,一个分子出20个候选构象,按MMFF能量排序后保留低能量且RMSD差异明显的5个。这样可以兼顾质量和多样性,模型在后面学到的几何变化集就比较完整。
2.2 第二环:选择等变表征模型而不是普通图网络
处理好输入构象之后,下一步就是选编码器。这个环节最常见的失误是:把2D图神经网络原封不动搬到3D数据上,只是额外加了距离特征。这么做不是完全不行,但问题很多。普通GNN的节点更新是建立在“消息是标量”这个假设上的,虽然距离是旋转不变的,但你丢弃了方向信息,模型其实是在用一个降维的、残缺的几何描述去做预测。
更好的选择是等变图神经网络。当年处理这个项目时,我在EGNN、PaiNN、GemNet这几个架构之间做了一番折腾,最终主力用的是EGNN的等变消息传递框架。原因有三点:第一,它直接把坐标作为模型的一部分参与更新,而不是当作额外的边特征,几何信息没有被降维处理;第二,它在做坐标预测和力场学习的时候天然满足等变性,不用额外加数据增强去“骗”模型学会旋转不变;第三,实现复杂度适中,比GemNet那些动辄几层球谐函数张量积的架构好维护得多。
等变网络的核心思想其实不复杂:在每一层更新里,既要更新原子特征(标量),又要更新坐标(向量),而且坐标更新的方式必须和旋转保持一致。你可以把这个过程理解成一个木偶戏:人体(标量特征)在动,线条(向量坐标)也跟着动,但不管你怎么转舞台,木偶和线条之间的相对关系是不变的。这么设计出来的模型,它在物理上的归纳偏置是内生的,不需要靠大量数据来领悟“旋转了分子,标签还应该一样”这种底层规律。
2.3 第三环:预训练任务与下游适配
架构选好之后,最关键的灵魂问题是:模型用什么目标去学几何表征?这也是标题里“不止生成”的核心含义。如果你的项目只做构象生成,那训练目标显然是最小化生成坐标和真实坐标的误差;但如果要做更好的表征,那生成本身只是辅助任务,你更应该思考的是“如何让模型学到既能表达几何细节、又能泛化到不同任务目标的特征”。
我在项目里用了三个训练阶段的组合。第一是无监督的几何去噪预训练,做法很简单:把构象坐标做微小随机扰动,模型要预测出原始无噪声的坐标。这个任务逼着模型编码器理解分子内部坐标的可信度,哪些原子位置是局部能量约束很强的,哪些是松散可动的。第二是自监督的掩码原子预测,随机遮盖一部分原子及其局部几何,让模型从剩余原子环境去恢复被遮盖环境的几何类型。第三才是生成式预训练,用扩散损失去学习构象空间的结构。实践证明,前两个任务对下游性质预测的提升更直接,生成任务更多是帮助模型在困难构象上建立泛化能力。
下游适配时又要注意一件事:纯拿预训练的embedding做聚类或线性探针,并不总是最优。我在验证一个活性预测任务时发现,冻结encder只训一个MLP头的效果,反而不如在encoder后面加一个轻量的任务适配层,用低学习率微调几轮来得好。原因是分子几何表征与化学性质的映射往往是非线性的,线性探针会丢掉关系信息;而全量微调又可能破坏预训练学到的几何规律,所以折中路径是只微调最后两层加投影头。
3. 实操手记:构建一个分子几何表征项目的完整流程
3.1 数据准备与几何特征化
第一步是从数据源准备干净的分子。如果你用的是SMILES字符串,建议先做标准化处理:去盐、去重、选择主组分、质子化状态统一。这一步很多人偷懒,结果后面全乱套。举个例子,同一个分子的中性形态和酸式盐形态如果在数据里并存,模型会被迫去学那些和物理性质无关的假变化。
接下来就是加氢和构象生成。我用的是RDKit,代码上大概是这样:
from rdkit import Chem from rdkit.Chem import AllChem def molecule_to_conformers(smiles, num_conf=20, seed=42): mol = Chem.MolFromSmiles(smiles) if mol is None: return None mol = Chem.AddHs(mol) params = AllChem.ETKDGv3() params.randomSeed = seed params.useRandomCoords = True params.numThreads = 0 cids = AllChem.EmbedMultipleConfs(mol, numConfs=num_conf, params=params) results = [] for cid in cids: try: converged = AllChem.MMFFOptimizeMoleculeConfs(mol, maxIters=200) # returns (cid, energy) except ValueError: continue conf = mol.GetConformer(cid) positions = conf.GetPositions() results.append((positions, converged[cid][1])) return mol, results这段代码里有两个容易被忽略的细节。第一个是AddHs一定要做,不加氢的分子连键长信息都不完整,很多机器学习的3D模型会把不饱和价电子状态搞混。第二个是MMFF优化那一行,如果不做力场精修,ETKDG直接吐出来的坐标在键长键角上会有明显的不自然感,模型学了这种畸变几何,后面去噪任务会不收敛。
构象拿到后,构建3D分子图。我建议以原子为节点、原子间空间距离为连边依据,采用“半径截断+前k近邻”混合策略:默认截断半径取4.5埃,再确保每个原子至少连到前8个近邻。边特征领域不要只用纯距离标量,可以尝试加入单位方向向量的特征投影,这在等变网络里是加分项。节点特征上,除了原子序数、形式电荷、杂化类型,我还会补充局部环境的平面性指标和手性标记,这能帮助模型区分分子构型的细微差别。
3.2 模型构建的关键实现细节
等变图神经网络的实现里,最容易写错的就是坐标更新分支。以EGNN为例,一个等变层的更新步骤大致是这样:
import torch import torch.nn as nn class EGCLayer(nn.Module): def __init__(self, hidden_dim): super().__init__() self.phi_m = nn.Sequential( nn.Linear(hidden_dim * 2 + 1, hidden_dim), nn.SiLU(), nn.Linear(hidden_dim, hidden_dim) ) self.phi_h = nn.Sequential( nn.Linear(hidden_dim * 2, hidden_dim), nn.SiLU(), nn.Linear(hidden_dim, hidden_dim) ) self.phi_x = nn.Sequential( nn.Linear(hidden_dim, hidden_dim), nn.SiLU(), nn.Linear(hidden_dim, 1, bias=False) ) def forward(self, h, x, edge_index): i, j = edge_index d_ij = (x[i] - x[j]).norm(dim=-1, keepdim=True) # (E, 1) d_ij = torch.sqrt(torch.clamp(d_ij**2, min=1e-6)) edge_input = torch.cat([h[i], h[j], d_ij], dim=-1) m_ij = self.phi_m(edge_input) agg = torch.zeros_like(h) agg.index_add_(0, i, m_ij) h_new = self.phi_h(torch.cat([h, agg], dim=-1)) coef = self.phi_x(m_ij) / d_ij # 关键:除以距离保证等变 delta_x = torch.zeros_like(x) delta_x.index_add_(0, i, (x[i] - x[j]) * coef) x_new = x + delta_x return h_new, x_new这里最需要注意的细节是coef的计算。坐标更新里出现了(x_i - x_j) / d_ij,这个表达式其实是一个单位方向的向量,它本身是等变的。从旋转角度看,只要坐标向量旋转,这个差值向量也会跟着旋转;从平移角度看,坐标系原点移动时,差值不变。所以这个设计天然保证了等变性。如果你只把坐标差值直接乘一个标量,不除以距离,那么更新幅度会严重依赖分子尺度的大小,训练起来会特别不稳定。
另外,索引归约这里用index_add_而不是手动循环,不然E条边的数量一上来,内存和时间都扛不住。我在项目里用的隐藏维度和层数是版本化的,底线上保证每层都有残差连接和层归一化,防止等变更新的梯度在深层消失。
3.3 训练策略、损失函数与超参数选择
几何表征模型的训练目标可以多任务组合,但损失之间的配比和权重需要谨慎。我用的是以下加权损失:
L = w_noise * L_noise + w_contrast * L_contrast + w_gen * L_genL_noise是坐标去噪损失,用预测噪声和真实噪声之间的均方误差;L_contrast是对比损失,把同一个分子不同构象的embedding拉近,同时把不同分子的embedding推开;L_gen则是扩散生成损失,让模型能重构训练集中的构象分布。权重上,我实际取w_noise=1.0, w_contrast=0.5, w_gen=0.1。生成损失权重不敢给大,因为一旦它主导,模型会退化成纯生成模型,特征空间里全是构象重建细节,反而对下游任务不利。
超参数方面,我整理了一份通用参考配置,你可以拿去做起点:
| 参数 | 推荐值 | 备注 |
|---|---|---|
| 原子隐藏维度 | 128 | 分子量不太大时够用,太大会导致显存爆炸 |
| 消息传递层数 | 5 | 层数太浅学不到长程几何约束 |
| 学习率 | 1e-3 | 用AdamW + cosine schedule |
| 批量大小 | 64 | 和构象采样数量有关,每个样本多构象时可以减半 |
| 梯度裁剪 | 10.0 | 防止坐标更新偶发剧烈震荡 |
| 坐标扰动噪声sigma | 0.05 Å | 去噪预训练专用,太大变成模糊 |
| 随机旋转增强 | on | 每个epoch对坐标做随机SE(3)旋转 |
训练过程中有一个经验可以分享:等变网络在用混合精度训练时,坐标更新那个分支特别容易发鬼火。fp16下的坐标差值计算,动态范围比标量特征大得多,动不动就溢出。我的方案是坐标分支保持在fp32计算,只对特征矩阵用混合精度;实现对PyTorch来说就是把等变层的坐标更新部分放到torch.autocast(dtype=torch.float32)作用域里,其他部分保持默认。这样既不损失太多训练速度,又稳定得多。
4. 质量评估:如何判断几何表征“好”还是“不好”
4.1 下游任务基准与定量结论
判断几何表征质量,我坚持只看下游任务上的实际提升,不看预训练loss降到多低。跑完预训练之后,固定encoder参数,用线性探针和微调两种方式分别在公开数据集上做定量评测。
常用的标准做法是在QM9数据上做分子性质预测,比如极化率、HOMO-LUMO能隙、零点能这些目标。正常配置下,一个5层等变编码器加上去噪预训练,在HOMO-LUMO这种复杂性质上的MAE能降到0.05 eV量级甚至更低,而不做几何预训练的基线模型通常要到0.08 eV以上。再一个是MD17这种分子动力学数据集,能量预测MAE可以控制在0.5 kcal/mol以内,力预测的MAE也有可观下降。这不是说我拿到的绝对数值有多厉害,重点在于同一个实验条件下,几何表征是否稳定地比拓扑基线更准。
我还建议加一个等变性检验:随机旋转一个测试集分子,然后比较模型输出和原始输出。对坐标预测任务,输出应该完美跟着旋转走;对标量性质任务,输出应该完全不变。如果模型对这个测试敏感,说明你的等变模块里有泄漏或规范化不均匀,得回头检查图构造阶段是否有非等变的特征混进去了。
4.2 消融实验:不是堆模型就有用
很多人工程会上来就把生成、对比、去噪三个损失全堆上,模型参数量也一路增加,最后表现却不一定最好。你需要做消融实验来精确判定每个模块的真实贡献。
我做的消融排序大概是这样的:同时去掉对比和去噪,只留下生成损失做训练,下游任务的预测误差上升最明显;只保留去噪,对比去掉,效果次之;只有对比没有去噪,效果垫底但也比完全没有几何预训练强。这说明在几何表征学习中,坐标去噪任务贡献了最大比例的信息,对比学习负责稳定特征空间结构,而生成损失更像是一个正则或数据增广的角色。
但是要注意,消融结论和数据集规模强相关。在大规模非标注构象库上,生成任务能更好地帮助模型建立构象流形;在小规模精细数据集上,对比和去噪反而更高效。我建议你们做项目时不要抄固定结论,花两天时间跑消融矩阵,每一列的改动都记录下游任务指标,这比闭着眼睛堆模块有价值得多。
提到稳定性,我还会测一个东西:表征对构象噪声的鲁棒性。把验证集的坐标加上0.2埃的随机扰动,看看下游预测指标的波动幅度。一个好的几何表征应该具备天然的平滑性,小扰动不改变预测结果;如果发现波动很大,说明模型有一部分在死记坐标位置,没有真正学到几何常量,这通常需要增强去噪任务权重或换更严格的构象预处理策略。
5. 常见问题与排查技巧实录
5.1 构象生成质量差导致训练崩溃
这是我遇到最频繁的问题,没有之一。现象是:训练loss一直在下降,但下游任务验证集误差高得离谱。排查下来,八成是训练输入的构象空间里混杂了大量高能量不合理构象,模型的去噪任务无法区分“合理的构象变化”和“物理上不可能的畸变”。
解决办法分两条线走。一条线是回数据源头,把距离窗口内的镜像故障构象过滤掉,比如某些原子对的距离小于共价半径之和就该直接删除该样本。另一条线是在训练时给构象加一个RMSD分布筛选,同一分子保留的多个构象互相之间的RMSD不能太小也不能太大,太小了模型见不到多样性,太大了说明中间肯定有坏构象。我用的是RMSD在0.5到2.0埃之间的过滤窗口,实测效果很好。
5.2 等变网络不收敛或发散
等变网络训练发散时,容易让人怀疑“等变结构是不是不行”。实际上,很多时候只是细节处理有问题。先检查坐标中心化状况:如果输入分子的质心不在原点,模型的更新里会多出一部分不需要学习的平移成分,导致隐藏维度里有一路神经元专门在补平移动态,训练起来特别不稳。我都习惯在做预处理时就把所有构象对齐到质心为零。
另一个是梯度裁剪加上去抖动技巧:在等变更新里,coef那个分支的梯度经常会产生瞬间的大数值,因为当两个原子距离接近零时,除以距离就会把更新幅度放大到离谱。所以我在距离计算里加了最小值钳制clamp(min=1e-6),同时所有坐标分支的梯度只允许在[-10, 10]范围内流动。如果你在训练日志里看到loss曲线每隔几十步就跳一下,十有八九是这里的问题。
5.3 数据泄漏与评估陷阱
做分子几何表征时,数据泄漏比你想的更隐蔽。最常见的泄漏是把同一个分子的多个构象同时放在训练集和验证集里。虽然构象不同,但它们来自同一分子骨架,模型只要记住分子身份就能在验证集上“作弊”,下游指标看起来很漂亮,换到新分子上立刻原型毕露。正确做法是按分子骨架拆分子集,确保验证集里的骨架在训练集中一次都没出现过。
另外还有一个隐患是构象生成的随机种子泄漏。如果训练集和验证集的分子都是同一批RDKit种子生成的构象,模型可能在隐空间里学到种子的bias。我后期调整方案把随机种子也要和数据拆分绑定,每个分子拆分位置重新生成随机种子。这个坑很冷门,但排查真会浪费一整周。
5.4 常见问题速查表
| 现象 | 可能原因 | 处理思路 |
|---|---|---|
| 预训练loss很低,下游精度差 | 预训练任务和下游任务目标错位 | 调整任务权重,增加去噪比例 |
| 等变性测试不过 | 图构建阶段混入绝对坐标特征 | 检查边特征,保证只用距离和相对方向 |
| 同一个分子不同构象embedding距离过大 | 对比损失权重过低 | 提高对比权重,降低生成权重 |
| 显存很快打满 | 边数量爆炸 | 缩小截断半径或k近邻数,考虑子图采样 |
| 坐标更新输出NaN | 距离接近零导致除零 | 检查距离clamp,降低学习率 |
6. 一些我在实操中的个人体会
搞完这一整套流程之后,我的一个明显体会是:分子几何表征的真正价值不是做出一个能生成漂亮构象的模型,而是让几何信息以符合物理规律的方式进入模型的语言系统。等变网络解决的问题就是“坐标系变了,语义不散架”,这件事一定不要靠数据增广硬学,要放在模型结构里。
具体操作层面,我最后还想分享一个很有用的细节:在预训练阶段对输入构象做随机旋转加微小抖动,这个做法被很多人以为是多此一举,但实测下来,它对于训练稳定性尤其是坐标更新分支的帮助非常大。原因在于,即使等变结构理论上是旋转等变的,数值计算里的浮点误差在大的旋转角度下还是会累积出差异,数据增强等于给模型打了一针疫苗,让它在很宽的输入范围内都保持一致的输出行为。
另外,有一个值得往下走的方向:如果把几何表征和力场物理约束直接嵌在一个损失里,比如要求模型预测的隐空间坐标变化方向接近分子动力学模拟的低频振动模式,这类混合方案未来可能会比纯数据驱动的几何表征更接近“理解分子”。我还没完全跑通,但在小规模体系上已经有了一些迹象,有兴趣的人可以往这个方向试试看。
这个项目的下一步,我打算把几何表征模型和构象生成器融合成一个Encoder-Decoder框架,共用一套等变encoder,但把decoder从扩散模型换成更轻量化的规范流模型,这样既能保持生成的多样性,也能确保encoder学到的表征不会在反向传播中被生成损失带偏。路还长,但每一步踩实了,后面的成果就水到渠成。