云计算百科
云计算领域专业知识百科平台

AlphaFold原理与蛋白质结构预测实战

AlphaFold原理与蛋白质结构预测实战

一、引言

"蛋白质折叠问题"被称为生物学50年未解难题。2020年 DeepMind 的 AlphaFold2 在 CASP14 中以原子级精度(中位误差0.96Å)预测蛋白质3D结构,震惊科学界。2024年 AlphaFold3 进一步扩展到所有生物分子。

本文将深入解析 AlphaFold2 的核心架构:MSA处理、Pair表示、结构模块(IPA)、回收机制,并提供完整推理代码。

二、蛋白质结构基础

# 蛋白质由20种氨基酸组成
AMINO_ACIDS = "ACDEFGHIKLMNPQRSTVWY" # 20种标准氨基酸

# 结构层次:
# 序列(1D) → 二级结构(α螺旋/β折叠) → 三级结构(3D坐标)
# 每个残基由主链原子(N, Cα, C) + 侧链原子组成

# 目标: 给定氨基酸序列 → 预测每个原子的3D坐标

三、AlphaFold2 架构

输入序列: "MKFLILFNILV…"

├──→ [基因数据库搜索] → MSA (多序列比对)
│ └── JackHMMER/Mmseqs2 → 同源序列

├──→ [结构数据库搜索] → Templates (结构模板)
│ └── HHSearch → 已知结构


┌──────────────────────────────┐
│ Evoformer (48层) │
│ ├── MSA Stack │ 行注意力/列注意力
│ │ ├── Row-wise Gated │
│ │ ├── Column-wise Gated │
│ │ └── Transition │
│ │ │
│ ├── Pair Stack │ Pair偏置更新
│ │ ├── Triangle Multiplicative│
│ │ ├── Triangle Self-Attention│
│ │ └── Transition │
│ │ │
│ └── MSA ↔ Pair 信息交换 │ Outer Product Mean
└──────────────┬───────────────┘


┌──────────────────────────────┐
│ Structure Module (8层) │
│ ├── IPA (Invariant Point │
│ │ Attention) │ SE(3)等变Transformer
│ │ │
│ └── Backbone Update │ 预测旋转+平移
│ (每个残基的刚体变换) │
└──────────────┬───────────────┘


3D坐标 (N, Cα, C per residue)

3.1 Evoformer核心:三角乘法更新

import torch
import torch.nn as nn
import torch.nn.functional as F

class TriangleMultiplication(nn.Module):
"""三角乘法更新:利用Pair矩阵的几何约束
zij ← zij + Σ_k z_ik · z_kj (行更新)
zij ← zij + Σ_k z_ki · z_jk (列更新)
"""

def __init__(self, c_z=128, c_hidden=128):
super().__init__()
self.layer_norm_in = nn.LayerNorm(c_z)
self.layer_norm_out = nn.LayerNorm(c_z)

# 门控线性层
self.left_projection = nn.Linear(c_z, c_hidden)
self.right_projection = nn.Linear(c_z, c_hidden)
self.left_gate = nn.Linear(c_z, c_hidden)
self.right_gate = nn.Linear(c_z, c_hidden)

self.output_projection = nn.Linear(c_hidden, c_z)
self.gate = nn.Linear(c_z, c_z)

def forward(self, z, mask=None):
"""z: [B, N_res, N_res, c_z] Pair表示"""
z = self.layer_norm_in(z)

# 左投影和门控
left_proj = self.left_projection(z) # [B,N,N,c_hidden]
left_gate = torch.sigmoid(self.left_gate(z))
left = left_proj * left_gate

# 右投影和门控
right_proj = self.right_projection(z)
right_gate = torch.sigmoid(self.right_gate(z))
right = right_proj * right_gate

# 行更新: out_ij = Σ_k left_ik * right_jk
# 等价于 out = left @ right^T(在通道维度求和)
# 实现为einops风格: b i k c, b j k c -> b i j c
left = left.permute(0, 3, 1, 2) # [B, c_hidden, N, N]
right = right.permute(0, 3, 2, 1) # [B, c_hidden, N, N]

