尧图网络 高端网站定制 · 原创设计
免费咨询热线
400-888-6620
免费获取方案
AAV 衣壳蛋白机器学习模型训练指南:基于 tf.Estimator 的 RNN/CNN/逻辑回归实战解析(google-research/aav)
人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载本文围绕 google-research 仓库 aav/model_training/README.md 展开系统讲解其用于学习 AAV2 衣壳蛋白包装表型packaging phenotype的三类模型——RNN、CNN 与逻辑回归LR的训练流程、输入编码与正则化机制。读者读完将掌握如何通过 train.py 命令行入口训练单个模型副本、固定长度与变长序列编码的差异及由来58 2 × (28 1)、基于EarlyStopper的早停训练策略以及如何组合多个副本构成集成模型。该目录是《Deep diversification of an AAV capsid protein by machine learning》Nature Biotechnology, 2020的配套代码之一项目级安装指引可参考 aav/README.md 及 aav/model_and_dataset_analysis/README.md。目录结构与三种模型架构aav/model_training/目录的定位非常明确为论文中评估的3 种不同架构提供完整可训练的实现每个架构一个文件外加一个统一的训练入口文件架构说明rnn.pyRNN多层 LSTM 循环神经网络cnn.pyCNN一维卷积神经网络分类器lr.pyLR逻辑回归线性基线train.py训练入口所有架构共用的训练 harnesstrain_utils.py工具提供EarlyStopper早停实现原文档明确说明所有模型均基于TensorFlow 1.3 API实现并专门使用tf.Estimator框架。具体来说每个模型都被实现为一个 TensorFlow Estimator 的model_fn如 rnn.py 的rnn_model_fn、cnn.py 的cnn_model_fn、lr.py 的logistic_regression_model_fn训练与推理均通过tf.Datasets消费数据。从源码结构看三者都遵循同一套约定分别构建reuseFalse的训练子图和reuseTrue的测试子图输出pred_labelstf.argmax与pred_probastf.nn.softmax并在eval_metric_ops中统一注册accuracy、precision、recall三个评估指标——这正是早停逻辑得以跨架构复用的基础。一个值得注意的共性参数训练时每个 step 的 batch size 在三种模型中均为 25对应train.py中三组默认超参的batch_size: 25这是论文设置中的统一约定。训练入口 train.py命令行参数全解析训练入口 通过 absl flags 暴露了全部可配置项。原文档概括了训练骨架源码则给出了完整参数定义Flag默认值说明--model_dir必填无默认checkpoint 保存目录--train_path默认指向内部路径训练数据集TFRecord路径--validation_path默认指向内部路径验证数据集TFRecord路径--model必填枚举rnn/cnn/logistic--seq_encoder必填枚举varlen-id变长/fixedlen-id固定长度--hparamsNone超参覆盖格式hparam1value1,hparam2value2--eval_metricprecision早停监控指标可选precision/recall/accuracy--train_steps_per_eval500每次评估间隔的训练步数需 ≥ 每个 epoch 的步数--max_train_steps100000最大训练步数--early_stopper_num_evals_to_wait10指标不再提升时最多等待的评估次数--ref_seqDEEEIRTTNPVATEQYGSVSTNLQRGNR突变编码的参考WT序列--eval_throttle_secs10训练开始或上次评估后最少等待的秒数main()中的逻辑清晰可循train.py根据--seq_encoder选择编码器varlen-id映射到dataset_utils.ONEHOT_VARLEN_SEQUENCE_ENCODERfixedlen-id映射到ONEHOT_FIXEDLEN_MUTATION_ENCODER见 dataset_utils.py根据--model选择对应的model_fn与默认超参并为 CNN/LR 计算residue_encoding_size与seq_encoding_length (len(FLAGS.ref_seq) 1) * 2用tf.contrib.training.HParams合并默认值与--hparams覆盖项并强制注入model、seq_encoder两个字段分别构造训练/验证input_fn训练集shuffleTrue、num_epochsNone无限重复验证集shuffleFalse、num_epochs1且均drop_partial_batchesTrue调用train_model()执行Experiment.continuous_train_and_eval()并传入早停谓词训练结束后把最终超参写入model_dir/hparams.json。源码中的一段注释train.py特别提醒train_steps_per_iteration必须不小于训练 epoch 的步数否则在Experiment的实现中训练输入队列会在每次评估后被重置可能导致只训练到训练数据的一个子集——这是复现实验时最容易被忽略的坑。三组默认超参数三种架构的默认超参直接定义在train.py中L123-L156RNN_DEFAULT_HPARAMS_RNNseq_encodervarlen-idbatch_size25learning_rate0.01num_classes2num_layers2num_units100。CNN_DEFAULT_HPARAMS_CNNseq_encoderfixedlen-idbatch_size25learning_rate0.001num_classes2pool_width2conv_depth12conv_depth_multiplier2conv_width7feature_axis-1fc_size128fc_size_multiplier0.5。逻辑回归_DEFAULT_HPARAMS_LOGISTICseq_encoderfixedlen-idbatch_size25learning_rate0.01num_classes2。三种模型的源码级实现细节RNN多层 LSTMrnn.py 的build_rnn_inference_subgraph使用tf.nn.rnn_cell.MultiRNNCell堆叠num_layers2层、每层num_units100的LSTMCell通过tf.nn.dynamic_rnn处理变长输入传入sequence_lengthfeatures[sequence_length]将全部时间步输出做reduce_mean后接一个线性dense层得到 logits。训练优化器为RMSPropOptimizer学习率 0.01损失为 softmax 交叉熵。CNN一维卷积 批归一化cnn.py 的build_nn_inference_subgraph将展平的固定长度特征 reshape 为[batch, seq_encoding_length, residue_encoding_size]后依次执行conv1d(conv_depth12, kernel7, relu) → batch_norm → max_pool(2)、再conv1d(24, kernel7) → batch_norm → max_pool(2)展平后接fc(128) → fc(64)两个带批归一化的全连接层最后输出 2 类 logits。训练优化器为AdamOptimizer学习率 0.001并通过tf.GraphKeys.UPDATE_OPS的 control dependency 确保批归一化 moving mean/variance 正确更新cnn.py。逻辑回归线性基线lr.py 是最简单的基线将固定长度编码展平为[batch, seq_encoding_length * residue_encoding_size]后直接做线性dense输出训练优化器为FtrlOptimizer学习率 0.01。输入序列编码one-hot 与两种长度约定所有序列在进入模型前都先转换为特征向量。原文档指出编码发生在util.dataset_utils中encode_varlen/encode_fixedlen单个残基被编码为 one-hot 向量实现位于util.residue_encoding。对应源码如下residue_encoding.py 的ResidueIdentityEncoder基于 20 字母表ACDEFGHIKLMNPQRSTVWY定义于 config.py生成 one-hot 向量例如字母表ACDE中C编码为[0, 1, 0, 0]encoding_size20mutation_encoding.py 提供两个序列级编码器DirectSequenceEncoder将每个残基逐位 one-hot输出形状(len(seq), 20)的变长表示供 RNN 使用MutationSequenceEncoder基于参考序列输出形状(len(ref_seq)1, 2, 20)的固定长度表示供 CNN/LR 使用dataset_utils.py 的encode_fixedlen默认编码mutation_sequence字段与encode_varlen默认编码sequence字段通过tf.py_func在数据管道内完成上述转换标签统一取is_viable。固定长度编码为什么是 58这是理解该编码体系的关键。固定长度编码中的突变序列mutation mask语法如下见 mutation_encoding.py 的文档注释_表示野生型未突变位置大写字母表示单残基替换小写字母表示单残基插入每个 WT 位置有 2 个 slot替换 slot 插入 slot另有一个前缀位置prefix insertion也占 2 个 slot。给定 28 个 WT AAV2 残基位置DEEEIRTTNPVATEQYGSVSTNLQRGNR即R1_TILE21_WT_SEQ对应 AAV2 serotype 2 的 561–588 位见 config.py每个位置 1 个 slot 且可附带 1 次插入因此最大可表示的残基数58 2 × (28 1)即seq_encoding_length (len(ref_seq) 1) * 228 个 WT 位置 1 个前缀位置每个位置含替换与插入两个 slot。CNN 与 LR 使用该固定长度表示RNN 则直接使用原始序列的变长 one-hot 表示。模型正则化基于验证集 precision 的早停原文档对正则化策略的描述非常具体模型通过train_utils.EarlyStopper实现早停正则化所有架构共享同一个 hold-out 验证序列集来监控训练进度当模型在验证集上的precision 连续 10 个评估周期未提升时停止训练评估周期为每 500 步一次。对照 train_utils.py 的实现EarlyStopper是一个有状态谓词类构造参数num_evals_to_wait默认 10、metric_key默认precision、epsilon默认 1e-3early_stop_predicate_fn(tf_eval_results)每次评估后被Experiment.continuous_train_and_eval调用记录当前指标值若best_so_far epsilon curr则刷新最优值并重置等待计数否则等待计数 1当计数达到num_evals_to_wait时返回False停止训练注意epsilon的存在意味着指标必须显著超过 1e-3提升才被视为改善避免噪声抖动拖延训练按约定框架首次调用会传入None此时直接返回True继续训练。这三个参数分别对应命令行中的--eval_metric、--early_stopper_num_evals_to_wait、--train_steps_per_eval500在 train.py 中完成接线读者可以按需调整监控指标precision/recall/accuracy与耐心值。集成模型Ensemble训练原文档特别强调train.py的每次调用只训练单个模型副本model replica。论文中使用的集成模型由N 个这样的副本组成每个副本使用不同的随机权重初始化独立训练。因此复现集成时需要为每个副本指定独立的--model_dir避免 checkpoint 互相覆盖对同一架构重复运行 N 次训练命令在推理阶段对 N 个副本的输出proba进行聚合如平均即构成集成预测。推理相关的聚合逻辑可在 util/inference_utils.py 中进一步查看。端到端运行示例基于 train.py 文档字符串中提供的官方示例训练前需先按 aav/model_and_dataset_analysis/README.md 的指引准备 TFRecord 格式的训练/验证数据集# 训练 CNN固定长度编码 python aav/model_training/train.py \ --modelcnn \ --seq_encoderfixedlen-id \ --model_dir/tmp/cnn_model_dir \ --train_path/path/to/datasets/train.tfrecord \ --validation_path/path/to/datasets/valid.tfrecord \ --alsologtostderr # 训练 RNN变长编码 python aav/model_training/train.py \ --modelrnn \ --seq_encodervarlen-id \ --model_dir/tmp/rnn_model_dir \ --train_path/path/to/datasets/train.tfrecord \ --validation_path/path/to/datasets/valid.tfrecord \ --alsologtostderr # 训练逻辑回归固定长度编码可附加超参覆盖 python aav/model_training/train.py \ --modellogistic \ --seq_encoderfixedlen-id \ --model_dir/tmp/lr_model_dir \ --train_path/path/to/datasets/train.tfrecord \ --validation_path/path/to/datasets/valid.tfrecord \ --hparamslearning_rate0.05 \ --alsologtostderr其中--model与--seq_encoder是必填 flagCNN/LR 必须搭配fixedlen-id、RNN 必须搭配varlen-id否则会在 train.py 中抛出NotImplementedError。训练完成后最终超参配置会以 JSON 形式落盘到--model_dir/hparams.json便于事后追溯实验设置。小结aav/model_training/以极简的架构呈现了一套完整的、基于 TensorFlow 1.3tf.Estimator的蛋白质序列表型预测训练流水线三种模型共享统一的训练入口、输入管道tf.Datasets与早停正则化机制仅通过--model与--seq_encoder两个 flag 切换架构与编码方式固定长度编码58 2 × (28 1)与变长编码的取舍、每 500 步评估一次并基于验证集 precision 的早停策略、以及单命令训练单副本、N 副本构成集成的训练范式共同构成了论文实验可复现的关键。对于希望将序列分类模型应用到突变体筛选任务的读者这份代码仓库提供了一个结构清晰、可逐文件研读的参考实现。赞分享人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载相关推荐基于Google Cloud AI Platform的机器学习模型训练与部署实战基于Google Cloud AI Platform的机器学习模型训练与部署实战 本文将通过一个出租车费用预测模型的案例详细介绍如何将本地开发的TensorF机器学习教程示例工程深度学习数据工程C-Learning 实战指南基于递归分类的目标条件强化学习google-research 源码解析与训练配置C Learning 实战指南基于递归分类的目标条件强化学习google research 源码解析与训练配置 C LearningC Learning人工智能深度学习NLP计算机视觉强化学习Machine-Learning-Interviews实战线性回归与逻辑回归代码解析Machine Learning Interviews实战线性回归与逻辑回归代码解析 在机器学习面试中 线性回归 和 逻辑回归 是两个最基础也是最重要的算法示例工程教程人工智能创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED

