#用 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 copy | mRNA 半衰期相关特征 | [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
###启动子通路
两级卷积 + 批标准化 + 最大池化,逐步提取从局部碱基模式到更大范围的调控特征:
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。
###注意力机制
在卷积特征上引入简单的注意力加权,让模型学会关注更重要的位置:
# 计算注意力权重
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 之间的权重,逐通道地对卷积特征做软性选择。直觉上,某些位置的卷积特征对表达预测更有信息量(比如转录因子结合位点附近),注意力机制帮助模型聚焦这些区域。
###半衰期通路
半衰期特征维度低,只需一层全连接升维:
self.fc_halflife = nn.Linear(8, 32)
###融合与分类
两路特征拼接后通过全连接层:
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 Size | 32 |
| Epochs | 10 |
###数据加载
数据以 HDF5 格式存储,包含 train.h5click to copy、valid.h5click to copy、test.h5click to copy 三个文件:
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 copy 和 DataLoaderclick to copy 构建数据管道:
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 | 整体分类准确率 |
| AUC | ROC 曲线下面积,衡量正负类区分能力 |
| 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 是我们读懂其中某些段落的放大镜。"