之前做图像分类项目时,总有个很直接的体会:单靠 CNN 抓局部纹理和边缘特征很顺手,但一遇到全局上下文、长距离依赖就有些吃力;换成纯 Transformer 后,全局关系能建模了,可小数据集上又容易过拟合,训练速度也明显变慢。后来把 CNN 和 Transformer 串联或并联到一起,再在中间加入特征融合模块,效果一下子好了很多。这个组合近两年在论文里出现频率非常高,也是很多“大小论文”的创新切入点。
这篇文章不打算讲空泛概念,而是把 CNN、Transformer、特征融合这三件事从原理、代码到实验设计完整串起来。文章会给出一个可以直接运行的 PyTorch 示例,用图像分类任务演示“CNN 分支 + Transformer 分支 + 特征融合”的整体流程。无论你是新手入门深度学习,还是已经在做实验、准备写论文,都能从中找到可以直接复用的思路。
1. 背景与核心概念
1.1 为什么要组合 CNN 和 Transformer
先来说说这两个网络各自的定位。
CNN(卷积神经网络)的核心假设是局部性和平移不变性。卷积核在一小块邻域内做加权求和,因此天然擅长提取边缘、纹理、形状等局部特征。参数共享又让 CNN 在图像任务里非常高效,不需要每个位置都单独学一套参数。
Transformer 的核心机制是自注意力。它会把输入序列中的每个元素与其他所有元素计算相关性,因此能直接建模长距离依赖,捕获全局上下文信息。这也是 Transformer 在 NLP 领域成功后,又被大量移植到视觉和时序任务中的原因。
两者单独使用都有短板:
- CNN 的感受野有限,虽然可以通过堆叠层数扩大,但深层特征对全局关系的建模仍然不够直接。
- Transformer 缺少 CNN 那种内在的局部归纳偏置,在小规模数据集上往往需要更多训练数据和更大的模型量级才能收敛到理想效果。
所以一个很自然的思路是:让 CNN 负责局部特征,让 Transformer 负责全局关系,最后再把两种特征融合起来。这样既保留了局部细节,又引入了全局语义。
1.2 特征融合要解决什么问题
单纯把 CNN 和 Transformer 组合在一起,只能算结构上的堆叠。真正让模型变强的是怎么把两个分支的特征融合起来。
特征融合(Feature Fusion)指的是把不同来源、不同尺度、不同语义层级的特征向量组合成一个更完整的表示。常见的融合方式有:
- 拼接(Concat):直接把维度拼接,简单但维度翻倍。
- 相加(Add):要求维度一致,操作轻量,类似残差思想。
- 门控融合(Gating):通过可学习的权重动态调节两个特征的贡献。
- 注意力融合(Attention Fusion):用注意力机制让模型学会从两个特征里“挑”它需要的信息。
在论文里,特征融合往往就是那个“创新点”所在。同样两个主干网络,融合模块设计得好不好,直接决定了实验效果的上限。
1.3 常见应用场景
这个组合的适用范围很广,常见任务包括:
- 图像分类、目标检测、语义分割。
- 视频理解中的时空特征建模。
- 时间序列预测、异常检测。
- 多模态任务,比如文本加图像、文本加表格数据。
- 恶意软件检测、医学图像分析等垂直领域。
实际做项目时,只要数据同时存在“局部模式”和“全局依赖”,这个组合都有尝试价值。
2. 环境准备与版本说明
由于后面的代码使用 PyTorch 编写,我们先准备 Python 环境和依赖。
conda create -n cnn_transformer python=3.9 -y conda activate cnn_transformer pip install torch torchvision tqdm matplotlib版本不需要刻意锁死,重点演示设计思路。本文代码在 PyTorch 2.x、Python 3.9 环境下测试过,如果你使用的是 PyTorch 1.x,注意两个地方:
nn.TransformerEncoderLayer需要设置batch_first=True,这个参数在较新版本中默认就是True。torchvision.datasets.MNIST下载可能需要稳定的网络环境,如果下载失败,可以手动下载后放到./data目录。
建议准备一块支持 CUDA 的显卡,虽然 CPU 也能跑,但训练速度会比较慢。如果只有 CPU,可以把 batch size 调小一点。
3. 核心原理拆解
3.1 CNN 分支:提取局部特征
CNN 分支的常规写法是“卷积 + 归一化 + 激活 + 池化”交替堆叠。
import torch.nn as nn class CNNBranch(nn.Module): def __init__(self, in_channels=1, feat_dim=128): super().__init__() self.features = nn.Sequential( nn.Conv2d(in_channels, 32, kernel_size=3, padding=1), nn.BatchNorm2d(32), nn.ReLU(inplace=True), nn.Conv2d(32, 64, kernel_size=3, padding=1), nn.BatchNorm2d(64), nn.ReLU(inplace=True), nn.AdaptiveAvgPool2d((1, 1)), ) self.fc = nn.Linear(64, feat_dim) def forward(self, x): x = self.features(x) x = x.flatten(1) x = self.fc(x) return x这里的AdaptiveAvgPool2d((1,1))会把任意大小的特征图压成1x1,这样后面的全连接层输入维度就固定了,不用手动计算卷积后的尺寸。
为什么用全局平均池化而不是直接flatten?因为flatten会把空间位置全部展开,特征图大小一变,全连接层的参数就全部失效。全局平均池化把每个通道汇总成一个数值,既降低了参数量,又保留了一定的空间鲁棒性。
3.2 Transformer 分支:建模全局依赖
Transformer 分支通常分两步走:先把图像切块并映射成 embedding,然后送入编码器。
图像切块可以手工reshape,也可以用Conv2d实现。用卷积层的思路很巧妙:一个kernel_size=patch_size, stride=patch_size的卷积,恰好等价于把图像分成不重叠的小块,再对每个小块做线性映射。
class TransformerBranch(nn.Module): def __init__(self, in_channels=1, img_size=28, patch_size=4, d_model=128, nhead=4, num_layers=2): super().__init__() self.patch_size = patch_size self.d_model = d_model num_patches = (img_size // patch_size) ** 2 self.patch_embed = nn.Conv2d(in_channels, d_model, kernel_size=patch_size, stride=patch_size) self.pos_embed = nn.Parameter(torch.zeros(1, num_patches, d_model)) encoder_layer = nn.TransformerEncoderLayer(d_model=d_model, nhead=nhead, batch_first=True) self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers) self.norm = nn.LayerNorm(d_model) def forward(self, x): x = self.patch_embed(x) # (B, d_model, ph, pw) x = x.flatten(2).transpose(1, 2) # (B, L, d_model) x = x + self.pos_embed x = self.transformer(x) x = x.mean(dim=1) # 全局平均池化 x = self.norm(x) return xTransformer 本身是“排列不变”的,也就是说它不会天然感知输入的顺序信息。所以位置编码必不可少。这里使用可学习位置编码,即一个形状为(1, num_patches, d_model)的nn.Parameter,让模型在训练过程中自己调整位置表示。
如果你追求更稳定的效果,也可以换成正弦位置编码或被广泛使用的相对位置编码。论文里经常对不同位置编码方式做对比实验,这也是一个很好的“凑实验”方向。
3.3 特征融合:把局部与全局信息整合
最直接也最稳定的融合方式就是拼接。将 CNN 输出的特征与 Transformer 输出的特征在最后一个维度拼接,得到更长的向量,再送入分类头。
class FeatureFusion(nn.Module): def __init__(self, cnn_dim, trans_dim, hidden_dim, num_classes): super().__init__() self.fusion = nn.Sequential( nn.Linear(cnn_dim + trans_dim, hidden_dim), nn.ReLU(inplace=True), nn.Dropout(0.1), nn.Linear(hidden_dim, num_classes), ) def forward(self, cnn_feat, trans_feat): fused = torch.cat([cnn_feat, trans_feat], dim=1) return self.fusion(fused)除了拼接,常见的融合设计还有:
| 融合方式 | 做法 | 优点 | 缺点 |
|---|---|---|---|
| Concat | 直接拼接特征向量 | 简单、稳定、信息不丢失 | 维度变大,计算量增加 |
| Add | 逐元素相加 | 参数少、实现简单 | 要求两个分支维度完全一致,可能互相干扰 |
| Gating | 学两个权重,加权相加 | 可动态调节分支贡献 | 多一层参数,训练稍微复杂 |
| Cross Attention | 一个分支的 token 与另一个分支的 token 做注意力 | 交互充分,特征融合效果好 | 计算量大,容易过拟合 |
如果做论文实验,建议至少对比 Concat 和 Cross Attention 两种方案,再配合消融实验说明每个模块的贡献。
4. 完整实战案例:图像分类中的 CNN + Transformer + 特征融合
这一节给出一个完整的可运行项目。任务选择 MNIST 手写数字分类,因为数据量小、训练速度快,方便你快速跑通。整个流程稍加修改也能用到 CIFAR-10 或你自己的数据集上。
4.1 创建项目结构
cnn_transformer_fusion/ ├── main.py └── data/main.py存放所有代码,data目录存放 MNIST 数据集。
4.2 编写完整模型代码
在main.py中依次写入下面的内容。
import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from torchvision import datasets, transforms # ---------- CNN 分支 ---------- class CNNBranch(nn.Module): def __init__(self, in_channels=1, feat_dim=128): super().__init__() self.features = nn.Sequential( nn.Conv2d(in_channels, 32, kernel_size=3, padding=1), nn.BatchNorm2d(32), nn.ReLU(inplace=True), nn.Conv2d(32, 64, kernel_size=3, padding=1), nn.BatchNorm2d(64), nn.ReLU(inplace=True), nn.AdaptiveAvgPool2d((1, 1)), ) self.fc = nn.Linear(64, feat_dim) def forward(self, x): x = self.features(x) x = x.flatten(1) x = self.fc(x) return x # ---------- Transformer 分支 ---------- class TransformerBranch(nn.Module): def __init__(self, in_channels=1, img_size=28, patch_size=4, d_model=128, nhead=4, num_layers=2): super().__init__() self.patch_size = patch_size self.d_model = d_model num_patches = (img_size // patch_size) ** 2 self.patch_embed = nn.Conv2d(in_channels, d_model, kernel_size=patch_size, stride=patch_size) self.pos_embed = nn.Parameter(torch.zeros(1, num_patches, d_model)) encoder_layer = nn.TransformerEncoderLayer( d_model=d_model, nhead=nhead, batch_first=True ) self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers) self.norm = nn.LayerNorm(d_model) def forward(self, x): x = self.patch_embed(x) x = x.flatten(2).transpose(1, 2) x = x + self.pos_embed x = self.transformer(x) x = x.mean(dim=1) x = self.norm(x) return x # ---------- 特征融合与分类头 ---------- class FeatureFusion(nn.Module): def __init__(self, cnn_dim, trans_dim, hidden_dim, num_classes): super().__init__() self.fusion = nn.Sequential( nn.Linear(cnn_dim + trans_dim, hidden_dim), nn.ReLU(inplace=True), nn.Dropout(0.1), nn.Linear(hidden_dim, num_classes), ) def forward(self, cnn_feat, trans_feat): fused = torch.cat([cnn_feat, trans_feat], dim=1) return self.fusion(fused) # ---------- 整体模型 ---------- class CNNTransformerFusion(nn.Module): def __init__(self, in_channels=1, img_size=28, patch_size=4, feat_dim=128, nhead=4, num_layers=2, hidden_dim=64, num_classes=10): super().__init__() self.cnn_branch = CNNBranch(in_channels=in_channels, feat_dim=feat_dim) self.transformer_branch = TransformerBranch( in_channels=in_channels, img_size=img_size, patch_size=patch_size, d_model=feat_dim, nhead=nhead, num_layers=num_layers, ) self.fusion = FeatureFusion( cnn_dim=feat_dim, trans_dim=feat_dim, hidden_dim=hidden_dim, num_classes=num_classes, ) def forward(self, x): cnn_feat = self.cnn_branch(x) trans_feat = self.transformer_branch(x) out = self.fusion(cnn_feat, trans_feat) return out这里的关键设计是让两个分支输出的特征维度一致,都是feat_dim=128。这样做的好处是融合之后维度对称,后续替换成 Add、Gating 等融合方式时,不需要大规模调整代码。
4.3 数据加载与训练循环
def main(): device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Using device: {device}") transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_dataset = datasets.MNIST(root="./data", train=True, download=True, transform=transform) test_dataset = datasets.MNIST(root="./data", train=False, download=True, transform=transform) train_loader = DataLoader(train_dataset, batch_size=128, shuffle=True, num_workers=2) test_loader = DataLoader(test_dataset, batch_size=256, shuffle=False, num_workers=2) model = CNNTransformerFusion().to(device) criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=1e-3) epochs = 5 for epoch in range(epochs): model.train() total_loss = 0.0 correct = 0 total = 0 for images, labels in train_loader: images, labels = images.to(device), labels.to(device) outputs = model(images) loss = criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() total_loss += loss.item() * images.size(0) _, predicted = outputs.max(1) total += labels.size(0) correct += predicted.eq(labels).sum().item() train_acc = 100.0 * correct / total print(f"Epoch [{epoch+1}/{epochs}] Loss: {total_loss / total:.4f} Acc: {train_acc:.2f}%") # 测试 model.eval() test_correct = 0 test_total = 0 with torch.no_grad(): for images, labels in test_loader: images, labels = images.to(device), labels.to(device) outputs = model(images) _, predicted = outputs.max(1) test_total += labels.size(0) test_correct += predicted.eq(labels).sum().item() test_acc = 100.0 * test_correct / test_total print(f"Test Acc: {test_acc:.2f}%") if __name__ == "__main__": main()4.4 运行与验证
在项目目录下执行:
python main.py如果环境配置正确,你会看到类似下面的输出(具体数值会因随机种子、硬件环境波动):
Using device: cuda Epoch [1/5] Loss: 0.3261 Acc: 90.35% Epoch [2/5] Loss: 0.1612 Acc: 95.67% Epoch [3/5] Loss: 0.1210 Acc: 96.71% Epoch [4/5] Loss: 0.0998 Acc: 97.33% Epoch [5/5] Loss: 0.0847 Acc: 97.68% Test Acc: 97.53%MNIST 本身比较简单,只看准确率可能看不出 CNN + Transformer 的优势。建议把代码迁移到 CIFAR-10 或者你自己的业务数据上,再对比“纯 CNN”“纯 Transformer”“CNN + Transformer + 特征融合”三组实验,差异就会明显得多。
4.5 结果说明
这个示例的价值不在于刷高 MNIST 准确率,而在于给你一个可以继续修改的基线。拿到代码后,你可以做几件事:
- 修改
patch_size,观察 Transformer 分支对输入序列长度的敏感度。 - 修改
num_layers,比较不同深度的 Transformer 对最终效果的影响。 - 把
FeatureFusion中的torch.cat改成x = cnn_feat + trans_feat,对比拼接与相加的效果。 - 在融合之前给每个分支的特征加一个
nn.LayerNorm,有时能稳定训练。
5. 常见问题与排查思路
5.1 维度不匹配
刚写完代码,最容易碰到size mismatch之类的报错。常见原因有三个:
- 图像输入尺寸不是
patch_size的整数倍,导致num_patches计算错误。 - 两个分支输出的特征维度不一致,拼接时对不上。
- 使用了不同尺寸的数据集,但模型里的
img_size没有改。
排查时先打印每个分支输出的shape:
print(cnn_feat.shape, trans_feat.shape)然后根据实际形状去调整全连接层或卷积层参数。
5.2 Transformer 在小型数据集上过拟合
Transformer 的参数数量通常比同等规模的 CNN 多,而且在数据少时更容易过拟合。常见表现是训练准确率很高、测试准确率低。
解决办法包括:
- 增加数据增强,比如随机裁剪、翻转、色彩抖动。
- 在 Transformer 分支里增加
dropout参数。 - 减少
num_layers或nhead。 - 引入预训练权重。
5.3 训练速度很慢
Transformer 的自注意力是平方复杂度,序列长度越长,计算越慢。如果patch_size=4,对于28x28图像,序列长度是7x7=49,问题不大。但如果换成224x224图像,序列长度变成56x56=3136,普通 GPU 都很难跑动。
可以这样优化:
- 增大
patch_size,减少 patch 数。 - 使用 Swin Transformer 等窗口注意力结构。
- 在 CNN 分支提取特征后,只在深层特征上使用 Transformer。
- 使用混合精度训练(
torch.cuda.amp)。
5.4 加了 Transformer 分支后效果反而变差
这种情况并不少见。原因是你的任务可能本身就不需要很强的全局建模能力,或者 Transformer 分支在训练初期不稳定,拖累了整个模型。
建议做消融实验:
- 只保留 CNN 分支。
- 只保留 Transformer 分支。
- 两个分支都保留,但不做融合,直接相加。
- 两个分支都保留,使用拼接融合。
通过这种对比,你能清楚看到哪个模块真正带来了提升。
6. 最佳实践与工程建议
6.1 模型设计建议
不要把两个分支设计得一样深。CNN 可以浅一点,负责底层特征;Transformer 分支放在更高层的语义特征上效果更好。在实际项目中,更合理的结构是“CNN 先降维,Transformer 后建模”。
特征融合模块不要一上来就设计得太复杂。先用最简单的 Concat 跑通基线,再逐步加入 Gate、注意力等机制。复杂模块在小数据集上很容易过拟合。
6.2 训练技巧
- 给两个分支设置不同的学习率。CNN 分支通常收敛快,Transformer 分支可以适当使用更小的学习率。
- 先冻结一个分支,训练另一个分支,然后再两个分支一起微调。这个思路在跨模态任务中很常见。
- 使用余弦退火学习率调度器,比固定学习率更稳。
- 记录训练日志时,除了准确率,还要保存每个分支输出的特征范数,方便观察两个分支是否“不平衡”。
6.3 论文实验设计建议
如果你准备围绕这个组合写论文,核心不是“我拼了两个网络”,而是“我为什么要拼、怎么拼、带来什么收益”。论文里建议至少包含以下实验:
- 基线实验:单独 CNN、单独 Transformer、常见主流模型。
- 消融实验:去掉特征融合、去掉 CNN 分支、去掉 Transformer 分支、替换不同融合方式。
- 参数分析:patch size、序列长度、融合维度对结果的影响。
- 可视化:CNN 特征图、注意力权重、融合后的特征分布。
- 复杂度对比:参数量、FLOPs、推理时间。
这五组实验是论文里最经典的论证框架。做实验时一定要保留随机种子、超参数记录和多次重复结果,保证实验可复现。这也是审稿人非常看重的一点。
6.4 工程落地注意事项
如果要把模型部署到生产环境,还需要考虑:
- 模型量化:CNN 和 Transformer 都能用 INT8 量化,但需要验证精度损失。
- 推理框架:ONNX Runtime 或 TensorRT 对 Transformer 的支持已经很好,但
torch.cat等动态 Shape 操作可能影响优化。 - 数据流:预处理、标准化和模型输入输出需要封装成统一接口,避免训练和推理行为不一致。
7. 总结与学习路线
本文从 CNN 和 Transformer 的原理差异出发,解释了为什么要在两者之间加入特征融合,并用一个完整的 PyTorch 示例演示了“CNN 分支 + Transformer 分支 + 特征融合 + 分类头”的实现方法。你可以直接跑通 MNIST,再迁移到自己的任务上。
接下来可以按下面顺序继续深入学习:
- 先完整阅读
nn.TransformerEncoderLayer源码,理解 Q、K、V 的计算过程。 - 把本文的 Concat 融合换成 Cross Attention,对比效果。
- 把模型迁移到时间序列预测任务中,CNN 用来提取局部波动特征,Transformer 用来建模长周期依赖。
- 阅读 ViT、Swin Transformer、DETR 等经典论文,了解 patch embedding、窗口注意力、多尺度特征在不同任务中的设计细节。
如果你正在准备自己的论文,千万别把“组合模型”当成万能框架。每个数据集、每个任务都有最适合的结构。把消融实验做扎实,把可视化做清晰,比堆叠更多模块更有说服力。希望这套组合能帮你顺利产出高质量的实验和论文。