深度学习入门:从CNN、RNN到Transformer的核心原理与实战对比

1次阅读
没有评论

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

image.webp

背景痛点:传统机器学习的局限

传统机器学习方法(如 SVM、随机森林)在处理图像和序列数据时会遇到明显瓶颈:

深度学习入门:从 CNN、RNN 到 Transformer 的核心原理与实战对比

  • 图像数据:像素级特征提取导致维度灾难,手工设计特征(如 SIFT、HOG)难以适应复杂场景
  • 序列数据:时间步之间的动态依赖关系无法被静态模型捕获,滑动窗口方法丢失长程上下文

这促使了 CNN、RNN 和 Transformer 等深度学习架构的发展,它们通过不同的方式实现了自动特征学习和上下文建模。

技术原理对比

1. CNN:局部感知的视觉专家

卷积神经网络 (Convolutional Neural Network) 的核心设计:

  • 局部感受野:3×3/5×5 卷积核只扫描局部区域(vs 全连接层的全局连接)
  • 参数共享:同一卷积核在图像不同位置复用,显著减少参数量
  • 层级抽象:通过多个卷积层逐步组合低阶边缘→纹理→物体部件

经典 LeNet- 5 结构示例:

Conv1(1->6,k=5)→AvgPool→Conv2(6->16,k=5)→AvgPool→FC120→FC84→Softmax

2. RNN:时序记忆的传承者

循环神经网络 (Recurrent Neural Network) 的特点:

  • 隐状态传递:$h_t = f(W_{xh}x_t + W_{hh}h_{t-1} + b)$
  • 梯度消失:长序列训练时梯度连乘导致早期步信息丢失(LSTM/GRU 通过门控缓解)

LSTM 的三大门控机制:

\begin{aligned}
f_t &= \sigma(W_f\cdot[h_{t-1},x_t]+b_f) \\
i_t &= \sigma(W_i\cdot[h_{t-1},x_t]+b_i) \\
o_t &= \sigma(W_o\cdot[h_{t-1},x_t]+b_o)
\end{aligned}

3. Transformer:全局交互的革命者

自注意力 (Self-Attention) 的核心计算:

  1. 将输入映射为 Q(Query)、K(Key)、V(Value)矩阵
  2. 计算注意力权重:$Attention(Q,K,V)=softmax(\frac{QK^T}{\sqrt{d_k}})V$
  3. 多头注意力并行计算不同子空间的语义关系

位置编码 (Positional Encoding) 示例:

def positional_encoding(pos, d_model):
    angle = pos / (10000 ** (2*(i//2)/d_model))
    return sin(angle) if i%2==0 else cos(angle)

实战代码对比

CNN 实现 MNIST 分类

import torch.nn as nn

class CNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(1, 32, 3)  # 输入通道 1, 输出 32 通道,3x3 卷积
        self.pool = nn.MaxPool2d(2, 2)
        self.fc = nn.Linear(32 * 12 * 12, 10)  # MNIST 最终 10 分类

    def forward(self, x):
        x = self.pool(F.relu(self.conv1(x)))  # 卷积→激活→池化
        x = x.view(-1, 32 * 12 * 12)  # 展平特征图
        return self.fc(x)

LSTM 股票预测

# 滑动窗口构造时序样本
def create_dataset(data, window=5):
    X, y = [], []
    for i in range(len(data)-window):
        X.append(data[i:i+window])
        y.append(data[i+window])
    return np.array(X), np.array(y)

# LSTM 模型定义
class LSTM(nn.Module):
    def __init__(self):
        super().__init__()
        self.lstm = nn.LSTM(input_size=1, hidden_size=50)
        self.linear = nn.Linear(50, 1)

Transformer 文本分类

from transformers import BertModel, BertTokenizer

model = BertModel.from_pretrained('bert-base-uncased')
tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')

inputs = tokenizer("Hello world!", return_tensors="pt")
outputs = model(**inputs)  # 自动处理 positional encoding

生产环境考量

指标 CNN RNN Transformer
FLOPs 中等
内存占用 中等
并行能力 极高
长序列处理 不支持 中等 优秀

小样本数据增强策略:

  • 图像:随机裁剪(Random Crop)、颜色抖动(Color Jitter)
  • 文本:同义词替换(Synonym Replacement)、回译(Back Translation)
  • 时序:窗口切片(Window Slicing)、添加噪声(Add Noise)

常见陷阱与解决方案

  1. CNN 学习率设置
  2. 大卷积核 (7×7) 需要更小的学习率(如 1e-4)
  3. 小卷积核 (3×3) 可尝试较大学习率(如 1e-3)

  4. RNN 序列 Padding

  5. 使用 pack_padded_sequence 跳过无效计算

    from torch.nn.utils.rnn import pack_padded_sequence
    packed = pack_padded_sequence(padded, lengths, batch_first=True)

  6. Transformer 复杂度优化

  7. 使用稀疏注意力(Sparse Attention)
  8. 采用分块计算(Blockwise Computation)
  9. 蒸馏为小型模型(Knowledge Distillation)

延伸阅读

在实际项目中,建议先用 CNN 处理图像任务,RNN 处理短序列任务,Transformer 处理需要长程依赖的场景。随着对模型理解的深入,可以尝试混合架构(如 CNN+Transformer)以获得更好效果。

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