145

ChemBERTa 论文精读:SMILES Transformer 的价值和局限

ChemBERTa 把 BERT 用在 SMILES 上。这篇讲清 SMILES 预训练的价值边界与和指纹基线的真实差距。

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 增强对所有基于字符串的模型都有效,成本低,值得默认开启。

延伸资源