CNN、RNN与Transformer架构对比:从时序建模到自注意力机制的技术演进

1次阅读
没有评论

共计 2635 个字符,预计需要花费 7 分钟才能阅读完成。

image.webp

技术背景

1. CNN 的局部感知特性

卷积神经网络(CNN)通过卷积核在输入数据上滑动,提取局部特征。这种局部感知特性使得 CNN 在处理图像等网格数据时非常高效。

CNN、RNN 与 Transformer 架构对比:从时序建模到自注意力机制的技术演进

  • 特征图可视化:CNN 的每一层都会生成多个特征图,每个特征图对应一个卷积核的响应。这些特征图展示了网络如何从低级特征(如边缘)逐步组合成高级特征(如物体部分)。
  • 局部连接:每个神经元只与输入数据的一个局部区域连接,大大减少了参数数量。

2. RNN 的长期依赖困境

循环神经网络(RNN)通过隐藏状态传递历史信息,理论上可以处理任意长度的序列。然而,RNN 在实践中面临长期依赖问题。

  • 梯度传播公式:RNN 的梯度通过时间反向传播(BPTT),在长序列中容易出现梯度消失或爆炸。
  • LSTM/GRU:通过门控机制缓解梯度消失,但仍无法完全解决长序列依赖问题。

3. Transformer 的全局建模能力

Transformer 通过自注意力机制(Self-Attention)实现了对序列数据的全局建模。

  • 注意力权重矩阵:展示了输入序列中每个位置对其他位置的关注程度,这种全局交互能力是 RNN 和 CNN 所不具备的。
  • 并行计算:自注意力机制可以并行计算,大大提高了训练效率。

对比实验

1. PyTorch 实现与 MNIST 时序分类

我们使用 PyTorch 实现了 CNN、RNN(LSTM)和 Transformer 在 MNIST 时序分类任务上的对比。

import torch
import torch.nn as nn

# CNN 实现
class CNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(1, 32, kernel_size=3, stride=1, padding=1)
        self.pool = nn.MaxPool2d(2, 2)
        self.fc = nn.Linear(32 * 14 * 14, 10)

    def forward(self, x):
        x = self.pool(torch.relu(self.conv1(x)))
        x = x.view(-1, 32 * 14 * 14)
        x = self.fc(x)
        return x

# LSTM 实现(带梯度裁剪)class LSTM(nn.Module):
    def __init__(self):
        super().__init__()
        self.lstm = nn.LSTM(28, 128, num_layers=2, batch_first=True)
        self.fc = nn.Linear(128, 10)

    def forward(self, x):
        x = x.squeeze(1)  # (batch, 28, 28)
        output, _ = self.lstm(x)
        output = output[:, -1, :]
        output = self.fc(output)
        return output

2. Benchmark 数据

在 V100 32GB 显存环境下测试:

  • CNN:训练时间最短,显存占用最低。
  • LSTM:需要梯度裁剪(阈值设置为 1.0),显存占用中等。
  • Transformer:训练时间最长,显存占用最高,但准确率最优。

生产实践

1. CNN-RNN 混合架构

结合 CNN 的局部特征提取和 RNN 的时序建模能力,常用于视频分析等任务。

class CNN_LSTM(nn.Module):
    def __init__(self):
        super().__init__()
        self.cnn = nn.Sequential(nn.Conv2d(3, 32, kernel_size=3),
            nn.ReLU(),
            nn.MaxPool2d(2)
        )
        self.lstm = nn.LSTM(32 * 13 * 13, 128, batch_first=True)
        self.fc = nn.Linear(128, 10)

    def forward(self, x):
        batch, seq = x.shape[:2]
        x = x.view(batch * seq, *x.shape[2:])
        x = self.cnn(x)
        x = x.view(batch, seq, -1)
        x, _ = self.lstm(x)
        x = x[:, -1, :]
        x = self.fc(x)
        return x

2. Transformer 的 KV 缓存

Transformer 在推理时需缓存 Key 和 Value 矩阵,导致内存占用随序列长度线性增长。解决方案:

  • 分块处理:将长序列分块,逐块计算注意力。
  • 内存优化:使用内存高效的注意力实现,如 FlashAttention。

3. 多 GPU 训练

数据并行注意事项:

  • 确保 batch size 能被 GPU 数量整除。
  • 使用 torch.nn.DataParallelDistributedDataParallel
  • 注意同步 BatchNorm 层的统计量。

避坑指南

1. RNN 变长序列处理

使用填充掩码(Padding Mask)忽略无效位置:

lengths = [len(seq) for seq in sequences]
max_len = max(lengths)
padded = torch.zeros(len(sequences), max_len)
for i, seq in enumerate(sequences):
    padded[i, :lengths[i]] = torch.tensor(seq[:lengths[i]])

# 创建掩码
mask = torch.arange(max_len).expand(len(lengths), max_len) < torch.tensor(lengths).unsqueeze(1)

2. Transformer 位置编码

使用正弦位置编码时,注意数值稳定性:

  • 将位置编码缩放到与词嵌入相近的范围。
  • 混合使用可学习的位置嵌入(Learned Positional Embedding)。

3. CNN 深度可分离卷积

深度可分离卷积(Depthwise Separable Convolution)的通道数设计原则:

  • 深度卷积(Depthwise Conv)的输出通道数等于输入通道数。
  • 逐点卷积(Pointwise Conv)的输出通道数根据任务需求设定。

总结与思考

本文对比了 CNN、RNN 和 Transformer 的核心特性及适用场景。在实际应用中,应根据任务需求选择合适的架构或混合架构。

开放性问题:当输入序列超过 10 万 token 时,如何改进现有架构?可能的思路包括:

  1. 稀疏注意力(Sparse Attention)
  2. 局部敏感哈希(LSH)注意力
  3. 分层次处理(Hierarchical Processing)

期待读者在实践中探索更多解决方案。

正文完
 0
评论(没有评论)