ChemBERTa(Chithrananda 等,2020)是最早把 BERT 式预训练应用到 SMILES 的工作之一。它证明了这条路可行——但也清楚地显示了「可行」与「优于传统方法」之间的距离。
方法
# 完全照搬 NLP 的 BERT 流程:
#
# 1) 分词
# 把 SMILES 切成 token
# 方案 A:字符级(C, c, 1, (, =, N, ...)
# 方案 B:BPE(字节对编码,学出常见的子串)
# 方案 C:正则表达式(按化学语义切,如 [nH]、Cl、Br 作为整体)
# → 【方案 C 通常最好】—— 保证了原子级 token 的正确性
#
# 2) 预训练
# 掩码语言建模:遮住 15% 的 token,预测它们
# 数据:PubChem 的 1000 万~7700 万 SMILES
#
# 3) 微调
# 在下游任务(性质预测)上微调
#
# ChemBERTa-2(2022)的改进:
# - 更大的数据(7700 万)
# - 增加【多任务回归预训练】:
# 除 MLM 外,还预测 200 个计算得到的分子描述符
# → 【这个改进的效果比单纯增大数据明显】
# → 说明「预训练任务的设计」比「数据规模」更关键
诚实的性能对比
| 方法 | 在 MoleculeNet 上的典型表现 |
|---|---|
| ECFP + 随机森林/LightGBM | 强基线,很多任务上难以超越 |
| Chemprop(D-MPNN) | 大数据集上占优(见 131《Chemprop / D-MPNN 论文精读》) |
| ChemBERTa | 与基线相当,部分任务略好或略差 |
| ChemBERTa-2 | 有改善,但仍非全面领先 |
论文本身对此比较坦诚——它没有宣称全面超越,而是把重点放在「证明这条路径可行」与「预训练规模的影响」上。这种诚实在领域中并不常见,值得肯定。
为什么 SMILES 预训练的收益有限
# 原因一:【SMILES 不是自然语言】
# 自然语言中,BERT 的成功依赖于:
# - 词的语义可以从上下文推断
# - 海量多样的文本
# - 下游任务需要复杂的语义理解
#
# 而 SMILES:
# - 语法严格且简单(不像自然语言那样歧义丰富)
# - 「语义」就是分子结构,已经完全确定
# - 【没有需要「理解」的隐含含义】
#
# → 分子的信息已经完全编码在结构中,
# 不需要像语言那样「推断」
#
# 原因二:【指纹已经很好地编码了结构】
# ECFP 直接枚举子结构,
# 这对多数性质预测任务已经足够
# → 预训练要提供的「额外信息」空间有限
#
# 原因三:【SMILES 的字符串形式带来噪声】
# - 同一分子多种写法(可用增强缓解)
# - 环闭合标记打破了局部性
# - 分支括号的嵌套
# → 模型要先「学会解析 SMILES」,
# 这部分能力对下游任务没有直接价值
#
# 原因四:【下游数据集太小】
# 预训练的价值在于「用无标注数据补充标注数据的不足」
# 但如果下游只有 1000 个样本,
# 模型的容量本身就用不上(见 132)
#
# 【什么时候预训练确实有价值】:
# - 下游数据极少(< 200 个样本)
# - 需要生成而非只是预测
# - 需要分子的连续表示做优化(见 149)
使用
from transformers import AutoTokenizer, AutoModelForSequenceClassification
import torch
model_name = "DeepChem/ChemBERTa-77M-MTR" # 多任务回归预训练版
tokenizer = AutoTokenizer.from_pretrained(model_name)
# ---- 提取表示 ----
from transformers import AutoModel
encoder = AutoModel.from_pretrained(model_name).eval()
def embed_smiles(smiles_list, batch_size=32):
out = []
for i in range(0, len(smiles_list), batch_size):
batch = smiles_list[i:i+batch_size]
enc = tokenizer(batch, padding=True, truncation=True,
max_length=256, return_tensors="pt")
with torch.no_grad():
h = encoder(**enc).last_hidden_state
# 用 attention_mask 做正确的平均池化
mask = enc["attention_mask"].unsqueeze(-1).float()
pooled = (h * mask).sum(1) / mask.sum(1)
out.append(pooled)
return torch.cat(out).numpy()
# ---- 微调 ----
model = AutoModelForSequenceClassification.from_pretrained(
model_name, num_labels=1, problem_type="regression")
# 【务必同时跑的基线】:
from rdkit import Chem
from rdkit.Chem import rdFingerprintGenerator
import numpy as np
import lightgbm as lgb
gen = rdFingerprintGenerator.GetMorganGenerator(radius=2, fpSize=2048)
def ecfp(smiles_list):
X = []
for s in smiles_list:
m = Chem.MolFromSmiles(s)
X.append(np.array(gen.GetFingerprint(m)) if m else np.zeros(2048))
return np.array(X)
baseline = lgb.LGBMRegressor(n_estimators=1000, learning_rate=0.05,
num_leaves=31, verbose=-1)
baseline.fit(ecfp(train_smiles), train_y)
# 【如果 ChemBERTa 打不过这个基线,就用基线】
# 训练时间:LightGBM 几秒,ChemBERTa 微调几十分钟
# 部署复杂度:LightGBM 一个 pickle 文件,ChemBERTa 需要 GPU
SMILES 增强:一个有效的技巧
# 同一个分子有多种 SMILES 写法
# 苯:c1ccccc1, C1=CC=CC=C1, c1ccc(cc1) ...
#
# 增强做法:
# 训练时随机生成不同的 SMILES 写法
# → 相当于数据增强
# → 让模型学会「不同写法是同一分子」
#
# 推理时:
# 对同一分子生成多个 SMILES,预测后取平均
# → 类似测试时增强(TTA),能提升稳定性
from rdkit import Chem
import random
def randomize_smiles(smiles, n=10, seed=0):
mol = Chem.MolFromSmiles(smiles)
if mol is None:
return []
rng = random.Random(seed)
out = set()
n_atoms = mol.GetNumAtoms()
for _ in range(n * 3):
idx = list(range(n_atoms))
rng.shuffle(idx)
renum = Chem.RenumberAtoms(mol, idx)
out.add(Chem.MolToSmiles(renum, canonical=False))
if len(out) >= n:
break
return list(out)
# 这个技巧对所有基于 SMILES 的模型都适用
# 【成本低,效果稳定,值得默认开启】
务实的结论
- 先跑 ECFP + LightGBM 基线:几分钟就能知道任务的难度与合理性能范围;
- 只在基线不够时考虑预训练模型,并且要求它显著超越(超过种子间方差);
- 拼接常常比替代好:ChemBERTa 表示 + ECFP + 描述符一起喂给树模型;
- 小数据时冻结表示,不要微调;
- SMILES 增强几乎总是有帮助,成本低;
- 部署成本要计入决策:GPU 依赖、模型体积、推理延迟。
关键要点
- SMILES 不是自然语言——分子信息已完全编码在结构中,不需要「推断隐含语义」;
- ChemBERTa-2 的多任务回归预训练比单纯增大数据效果更明显——任务设计比规模关键;
- 与 ECFP + LightGBM 基线相比没有全面优势,而成本高得多;
- SMILES 增强对所有基于字符串的模型都有效,成本低,值得默认开启。
延伸资源
- MolBERT:146《MolBERT 论文精读》;MolFormer:147《MolFormer 论文精读》;Uni-Mol:142《Uni-Mol 论文精读》;
- Chemprop:131《Chemprop / D-MPNN 论文精读》;分子指纹:035《分子指纹是什么》;GNN:159《图神经网络 GNN》。