z_update = torch.einsum('bcij,bcjk->bcik', left, right)
z_update = z_update.permute(0, 2, 3, 1) # [B, N, N, c_hidden]

# 输出门控
z_update = self.output_projection(z_update)
gate = torch.sigmoid(self.gate(self.layer_norm_out(z)))

return z + gate * z_update

3.2 IPA(不变点注意力)

class InvariantPointAttention(nn.Module):
"""SE(3)等变的不变点注意力
核心: 用3D坐标(r,t)投影到注意力的key/query空间
"""

def __init__(self, c_s=384, c_z=128, n_heads=12, n_query_points=4, n_value_points=8):
super().__init__()
self.n_heads = n_heads
self.n_query_points = n_query_points
self.n_value_points = n_value_points

# 线性投影
self.linear_q = nn.Linear(c_s, c_s)
self.linear_kv = nn.Linear(c_s, 2 * c_s)

# 点投影(不变性来自刚体变换下的几何距离)
self.linear_q_points = nn.Linear(c_s, n_heads * n_query_points * 3)
self.linear_kv_points = nn.Linear(c_s, n_heads * (n_query_points + n_value_points) * 3)

self.linear_b = nn.Linear(c_z, n_heads)
self.linear_out = nn.Linear(c_s, c_s)

def forward(self, s, z, rigids, mask=None):
"""
s: [B, N, c_s] 单表示
z: [B, N, N, c_z] Pair表示
rigids: 每个残基的刚体变换(旋转矩阵R + 平移t)
"""

B, N, c_s = s.shape
H = self.n_heads
d_h = c_s // H

# 1. 标准注意力投影
q = self.linear_q(s).reshape(B, N, H, d_h) # [B,N,H,d_h]
k, v = self.linear_kv(s).chunk(2, dim=1)
k = k.reshape(B, N, H, d_h)
v = v.reshape(B, N, H, d_h)

# 2. 不变点注意力的点投影
# 为每个残基生成3D查询点和关键点
q_pts = self.linear_q_points(s).reshape(B, N, H, self.n_query_points, 3)
k_pts = self.linear_kv_points(s).reshape(B, N, H, self.n_query_points + self.n_value_points, 3)
k_pts, v_pts = torch.split(k_pts, [self.n_query_points, self.n_value_points], dim=3)

# 3. 将点投影到全局坐标(应用刚体变换 R·p + t)
R, t = rigids # R: [B,N,3,3], t: [B,N,3]

# 全局查询点
q_pts_global = torch.einsum('bnrc,bnhpc->bnhpr', R, q_pts) + t[:,:,None,None,:]
# 全局关键点
k_pts_global = torch.einsum('bnrc,bnhpc->bnhpr', R, k_pts) + t[:,:,None,None,:]

# 4. 计算基于点距离的注意力偏置
# dist_ij = ||q_i_p – k_j_p||²
q_expand = q_pts_global.unsqueeze(2) # [B,N,1,H,P_q,3]
k_expand = k_pts_global.unsqueeze(1) # [B,1,N,H,P_k,3]

# 成对平方距离
dist_sq = ((q_expand k_expand) ** 2).sum(dim=1) # [B,N,N,H,P_q,P_k]
point_bias = 0.5 * dist_sq.mean(dim=(1, 2)) # 平均距离作为偏置

# 5. 标准注意力计算
attn_logits = torch.einsum('bihd,bjhd->bhij', q, k) * (d_h ** 0.5)

# 加入Pair偏置
pair_bias = self.linear_b(z).permute(0, 3, 1, 2) # [B,H,N,N]
attn_logits = attn_logits + pair_bias + point_bias

# 软注意力
if mask is not None:
attn_logits = attn_logits.masked_fill(~mask[:,None,None,:], 1e9)
attn = F.softmax(attn_logits, dim=1)

# 6. 聚合Value
out = torch.einsum('bhij,bjhd->bihd', attn, v)
out = out.reshape(B, N, c_s)
return self.linear_out(out)

