用 CNN 预测基因表达:从启动子序列到二分类模型
技术#深度学习#CNN

用 CNN 预测基因表达:从启动子序列到二分类模型

>
~5 min read

#用 CNN 预测基因表达:从启动子序列到二分类模型

基因表达调控是分子生物学的核心问题之一。启动子区域的 DNA 序列蕴含了大量调控信息——转录因子结合位点、CpG 岛、核心启动子元件……这些特征直接决定了一个基因是否会被转录。

问题是:给定一段启动子序列,我们能预测它对应的基因是否表达吗?

本文介绍一个基于 CNN + 注意力机制 的二分类模型,输入启动子序列和 mRNA 半衰期特征,输出基因的表达/不表达预测。模型简洁,训练 10 轮后验证集准确率达到 81%

项目地址:github.com/Linmoqian/DNA_CNN_predict

##问题定义

###任务

二分类:给定一个基因的启动子序列和半衰期特征,预测该基因是否表达。

###数据

字段说明形状
promoterclick to copy启动子区域 DNA 序列的 one-hot 编码[20000, 4]click to copy
halflifeclick to copymRNA 半衰期相关特征[8]click to copy
labelclick to copy表达标签(0/1)标量

启动子序列长度 20000 bp,4 个通道对应 A、T、C、G 四种碱基的 one-hot 编码。这是模型的主要输入。

数据来源:AISCCC 数据库

##模型架构

模型采用双路融合设计:一条通路处理启动子序列,另一条通路处理半衰期特征,最后在全连接层合并。

启动子序列 半衰期特征 [batch, 20000, 4] [batch, 8] | | permute 转置 | [batch, 4, 20000] | | | +---------v---------+ | | Conv1d(4→32, k=3) | | | + BatchNorm + ReLU | | +---------+---------+ | | | MaxPool1d(2) | | | +---------v---------+ | | Conv1d(32→64, k=3) | | | + BatchNorm + ReLU | | +---------+---------+ | | | MaxPool1d(2) | | | +---------v---------+ +---------v---------+ | 注意力机制加权 | | Linear(8→32)+ReLU | | sigmoid(FC(64)) | +---------+---------+ +---------+---------+ | | | Flatten | | | [batch, 640000] [batch, 32] | | +----------------+---------------+ | +--------v--------+ | Concatenate | | [batch, 640032] | +--------+--------+ | +--------v--------+ | Linear(→128)+ReLU| +--------+--------+ | Dropout(0.8) | +--------v--------+ | Linear(→2) | +--------+--------+ | 二分类输出 click to copy

###启动子通路

两级卷积 + 批标准化 + 最大池化,逐步提取从局部碱基模式到更大范围的调控特征:

>hljs python.0 lines
self.conv1_promoter = nn.Conv1d(in_channels=4, out_channels=32, kernel_size=3, padding=1)
self.bn1_promoter = nn.BatchNorm1d(32)
self.conv2_promoter = nn.Conv1d(in_channels=32, out_channels=64, kernel_size=3, padding=1)
self.bn2_promoter = nn.BatchNorm1d(64)
self.pool_promoter = nn.MaxPool1d(kernel_size=2, stride=2)

为什么用 kernel_size=3click to copy?因为生物序列中的功能 motif 通常是短片段(如 TATA box 的 TATAAA),小卷积核更适合捕捉这类局部模式。两层池化后,序列长度从 20000 压缩到 5000,通道数从 4 增长到 64。

###注意力机制

在卷积特征上引入简单的注意力加权,让模型学会关注更重要的位置:

>hljs python.0 lines
# 计算注意力权重
attention_weights = torch.sigmoid(self.attention_fc(x_promoter.permute(0, 2, 1)))
# 加权
x_promoter = x_promoter * attention_weights.permute(0, 2, 1)

这里用了一个 Linear(64, 64)click to copy + sigmoidclick to copy 生成 0 到 1 之间的权重,逐通道地对卷积特征做软性选择。直觉上,某些位置的卷积特征对表达预测更有信息量(比如转录因子结合位点附近),注意力机制帮助模型聚焦这些区域。

###半衰期通路

