- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
Dedal(Deep embedding and alignment of protein sequences)是 Google Research 在 google-research 仓库中开源的一套蛋白质序列比对与同源检测工具包,其核心思想是让 Transformer 编码器在监督信号(掩码语言建模、成对序列比对、同源检测)下学习序列嵌入,并借助可微的"扰动 Smith-Waterman"(Perturbed Smith-Waterman)把传统动态规划比对变成端到端可训练的网络层。本文以 dedal/README.md 为骨架,结合仓库源码与 Gin 配置,完整讲解环境安装、TensorFlow Hub 预训练模型调用、推理管线原理,以及从零开始多阶段训练的全部操作细节。
安装与环境准备
Dedal 并不是一个独立的 pip 包,而是 google-research 仓库中的一个子目录。安装的第一步是克隆整个仓库:
git clone https://github.com/google-research/google-research.git然后在仓库根目录(即google_research文件夹)下,通过 dedal/requirements.txt 安装全部依赖:
pip install -r dedal/requirements.txtrequirements.txt 中声明的核心依赖包括:
| 依赖 | 版本要求 | 用途 |
|---|---|---|
| absl-py | >=0.7.0 | 命令行 flag 解析与日志(main.py依赖) |
| gin-config | >=0.4.0 | 声明式超参数/架构配置(全部.gin文件依赖) |
| numpy | >=1.18.4 | 数值计算与后处理 |
| tensorflow | >=2.3.0 | 模型构建与训练 |
| tensorflow_datasets | >=3.0.0 | 数据集基础设施 |
| tensorflow_probability | >=0.1.0 | 扰动 SW 所需的分布采样(Gumbel/Normal) |
| tf-models-nightly | 最新 | TF Hub 加载与模型工具链 |
需要特别指出的是,源码中的导入语句形如from dedal import infer(参见 infer.py),要求以仓库根目录为工作区并保证dedal包可见,因此安装与后续所有命令都应在克隆出的google_research文件夹内执行。
使用 TensorFlow Hub 预训练模型
DEDAL 预训练模型以 SavedModel 形式发布在 TensorFlow Hub 上(模块名google/dedal/3)。模型已经包含了从原始蛋白质序列到比对结果、同源判定 logits 以及逐位嵌入的完整计算图。
输入输出协议
模型输入是一个tf.Tensor<tf.int32>[2B, 512],表示一批 B 对、按最大长度 512 右填充的序列对(512 已包含特殊的 EOS 结束符)。批次中序列对必须连续排列:inputs[2*b]与inputs[2*b + 1]构成第 b 个序列对(b 从 0 到 B-1)。
默认情况下模型运行在alignment(比对)模式,返回一个 Python dict,包含:
sw_scores:tf.Tensor<tf.float32>[B],比对分数(Smith-Waterman 得分);homology_logits:tf.Tensor<tf.float32>[B],同源检测 logits;paths:tf.Tensor<tf.float32>[B, 512, 512, 9],预测的比对路径;sw_params:由三个tf.Tensor<tf.float32>[B, 512, 512]组成的元组,分别对应上下文相关的 Smith-Waterman 参数——替换得分(substitution scores)、gap open 罚分与 gap extend 罚分。
此外模型还提供额外的签名以运行embedding(嵌入)模式,此时仅返回单个tf.Tensor<tf.float32>[2B, 512, 768],即每条输入序列每个位置的嵌入向量。
最小推理示例
README 中的完整示例(以论文 Figure 3 的"Gorilla"与"Mallard"两条序列为例)如下:
import tensorflow as tf import tensorflow_hub as hub from dedal import infer # Requires google_research/google-research. dedal_model = hub.load('https://tfhub.dev/google/dedal/3') # "Gorilla" and "Mallard" sequences from [1, Figure 3]. protein_a = 'SVCCRDYVRYRLPLRVVKHFYWTSDSCPRPGVVLLTFRDKEICADPRVPWVKMILNKL' protein_b = 'VKCKCSRKGPKIRFSNVRKLEIKPRYPFCVEEMIIVTLWTRVRGEQQHCLNPKRQNTVRLLKWY' # Represents sequences as `tf.Tensor<tf.float32>[2, 512]` batch of tokens. inputs = infer.preprocess(protein_a, protein_b) # Aligns `protein_a` to `protein_b`. align_out = dedal_model(inputs) # Retrieves per-position embeddings of both sequences. embeddings = dedal_model.call(inputs, embeddings_only=True) # Postprocesses output and displays alignment. output = infer.expand( [align_out['sw_scores'], align_out['paths'], align_out['sw_params']]) output = infer.postprocess(output, len(protein_a), len(protein_b)) alignment = infer.Alignment(protein_a, protein_b, *output) print(alignment) # Displays the raw Smith-Waterman score and the homology detection logits. print('Smith-Waterman score (uncorrected):', align_out['sw_scores'].numpy()) print('Homology detection logits:', align_out['homology_logits'].numpy())代码中infer.preprocess、infer.expand、infer.postprocess与infer.Alignment全部定义在 infer.py 中,值得逐层拆解其内部机制。
推理管线源码级解析
preprocess:文本到 token 批次
infer.preprocess 将两条蛋白质字符串依次经过三个数据变换:
transforms.Encode(vocab=vocabulary.seqio_vocab, on=keys):用 SentencePiece 词表把氨基酸序列编码为 token id;transforms.EOS(vocab=vocabulary.seqio_vocab, on=keys):在序列末尾追加 EOS 特殊 token;transforms.CropOrPad(size=max_length, ...):裁剪或填充到固定长度 512(默认max_length=512)。
最终通过tf.stack([seqs['left'], seqs['right']], axis=0)拼成[2, 512]的批次,与 TF Hub 模块的输入协议严格对应。
expand 与 postprocess:从扁平 dict 恢复嵌套结构
TF Hub SavedModel 的输出往往被序列化为扁平 key(形如output_4_1_1)。infer.expand 的作用正是按 key 中的位置编号把扁平 dict 递归还原成嵌套 tuple;非 dict 输入则原样返回(no-op)。
infer.postprocess 则完成三项关键后处理:
- 用
tf.squeeze(axis=0)去掉 batch 维; - 对 gap open / gap extend 罚分取负号并广播成与替换得分相同的形状,三者沿最后一维堆叠为
[L1, L2, 3]的参数张量(负号是为了在展示时表达"罚分"语义); - 裁剪掉
length_1 × length_2之外的填充区域; - 通过
alignment.paths_to_state_indicators把 9 状态路径表示转换为 match / gap_open / gap_extend 三态指示。
Alignment 对象:可读的比对展示
infer.Alignment 在构造时调用expand()把路径张量解析为三条等长字符串:left_match(左序列)、right_match(右序列)与matches(连接符行)。连接符语义在_position_to_char中定义:match 状态下,若两字符完全相同输出|,若替换得分为正则输出:,否则输出.;gap 状态输出空格。__str__最终把它们排版为经典的三行比对视图,并附带首尾位置坐标。
该对象还暴露了几个直接可用的派生属性:
identity:matches中|的个数(精确匹配数);similarity:identity加上:的个数(考虑替换得分后的相似数);gaps:左右序列中出现-的位置总数。
架构原理:可微 Smith-Waterman 与多任务模型
理解预训练模型的行为,需要看其网络结构配置 configs/model/dedal.gin 与模型实现 models/dedal.py。
顶层模型dedal.Dedal由三部分组成:
- encoder:
encoders.TransformerEncoder,默认emb_dim=768、num_layers=6、num_heads=12、mlp_dim=3072、各类 dropout 均为 0.1,使用绝对位置编码(max_len=1024); - aligner:
aligners.SoftAligner,负责把一对序列嵌入转成 SW 参数与比对分数。其中替换得分由PairwiseBilinearDense(双线性相似度)产生,gap 罚分由ContextualGapPenalties产生,并分别用 bias 初始化为 11.0(gap open)与 0.0(gap extend),激活函数为 softplus,配合mask_penalty=1e9屏蔽填充区域; - heads:多任务输出头,包括嵌入头(掩码 LM 的逐 token 输出头
DensePerTokenOutputHead)与比对头(dedal.Selector直通 + 同源头homology.LogCorrectedLogits)。
比对分数计算的核心是 smith_waterman.py 中的扰动 Smith-Waterman(smith_waterman.perturbed_alignment_score,默认sigma=0.1;评测时切换为无扰动的unperturbed_alignment_score,soft 版本温度temp=0.1)。该实现通过wavefrontify/unwavefrontify把动态规划矩阵按反对角线重排(smith_waterman.py),实现可向量化、可求导的波前算法,从而让"比对分数"本身可以参与反向传播——这正是 DEDAL 能把比对任务端到端训练的关键。
训练数据与预处理
DEDAL 训练与评测使用的公开数据为:
- UniRef50:2018 年 3 月版本,用于掩码语言建模预训练;
- Pfam 34.0:用于成对比对与同源检测训练。
preprocessing/子目录包含复现论文预处理流程的全部代码。论文中各 Pfam-A seed 划分对应的序列标识符可下载获取(见 README 原文)。三个数据任务的 Gin 配置分别声明了各自的目录布局与必需绑定:
- configs/data/uniref50.gin:要求
UNIREF50_DATA_DIR下含train/validation/test三个子目录,内部为带表头的 CSV 文件;序列长度 1024,shuffle buffer 32768; - configs/data/pfam34_alignment.gin:要求
PFAM34_ALIGNMENT_DATA_DIR下含train/iid_validation/iid_test/ood_validation/ood_test五个子目录,内部为带表头的 TSV 文件;序列长度 512,比对状态长度 1025; - configs/data/pfam34_homology.gin:目录布局与 alignment 相同,特征转换器
HomologyFeatureConverter默认fine_grained_labels=False。
三个数据配置还共同依赖三个必需词表绑定(MAIN_VOCAB_PATH、TOKEN_REPLACE_VOCAB_PATH、ALIGNMENT_PATH_VOCAB_PATH),指向 SeqIO SentencePiece 词表文件。
从零训练:命令与 Gin 配置详解
完整训练命令
在google_research文件夹下,同时训练掩码语言建模、成对序列比对与同源检测三个任务:
python3 -m dedal.main -- \ --base_dir /tmp/dedal \ --gin_config data/uniref50.gin \ --gin_config data/pfam34_alignment.gin \ --gin_config data/pfam34_homology.gin \ --gin_config model/dedal.gin \ --gin_config task/finetune.gin \ --gin_bindings UNIREF50_DATA_DIR=/path/to/masked_lm/data \ --gin_bindings PFAM34_ALIGNMENT_DATA_DIR=/path/to/alignment/data \ --gin_bindings PFAM34_HOMOLOGY_DATA_DIR=/path/to/homology/data \ --gin_bindings MAIN_VOCAB_PATH=/path/to/main/vocab \ --gin_bindings TOKEN_REPLACE_VOCAB_PATH=/path/to/token_replace/vocab \ --gin_bindings ALIGNMENT_PATH_VOCAB_PATH=/path/to/alignment_path/vocab \ --task train \ --alsologtostderrmain.py 定义了该入口的全部 flag:
| Flag | 默认值 | 说明 |
|---|---|---|
--base_dir | None | 保存 checkpoint 与日志的根目录 |
--reference_dir | None | 参考模型读取目录(若存在) |
--eval_in_train_job | True | 训练任务中是否同时运行 eval(对其他 task 忽略) |
--task | train | 可选train/eval/downstream |
--gin_config | [] | 可重复的 Gin 配置文件路径列表 |
--gin_bindings | [] | 换行分隔的 Gin 参数绑定 |
--config_path | dedal/configs | Gin 配置所在目录 |
main函数会把--gin_config中的相对路径拼接到--config_path下解析(main.py),随后构建分布策略并进入training_loop.TrainingLoop。值得注意的是,主循环对 worker 抢占(preemption)做了容错:捕获tf.errors.UnavailableError后会自动恢复并继续训练(main.py)。
多任务微调配置(task/finetune.gin)
configs/task/finetune.gin 是三个任务联合微调的核心配置,关键参数包括:
- 训练循环:
batch_size=(128, 128, 256)(分别对应 masked_lm / alignment / homology 三个数据构建器)、num_steps=1_000_000、num_eval_steps=100、num_steps_per_train_iteration=16; - 损失权重:掩码 LM 与同源检测各
weight=20.0,比对任务weight=1.0;三者分别使用SparseCategoricalCrossentropy(from_logits)、SmithWatermanLoss与BinaryCrossentropy; - 优化器:Adam,
epsilon=1e-08、clipnorm=1.0,学习率采用InverseSquareRootDecayWithWarmup(lr_max=1e-4、warmup_steps=8000); - 日志与 checkpoint:
PERIOD=2000步记录一次标量并保存一次 checkpoint,MAX_TO_KEEP=10; - 评测指标:掩码 LM 的准确率/交叉熵/困惑度,比对的
AlignmentPrecisionRecall(按序列一致性分桶统计)、AlignmentMSE、AlignmentStats、AlignmentScore、SWParamsStats,同源检测的BinaryAccuracy与 ROC/PR 的 AUC。
仓库还提供了其他任务变体配置,如 configs/task/masked_lm_pretraining.gin(仅预训练,num_steps=2_000_000、lr_max=1e-3、aligner_cls=None)、configs/task/finetune_alignment_only.gin、configs/task/finetune_without_homology.gin、configs/task/finetune_without_masked_lm.gin,以及 TAPE 基准系列(tape_fluorescence、tape_proteinnet、tape_remote_homology、tape_secondary_structure、tape_stability),可根据研究目标灵活组合。
训练运行与可视化
--base_dir指定的目录(上例为/tmp/dedal)用于写入 checkpoint 与日志指标。可视化指标只需在对应目录启动 TensorBoard:
%tensorboard --logdir /tmp/dedal训练中断后,直接重跑同一命令不会从头开始,而是从最近一个 checkpoint 恢复训练。checkpoint 保存频率与日志频率可在task系列 Gin 配置文件中调整,例如 configs/task/finetune.gin 中的PERIOD与MAX_TO_KEEP。
task 模式的三种语义
--taskflag 控制运行模式(见 main.py 与 README):
train:训练模式(默认,可同时附带 eval);eval:评测模式。评测时训练 checkpoint 会被动态加载直到最后一个,因此可以并行运行一个训练进程与一个评测进程,评测不会拖慢训练;downstream:下游任务训练,携带其自身的 eval。
如果希望训练与评测交替执行而非并行,可在训练循环中将separate_eval=False。
关于训练资源的说明
Transformer 架构在 CPU 上训练会非常缓慢,使用加速器可显著提升训练速度。README 明确说明:论文中的预训练 DEDAL 模型使用32 个 TPU v3 核心训练完成。如果只是复现或调参,建议在具备多卡 GPU 或 TPU 的环境中进行,并可按需调低num_steps等超参。
小结
DEDAL 的意义在于把序列比对这一经典算法任务纳入了端到端深度学习的框架:Transformer 编码器学习上下文相关的替换得分与 gap 罚分,可微 Smith-Waterman 层负责生成比对分数与路径,同源检测头在此基础上给出二分类 logits。通过本文介绍的 TF Hub 推理流程,可以在数行代码内获得可视化比对与同源打分;通过 Gin 配置体系与 main.py 入口,也可以完整复现掩码预训练、多任务微调、评测与下游任务训练的全流程。所有源码与配置均可从仓库的 dedal/ 目录(含 infer.py、models/dedal.py、smith_waterman.py、configs/)进一步深入研读。
说明:本文内容基于 dedal/README.md 及其对应源码与配置整理;DEDAL 遵循 Apache 2.0 许可,且"这不是 Google 的官方产品"(见 README 免责声明)。
- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
相关推荐
如何用 Python 实现 Smith-Waterman 算法完成 DNA 与蛋白质的局部序列比对
如何用 Python 实现 Smith Waterman 算法完成 DNA 与蛋白质的局部序列比对 如果你需要判断两段生物序列(DNA 碱基序列或蛋白质氨基酸序
示例工程西工大软院大一线性代数:nwpu-cram知识点与习题答案完整指南
西工大软院大一线性代数:nwpu cram知识点与习题答案完整指南 西北工业大学软件学院的大一线性代数课程是计算机科学专业的重要基础课,掌握好线性代数知识点和习
教程知识库教育MUMmer基因序列比对工具:快速完成DNA与蛋白质序列分析的终极指南
MUMmer基因序列比对工具:快速完成DNA与蛋白质序列分析的终极指南 MUMmer是一款强大的基因序列比对工具,专为快速比对DNA和蛋白质序列而设计。无论是处
数据分析
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考