143

ESM-2 论文精读:蛋白语言模型为什么能预测结构,但不是 AlphaFold 替代品

ESM-2 证明了语言模型能隐式学到蛋白结构信息。这篇讲清它的能力范围,以及它不能替代 AlphaFold 的原因。

ESM-2(Lin 等,Science 2023)是规模最大、应用最广的蛋白语言模型。它最重要的科学结论是:仅通过在海量序列上做掩码语言建模,模型的内部表示中就会涌现出结构信息——不需要任何结构标注。

训练方式与规模

# 预训练任务:掩码语言建模(MLM)
#   随机遮住约 15% 的氨基酸,让模型预测
#   MKTA[MASK]IAKQ → 预测 [MASK] = Y
#
# 为什么这能学到结构?
#   要正确预测被遮住的氨基酸,模型必须理解:
#     - 该位置的物理化学约束(疏水核心还是表面?)
#     - 与其它位置的协同约束(共进化)
#     - 二级结构倾向
#     - 功能位点的保守性
#   → 这些约束本质上来自结构与功能的要求
#
# 模型规模(参数量):
#   esm2_t6_8M      8M      快,表示较弱
#   esm2_t12_35M    35M
#   esm2_t30_150M   150M
#   esm2_t33_650M   650M    【最常用,性价比好】
#   esm2_t36_3B     3B      ESMFold 使用
#   esm2_t48_15B    15B     最大
#
# 【规模效应】:
#   模型越大,从注意力图中恢复接触图的准确度越高
#   → 这是「结构信息随规模涌现」的直接证据
#
# 训练数据:UniRef50/90,约数千万到数亿条序列

为什么不是 AlphaFold 的替代品

AlphaFold2 ESM-2 / ESMFold
结构信息来源 显式的 MSA 隐式的语言模型表示
准确度(有深 MSA) 明显更高 较低
准确度(无 MSA) 大幅下降 相对稳健
速度 慢(MSA 搜索) 快数十到数百倍
复合物 支持(Multimer) 不支持
适用场景 需要高质量结构 大批量、MSA 稀少

核心区别:显式 MSA 提供的共进化信号,比语言模型隐式学到的更精确。语言模型学到的是「统计上的一般规律」,而 MSA 提供的是「这个特定蛋白家族的具体约束」——后者对精确建模更有价值。

ESM-2 表示的真正强项

# 结构预测不是 ESM-2 最有价值的应用 ——
# 【蛋白性质预测才是】
#
# 因为:
#   1) 这些任务没有 AlphaFold 这样的强对手
#   2) 表示可以直接喂给简单的下游模型
#   3) 速度快,适合大规模筛选
#
# 典型应用:
#   - 蛋白稳定性预测(Tm、ΔΔG)
#   - 溶解度与表达量预测
#   - 变异效应预测(某个突变是否有害)
#   - 抗体可开发性(见 367)
#   - 酶活性与底物特异性
#   - 亚细胞定位、信号肽预测
#   - 蛋白功能注释

import torch
import esm

model, alphabet = esm.pretrained.esm2_t33_650M_UR50D()
batch_converter = alphabet.get_batch_converter()
model.eval()
if torch.cuda.is_available():
    model = model.cuda()

def get_embeddings(sequences, layer=33, batch_size=8):
    """提取序列级表示"""
    out = []
    for i in range(0, len(sequences), batch_size):
        chunk = sequences[i:i+batch_size]
        data = [(f"p{j}", s) for j, s in enumerate(chunk)]
        _, _, tokens = batch_converter(data)
        if torch.cuda.is_available():
            tokens = tokens.cuda()
        with torch.no_grad():
            res = model(tokens, repr_layers=[layer])
        reps = res["representations"][layer]
        for j, s in enumerate(chunk):
            # 跳过起止 token,对残基取平均
            out.append(reps[j, 1:len(s)+1].mean(0).cpu().numpy())
    return out

# 下游:直接用简单模型
from sklearn.ensemble import RandomForestRegressor
X = get_embeddings(train_seqs)
model_rf = RandomForestRegressor(n_estimators=500).fit(X, train_y)

# 【实用提示】:
#   1) 取哪一层的表示影响很大
#      → 通常最后一层或倒数第二层最好,但应该试
#   2) 平均池化 vs CLS token vs 注意力池化
#      → 平均池化通常最稳健
#   3) 【冻结表示 + 简单下游模型】常常够用
#      → 微调整个 ESM 成本高得多,收益不一定大
#   4) 表示维度很高(650M 模型是 1280 维)
#      → 小数据时容易过拟合,可先降维

零样本变异效应预测

# ESM-2 的一个漂亮应用:不需要任何训练数据
#
# 原理:
#   模型预测每个位置各氨基酸的概率
#   如果野生型氨基酸概率高而突变型概率低,
#   说明该突变「不符合进化规律」→ 可能有害
#
#   得分 = log P(突变型) - log P(野生型)

import torch
import esm

def mutation_effect(sequence, position, wt_aa, mt_aa, model, alphabet):
    """零样本预测单点突变的效应(position 从 1 开始)"""
    batch_converter = alphabet.get_batch_converter()
    _, _, tokens = batch_converter([("seq", sequence)])

    # 遮住目标位置
    idx = position  # tokens[0][0] 是 CLS
    tokens_masked = tokens.clone()
    tokens_masked[0, idx] = alphabet.mask_idx

    with torch.no_grad():
        logits = model(tokens_masked)["logits"]
    log_probs = torch.log_softmax(logits[0, idx], dim=-1)

    return (log_probs[alphabet.get_idx(mt_aa)].item()
            - log_probs[alphabet.get_idx(wt_aa)].item())

# 用途:
#   - 快速筛选大量突变体(见 366)
#   - 蛋白工程中的初步排序
#   - 临床变异的致病性初判
#
# 【局限】:
#   反映的是「进化上是否常见」,
#   不等于「功能上是否有害」
#   → 有些突变进化上罕见但功能正常
#   → 有些改善功能的突变(工程目标)恰恰是罕见的
#   → 【做定向进化时,这个信号可能指向错误方向】

使用中的注意事项

  • 长序列的显存需求:注意力是二次复杂度,超过约 1000 残基需要分块处理;
  • 对人工设计蛋白可能失效ESM 学的是自然进化的规律,人工设计蛋白没有进化史——表示可能不适用;
  • 不同层的表示差异大:浅层偏向局部化学性质,深层偏向全局与功能——应该在验证集上比较不同层
  • 不要过度解读注意力图:虽然某些注意力头对应接触,但这是统计发现,不能逐案例当作结构证据;
  • 与传统特征比较对某些任务,简单的氨基酸组成 + 理化性质特征已经很强——应该做这个对照。

关键要点

  • 结构信息随模型规模涌现——这是纯序列自监督学习的重要科学结论;
  • 显式 MSA 比语言模型隐式学到的共进化信号更精确,所以不能替代 AlphaFold;
  • 它真正的强项是蛋白性质预测,冻结表示 + 简单下游模型常常够用;
  • 零样本变异效应反映「进化上是否常见」,做定向进化时可能指向错误方向

延伸资源