半衰期特征维度低,只需一层全连接升维:

>hljs python.0 lines
self.fc_halflife = nn.Linear(8, 32)

###融合与分类

两路特征拼接后通过全连接层:

>hljs python.0 lines
self.fc1 = nn.Linear(self.promoter_out_size + 32, 128)
self.fc2 = nn.Linear(128, num_classes)  # num_classes = 2
self.dropout = nn.Dropout(0.8)

注意 Dropout 比率设为 0.8,非常高。对于这个规模的数据集,高 Dropout 有助于防止过拟合。

##训练流程

###超参数

参数
优化器Adam
学习率1e-4
权重衰减1e-6
学习率调度ReduceLROnPlateau (patience=3, factor=0.1)
损失函数CrossEntropyLoss
Batch Size32
Epochs10

###数据加载

数据以 HDF5 格式存储,包含 train.h5click to copyvalid.h5click to copytest.h5click to copy 三个文件:

>hljs python.0 lines
def load_hdf5(file_path):
    with h5py.File(file_path, 'r') as f:
        gene_ids = list(f['gene_id'])
        halflife = torch.tensor(np.array(f['halflife']), dtype=torch.float32)
        promoter = torch.tensor(np.array(f['promoter']), dtype=torch.float32)
        labels = torch.tensor(np.array(f['label']), dtype=torch.long)
        return gene_ids, halflife, promoter, labels

通过 PyTorch 的 TensorDatasetclick to copyDataLoaderclick to copy 构建数据管道:

>hljs python.0 lines
train_loader = create_dataloader(train_promoters, train_half_lives, train_labels)
valid_loader = create_dataloader(valid_promoters, valid_half_lives, valid_labels)
test_loader = create_dataloader(test_promoters, test_half_lives, test_labels)

###评估指标

使用三个指标全面评估模型:

指标说明
Accuracy整体分类准确率
AUCROC 曲线下面积,衡量正负类区分能力
F1 Score精确率和召回率的调和平均

##结果

10 轮训练后:

  • 验证集准确率:81%
  • 测试集同步评估 Accuracy、AUC 和 F1

训练过程可视化包含损失曲线和准确率曲线,帮助判断是否过拟合或欠拟合。

##设计思路与讨论

###为什么选择 CNN

DNA 序列本质上是一维离散信号。CNN 的卷积操作天然适合捕捉局部模式——在 DNA 中,这些局部模式就是转录因子结合位点、启动子元件等功能 motif。多层卷积 + 池化可以逐步扩大感受野,从碱基级别提取到更高级的序列特征。

###为什么加入注意力

单纯的 CNN 对所有位置一视同仁。但启动子序列中,真正影响基因表达的只是少数关键区域(核心启动子、增强子元件)。注意力机制让模型学会区分重要和不重要的位置。

###为什么融合半衰期

基因表达水平不仅取决于转录(启动子控制),还取决于 mRNA 的降解速度(半衰期控制)。将半衰期特征作为辅助输入,给模型提供了更完整的生物学上下文。

###可改进方向

方向思路
更深的网络ResNet 残差连接,缓解梯度消失
Transformer替代 CNN,用自注意力捕捉长距离依赖
多尺度卷积并行使用不同 kernel_size,捕捉不同长度的 motif
位置编码让模型感知碱基的绝对位置信息
更丰富的特征加入 CpG 岛、染色质开放性等表观遗传特征

##总结

这个项目展示了如何将深度学习应用于基因组学的一个经典问题。模型不复杂——两级 CNN 加注意力、双路特征融合、标准训练流程——但在 10 轮训练内就达到了 81% 的准确率。

核心收获:

  • DNA 序列的 one-hot 编码可以作为 CNN 的有效输入
  • 注意力机制帮助模型聚焦关键调控区域
  • 多源特征融合(序列 + 半衰期)提升预测能力

这是一个足够简洁的起点,可以在此基础上尝试更复杂的架构和更丰富的生物学特征。

"DNA 是一本用四种字母写成的书,CNN 是我们读懂其中某些段落的放大镜。"

> tags
#深度学习#CNN#基因组学#PyTorch#注意力机制
> related_posts
>cd /blog_