四、完整推理流程

from alphafold.model import model
from alphafold.data import pipeline, templates

class AlphaFoldInference:
def __init__(self, model_params_path, database_path):
# 加载模型
self.model_config = model.config(model_name="model_1")
self.model_params = model.get_model_haiku_params(model_name="model_1", data_dir=model_params_path)
self.model_runner = model.RunModel(self.model_config, self.model_params)

# 数据管道
self.data_pipeline = pipeline.DataPipeline(
jackhmmer_binary_path=f"{database_path}/jackhmmer",
hhblits_binary_path=f"{database_path}/hhblits",
uniref90_database_path=f"{database_path}/uniref90",
mgnify_database_path=f"{database_path}/mgnify",
bfd_database_path=f"{database_path}/bfd",
uniclust30_database_path=f"{database_path}/uniclust30",
pdb70_database_path=f"{database_path}/pdb70",
template_featurizer=templates.HhsearchHitFeaturizer(...),
)

def predict(self, sequence, name="query"):
"""端到端预测"""
features = self.data_pipeline.process(
input_fasta_path=f">{name}\\n{sequence}",
msa_output_dir="/tmp/msa"
)

# 运行模型(回收3次)
prediction = self.model_runner.process_features(
features,
random_seed=42
)

# 提取结果
result = {
"plddt": prediction["plddt"], # 预测置信度
"distogram": prediction["distogram"],
"structure_module": {
"final_atom_positions": prediction["structure_module"]["final_atom_positions"],
"final_atom_mask": prediction["structure_module"]["final_atom_mask"],
},
"predicted_aligned_error": prediction.get("predicted_aligned_error"),
}

return result

def save_pdb(self, prediction, output_path):
"""保存为PDB格式"""
from alphafold.common import protein

plddt = prediction["plddt"]
atom_positions = prediction["structure_module"]["final_atom_positions"]

# 转为Protein对象
unrelaxed_protein = protein.from_prediction(
features=None,
prediction=prediction,
)

with open(output_path, 'w') as f:
f.write(protein.to_pdb(unrelaxed_protein))

print(f"PDB saved: {output_path}")

# 使用
af = AlphaFoldInference("./params", "./databases")
result = af.predict("MKFLILFNILVCLAVAA…")
af.save_pdb(result, "prediction.pdb")
print(f"Mean pLDDT: {result['plddt'].mean():.2f}")

五、AlphaFold3 新特性

# AlphaFold3关键改进:
# 1. 扩散模块替代结构模块 → 直接生成原子坐标
# 2. 支持所有生物分子(蛋白质/DNA/RNA/配体/离子)
# 3. Pairformer替代Evoformer → 简化MSA处理
# 4. 置信度预测更准确

# 使用AlphaFold3 (通过API)
import requests

def predict_alphafold3(sequences):
response = requests.post(
"https://alphafold3-api.example.com/predict",
json={
"sequences": [
{"protein": {"sequence": "MKFLIL…", "id": "A"}},
{"ligand": {"smiles": "CC(=O)OC1=CC=CC=C1C(=O)O", "id": "B"}}
]
}
)
return response.json()

六、关键指标解读

指标含义好阈值
pLDDT 预测局部距离误差 >90 高置信, <50 低置信
PAE 预测对齐误差 <5Å 可信, >15Å 不可靠
pTM 预测模板建模分数 >0.8 高置信
ipTM 界面pTM(复合体) >0.8 高置信

七、总结

AlphaFold2的核心设计:

  • MSA + Pair → Evoformer — 48层深度处理共进化信息
  • IPA(不变点注意力) — SE(3)等变的3D结构生成
  • 回收机制 — 3次迭代逐步精化预测
  • 端到端训练 — 从序列直接预测3D坐标
  • AlphaFold3引入扩散模型和全原子支持,将蛋白质结构预测推向了新的高度。

    赞(0)
    未经允许不得转载:网硕互联帮助中心 » AlphaFold原理与蛋白质结构预测实战
    分享到: 更多 (0)

    评论 抢沙发

    评论前必须登录!