相关推荐

把 Agent Skills 的模型调用改到 TaoToken 通道,企业业务流程智能体化先跑哪个最小案例?

把 Agent Skills 的模型调用改到 TaoToken 通道,企业业务流程智能体化先跑哪个最小案例?

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

📅 2026/9/20 1:19:02
OneUptime Monitor Secrets 实战指南:加密密钥的创建、注入与访问控制

OneUptime Monitor Secrets 实战指南:加密密钥的创建、注入与访问控制

OneUptime Monitor Secrets 实战指南:加密密钥的创建、注入与访问控制 【免费下载链接】oneuptime Complete open-source monitoring and observability platform. 项目地址: https://gitcode.com/GitHub_Trending/on/oneuptime Monitor Secrets 是 OneUptim…

📅 2026/9/20 1:19:02
品牌命名实战:从案例提取特征,用检查与打分器评估新名字

品牌命名实战:从案例提取特征,用检查与打分器评估新名字

简介:全球著名品牌的产品命名案例是一份面向品牌策划、市场营销及产品经理的实战学习文档,系统梳理了产品命名在品牌建设中的关键作用与决策逻辑。内容从《定位》中的命名理论切入,深入拆解宏碁、柯达、索尼等企业的更名和起名过程&#xff1…

📅 2026/9/20 1:14:02
MORE NEWS

更多资讯

📰

FanControl 风扇控制静音调校:两档温度曲线,让电脑安静待机、满载即跟

FanControl 风扇控制静音调校:两档温度曲线,让电脑安静待机、满载即跟 【免费下载链接】FanControl.Releases This is the release repository for Fan Control, a highly customizable fan controlling software for Windows. 项目地址: https://gitc…

📰

小学生学C++,有必要先学python吗

小学生学C,完全没有“必须先学Python”的硬性要求,要不要先学Python,核心看孩子的基础能力和最终目标,适配你家四年级孩子的最优选择分两种情况: ✅ 完全可以直接跳过Python,直接学C 如果孩子已经通过之前的…

📰

Page Assist:把本地 AI 装进浏览器侧边栏的 5 分钟指南

Page Assist:把本地 AI 装进浏览器侧边栏的 5 分钟指南 【免费下载链接】page-assist Use your locally running AI models to assist you in your web browsing 项目地址: https://gitcode.com/GitHub_Trending/pa/page-assist Page Assist 是一款开源的本地…

📰

Azure MCP 内置进 Visual Studio 2026,Copilot Chat 的模型通道改走 TaoToken 行不行?

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

📰

CANN Runtime Stream管理接口详解:创建、同步、销毁与高级特性

CANNAscend人工智能任务调度 【免费下载链接】runtime 本项目提供CANN运行时组件和维测功能组件。 项目地址: https://gitcode.com/cann/runtime 点击查看 免费下载 CANN(Compute Architecture for Neural Networks)Runtime 的 Stream&#…

📰

Hugging Face 热门权重:Qwen3.7 Flash 接到 TaoToken 后跑长文档问答

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

TODAY

今日更新

THIS WEEK

本周精选

THIS MONTH

本月热门

读完文章,想聊聊您的网站?

告诉我们您的行业与需求,资深顾问一对一梳理方案与报价,全程免费。

📞 💬