深度学习八大核心神经网络:从入门到精通(CNN/RNN/LSTM/GAN/GNN/DQN/Transformer)

一、开篇总览:神经网络演进与核心定位

深度学习的核心是用多层神经网络自动学习数据特征,八大主流架构各司其职,覆盖视觉、序列、生成、图结构、强化学习、大模型六大核心场景,是从入门到进阶的必经之路。本文从核心原理→结构拆解→优缺点→实战场景→极简代码,一站式吃透八大网络,零基础也能看懂!

八大网络核心定位速览

网络 核心能力 处理数据 标志性应用
CNN 空间特征提取 图像/网格数据 图像分类、目标检测
RNN 序列记忆(短期) 文本/时间序列 简单文本分类、时序预测
LSTM 序列记忆(长期) 长文本/语音 机器翻译、语音识别
Transformer 全局序列建模 文本/多模态 BERT、GPT、大模型
GAN 数据生成 图像/文本 图像生成、换脸、超分辨率
GNN 图结构建模 社交网络/分子结构 推荐系统、药物研发
DQN 强化学习决策 游戏/交互场景 Atari游戏、自动驾驶决策

二、卷积神经网络(CNN):视觉领域“扛把子”

1. 核心原理

CNN(Convolutional Neural Network)专为**网格数据(图像)**设计,核心思想:局部感知+权值共享,模拟人类视觉“先看边缘→再看纹理→最后识别整体”的逻辑。

  • 局部感知:不用全图连接,只关注局部区域(减少参数)
  • 权值共享:同一卷积核在全图滑动,提取相同特征(如边缘)

2. 核心结构拆解

(1)卷积层(Conv Layer)
  • 作用:用**卷积核(滤波器)**滑动提取特征(边缘、纹理、形状)
  • 关键参数:
    • 卷积核大小:3×3(最常用)、5×5
    • 步长(Stride):滑动距离,步长越大输出越小
    • 填充(Padding):Same(输出同尺寸)、Valid(无填充)
  • 计算公式:Output(i,j)=∑m∑nInput(i+m,j+n)⋅Kernel(m,n)Output(i,j) = \sum_m\sum_n Input(i+m,j+n)·Kernel(m,n)Output(i,j)=mnInput(i+m,j+n)Kernel(m,n)
(2)池化层(Pooling Layer)
  • 作用:降维+保留关键特征,减少计算量,增强平移不变性
  • 常用类型:
    • 最大池化(Max Pooling):取局部最大值(保留纹理)
    • 平均池化(Avg Pooling):取局部平均值(保留背景)
(3)全连接层(FC Layer)
  • 作用:将卷积/池化提取的特征整合映射到输出维度(如分类10个类别)

3. 经典架构

  • LeNet(1998):首个成功CNN,手写数字识别(MNIST)
  • VGG(2014):3×3卷积堆叠,结构简洁
  • ResNet(2015):残差连接解决深层网络梯度消失,深度可达1000+层

4. 优缺点

  • ✅ 优点:参数少、计算高效、抗平移、特征提取能力强
  • ❌ 缺点:全局依赖弱(只能看局部)、不擅长序列数据

5. 实战场景

图像分类、目标检测(YOLO)、图像分割(U-Net)、人脸识别、医学影像分析

6. 极简PyTorch代码(MNIST分类)

import torch
import torch.nn as nn

class SimpleCNN(nn.Module):
    def __init__(self):
        super().__init__()
        # 卷积层+激活+池化
        self.conv1=nn.Conv2d(1, 16, 3, padding=1)  # 输入1通道,输出16通道
        self.relu=nn.ReLU()
        self.pool=nn.MaxPool2d(2, 2)  # 池化降维
        self.conv2=nn.Conv2d(16, 32, 3, padding=1)
        # 全连接层
        self.fc1=nn.Linear(32*7*7, 10)  # MNIST输入28×28,池化后7×7

    def forward(self, x):
        # 前向传播
        x=self.pool(self.relu(self.conv1(x)))
        x=self.pool(self.relu(self.conv2(x)))
        x=x.flatten(1)  # 展平特征
        x=self.fc1(x)
        return x

三、循环神经网络(RNN):序列数据“记忆大师”

1. 核心原理

RNN(Recurrent Neural Network)专为序列数据(有先后顺序)设计,核心是循环结构+隐藏状态传递,处理当前输入时结合历史信息,像“带记忆的细胞”。

  • 序列数据:文本、语音、股票、传感器数据
  • 隐藏状态hth_tht:存储历史信息,随时间步传递

2. 结构与计算

(1)基础结构

每个时间步ttt的输入xtx_txt,结合上一步隐藏状态ht−1h_{t-1}ht1,计算当前隐藏状态hth_tht和输出yty_tyt
ht=tanh⁡(Wxhxt+Whhht−1+bh)yt=Whyht+by \begin{align*} h_t &= \tanh(W_{xh}x_t+W_{hh}h_{t-1}+b_h) \\ y_t &= W_{hy}h_t+b_y \end{align*} htyt=tanh(Wxhxt+Whhht1+bh)=Whyht+by

请添加图片描述

3. 致命痛点:梯度消失/爆炸

处理长序列(如长文本)时,梯度随时间步传递指数衰减或爆炸,导致早期信息丢失(记不住开头),仅能捕捉短期依赖。

4. 优缺点

  • ✅ 优点:能处理任意长度序列、结构简单、适合短期依赖
  • ❌ 缺点:长序列记忆差、无法并行计算(必须按顺序)

5. 实战场景

简单文本分类、短期时序预测、语音片段识别

6. 极简PyTorch代码

import torch
import torch.nn as nn

class SimpleRNN(nn.Module):
    def __init__(self, input_dim, hidden_dim, output_dim):
        super().__init__()
        self.rnn=nn.RNN(input_dim, hidden_dim, batch_first=True)  # 批量优先
        self.fc=nn.Linear(hidden_dim, output_dim)

    def forward(self, x):
        # x形状:(batch_size, seq_len, input_dim)
        _, h_n=self.rnn(x)  # h_n:最后时间步隐藏状态
        out=self.fc(h_n.squeeze(0))
        return out

四、长短期记忆网络(LSTM):RNN“进化版”,解决长期依赖

1. 核心痛点

RNN长序列梯度消失,记不住早期信息,LSTM(Long Short-Term Memory)1997年提出,用门控机制+细胞状态解决该问题。

2. 核心结构:三大门+细胞状态

(1)细胞状态CtC_tCt(核心记忆线)
  • 作用:长期信息载体,类似“传送带”,信息直接流过,梯度不易消失
  • 特点:线性传递,修改少,长期记忆稳定
(2)三大门控(控制信息流动)
  • 遗忘门(Forget Gate):决定丢弃上一步细胞状态的信息(如无用历史)
    ft=σ(Wxfxt+Whfht−1+bf)f_t = \sigma(W_{xf}x_t+W_{hf}h_{t-1}+b_f)ft=σ(Wxfxt+Whfht1+bf)
  • 输入门(Input Gate):决定更新当前细胞状态的信息(如新信息)
    it=σ(Wxixt+Whiht−1+bi),C~t=tanh⁡(Wxcxt+Whcht−1+bc)i_t = \sigma(W_{xi}x_t+W_{hi}h_{t-1}+b_i), \quad \tilde{C}_t = \tanh(W_{xc}x_t+W_{hc}h_{t-1}+b_c)it=σ(Wxixt+Whiht1+bi),C~t=tanh(Wxcxt+Whcht1+bc)
  • 输出门(Output Gate):决定输出细胞状态的信息到隐藏状态
    ot=σ(Wxoxt+Whoht−1+bo),ht=ot⋅tanh⁡(Ct)o_t = \sigma(W_{xo}x_t+W_{ho}h_{t-1}+b_o), \quad h_t=o_t · \tanh(C_t)ot=σ(Wxoxt+Whoht1+bo),ht=ottanh(Ct)

请添加图片描述

3. 简化变体:GRU

GRU(Gated Recurrent Unit)合并LSTM三大门为更新门+重置门,参数更少、计算更快,性能接近LSTM。

4. 优缺点

  • ✅ 优点:解决长期依赖、记忆能力强、适合长序列
  • ❌ 缺点:结构复杂、参数多、计算慢、无法并行

5. 实战场景

机器翻译、语音识别、长文本情感分析、时间序列预测(股价、气温)

6. 极简PyTorch代码

import torch
import torch.nn as nn

class SimpleLSTM(nn.Module):
    def __init__(self, input_dim, hidden_dim, output_dim):
        super().__init__()
        self.lstm=nn.LSTM(input_dim, hidden_dim, batch_first=True)
        self.fc=nn.Linear(hidden_dim, output_dim)

    def forward(self, x):
        # x形状:(batch_size, seq_len, input_dim)
        _, (h_n, _)=self.lstm(x)  # h_n:最后时间步隐藏状态
        out=self.fc(h_n.squeeze(0))
        return out

五、Transformer:大模型基石,全局序列建模之王

1. 革命性突破

2017年Google提出,完全抛弃RNN/LSTM循环结构,用**自注意力机制(Self-Attention)**实现:

  • 并行计算:所有时间步同时处理(比RNN快N倍)
  • 全局依赖:直接关联任意两个位置(无论距离多远)
  • 多头注意力:从不同角度捕捉特征(语法、语义、语序)

请添加图片描述

2. 核心组件

(1)自注意力机制(Self-Attention)

计算每个输入与所有输入的关联权重,动态聚合信息,公式:
Attention(Q,K,V)=softmax(QKTdk)VAttention(Q,K,V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)VAttention(Q,K,V)=softmax(dk QKT)V

  • Q(Query):查询向量,K(Key):键向量,V(Value):值向量
  • dk\sqrt{d_k}dk :缩放因子,防止维度高时权重过大
(2)多头注意力(Multi-Head Attention)

并行计算多个注意力头,拼接结果,捕捉多维度依赖(如语法、语义)。

(3)位置编码(Positional Encoding)

Transformer无循环结构,无法感知顺序,用位置编码注入序列顺序信息。

(4)编码器+解码器
  • 编码器:多层堆叠,处理输入序列(如BERT只用编码器)
  • 解码器:多层堆叠,生成输出序列(如GPT只用解码器)

3. 优缺点

  • ✅ 优点:全局依赖、并行计算、长序列友好、多任务适配
  • ❌ 缺点:计算量大、参数多、小数据集易过拟合

4. 实战场景

NLP(BERT、GPT、T5)、多模态(ViT、CLIP)、语音、代码生成、大模型基座

5. 极简PyTorch代码(自注意力)

import torch
import torch.nn as nn

class SelfAttention(nn.Module):
    def __init__(self, embed_dim, num_heads):
        super().__init__()
        self.embed_dim=embed_dim
        self.num_heads=num_heads
        self.head_dim=embed_dim//num_heads
        assert self.head_dim*num_heads == embed_dim, "Embedding维度必须能被头数整除"
        
        # Q/K/V投影层
        self.q_proj=nn.Linear(embed_dim, embed_dim)
        self.k_proj=nn.Linear(embed_dim, embed_dim)
        self.v_proj=nn.Linear(embed_dim, embed_dim)
        self.out_proj=nn.Linear(embed_dim, embed_dim)

    def forward(self, x):
        batch_size, seq_len, _=x.shape
        # 投影并分头
        Q=self.q_proj(x).view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1,2)
        K=self.k_proj(x).view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1,2)
        V=self.v_proj(x).view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1,2)
        
        # 自注意力计算
        attn_scores=torch.matmul(Q, K.transpose(-2,-1))/torch.sqrt(torch.tensor(self.head_dim, dtype=torch.float32))
        attn_weights=torch.softmax(attn_scores, dim=-1)
        attn_out=torch.matmul(attn_weights, V)
        
        # 拼接多头输出
        attn_out=attn_out.transpose(1,2).contiguous().view(batch_size, seq_len, self.embed_dim)
        out=self.out_proj(attn_out)
        return out

六、生成对抗网络(GAN):数据生成神器,“以假乱真”

1. 核心原理

2014年提出,对抗训练:生成器(Generator)vs 判别器(Discriminator),零和博弈,最终生成器能生成以假乱真的数据

  • 生成器G:输入随机噪声,生成伪造数据(如假图)
  • 判别器D:判断输入是真实数据还是伪造数据(二分类)
  • 目标:G尽量骗过D,D尽量区分真假,互相提升

2. 核心结构

(1)生成器(G)
  • 输入:随机噪声zzz
  • 输出:伪造数据G(z)G(z)G(z)(如生成256×256图像)
  • 常用结构:反卷积(转置卷积)、上采样
(2)判别器(D)
  • 输入:真实数据xxx或伪造数据G(z)G(z)G(z)
  • 输出:概率值(0=假,1=真)
  • 常用结构:CNN(图像)、全连接(低维数据)

请添加图片描述

3. 损失函数(对抗损失)

min⁡Gmax⁡DV(D,G)=Ex∼pdata[log⁡D(x)]+Ez∼pz[log⁡(1−D(G(z)))]\min_G\max_D V(D,G) = \mathbb{E}_{x\sim p_{data}}[\log D(x)] + \mathbb{E}_{z\sim p_z}[\log(1-D(G(z)))]GminDmaxV(D,G)=Expdata[logD(x)]+Ezpz[log(1D(G(z)))]

  • D最大化:真实数据概率高,伪造数据概率低
  • G最小化:让D认为伪造数据是真实数据(D(G(z))→1D(G(z))→1D(G(z))1

4. 经典变体

  • DCGAN:CNN+反卷积,生成高清图像
  • StyleGAN:控制生成风格(人脸年龄、性别)
  • CycleGAN:无配对数据转换(照片→油画、马→斑马)

5. 优缺点

  • ✅ 优点:无需标注、生成效果强、数据增强神器
  • ❌ 缺点:训练不稳定、模式崩溃(生成单一数据)、梯度消失

6. 实战场景

图像生成、超分辨率、换脸、风格迁移、文本生成、数据增强

7. 极简PyTorch代码(GAN)

import torch
import torch.nn as nn

# 生成器
class Generator(nn.Module):
    def __init__(self, latent_dim, img_dim):
        super().__init__()
        self.model=nn.Sequential(
            nn.Linear(latent_dim, 128),
            nn.ReLU(),
            nn.Linear(128, 256),
            nn.ReLU(),
            nn.Linear(256, img_dim),
            nn.Tanh()  # 输出归一化到[-1,1]
        )

    def forward(self, z):
        return self.model(z)

# 判别器
class Discriminator(nn.Module):
    def __init__(self, img_dim):
        super().__init__()
        self.model=nn.Sequential(
            nn.Linear(img_dim, 256),
            nn.ReLU(),
            nn.Linear(256, 128),
            nn.ReLU(),
            nn.Linear(128, 1),
            nn.Sigmoid()  # 输出概率[0,1]
        )

    def forward(self, x):
        return self.model(x)

七、图神经网络(GNN):图结构数据建模专家

1. 核心原理

GNN(Graph Neural Network)专为图结构数据(节点+边)设计,核心是消息传递(Message Passing):节点聚合邻居节点特征,迭代更新自身特征,捕捉图的拓扑关系。

  • 图数据:社交网络(用户=节点,好友=边)、分子结构、知识图谱、推荐系统

2. 核心组件

(1)节点特征hvh_vhv

每个节点vvv的初始特征(如用户属性、原子类型)

(2)消息传递函数

节点vvv聚合邻居uuu的特征,更新自身特征:
hv(k)=AGG({hu(k−1),u∈N(v)},hv(k−1))h_v^{(k)} = \text{AGG}\left(\{h_u^{(k-1)}, u \in \mathcal{N}(v)\}, h_v^{(k-1)}\right)hv(k)=AGG({hu(k1),uN(v)},hv(k1))

  • N(v)\mathcal{N}(v)N(v):节点vvv的邻居集合
  • AGG:聚合函数(均值、求和、最大池化)
(3)经典变体
  • GCN(图卷积网络):用卷积聚合邻居,最常用
  • GAT(图注意力网络):用注意力分配邻居权重
  • GraphSAGE:归纳式GNN,能处理新节点(GCN是直推式)

请添加图片描述

3. 优缺点

  • ✅ 优点:捕捉拓扑关系、适配非欧数据、节点/边/图级任务
  • ❌ 缺点:计算复杂度高、大规模图需采样、动态图适配难

4. 实战场景

推荐系统(用户-物品图)、药物研发(分子结构)、社交网络分析、知识图谱推理、欺诈检测

5. 极简PyTorch代码(GCN)

import torch
import torch.nn as nn
import torch.nn.functional as F

class GCNLayer(nn.Module):
    def __init__(self, in_dim, out_dim):
        super().__init__()
        self.linear=nn.Linear(in_dim, out_dim)

    def forward(self, x, adj):
        # x:节点特征 (num_nodes, in_dim)
        # adj:邻接矩阵 (num_nodes, num_nodes)
        x=self.linear(x)
        x=torch.matmul(adj, x)  # 聚合邻居
        x=F.relu(x)
        return x

class SimpleGCN(nn.Module):
    def __init__(self, in_dim, hidden_dim, out_dim):
        super().__init__()
        self.gcn1=GCNLayer(in_dim, hidden_dim)
        self.gcn2=GCNLayer(hidden_dim, out_dim)

    def forward(self, x, adj):
        x=self.gcn1(x, adj)
        x=self.gcn2(x, adj)
        return x

八、深度强化学习(DQN):决策智能体,游戏/自动驾驶核心

1. 核心原理

DQN(Deep Q-Network)是强化学习+CNN的结合,解决高维状态空间(如游戏像素画面)的决策问题,核心是用神经网络拟合Q值函数(状态-动作价值)。

  • 强化学习三要素:智能体(Agent)、环境(Environment)、奖励(Reward)
  • Q值Q(s,a)Q(s,a)Q(s,a):在状态sss下执行动作aaa长期累积奖励

2. 核心结构

(1)Q网络(CNN)
  • 输入:环境状态(如游戏画面像素)
  • 输出:每个动作的Q值(如上下左右、攻击)
  • 结构:CNN提取图像特征→全连接层输出Q值
(2)经验回放(Replay Buffer)
  • 存储历史经验(s,a,r,s′s,a,r,s's,a,r,s),随机采样训练,打破数据相关性,稳定训练
(3)目标网络(Target Network)
  • 延迟更新的Q网络,计算目标Q值,稳定训练目标,防止梯度爆炸

请添加图片描述

3. 核心算法流程

  1. 初始化Q网络和目标网络(权重相同)
  2. 智能体与环境交互,收集经验存入回放缓冲区
  3. 随机采样批量经验,计算当前Q值和目标Q值
  4. 最小化损失(均方误差),更新Q网络
  5. 定期同步目标网络权重

4. 优缺点

  • ✅ 优点:端到端决策、适配高维状态、游戏/自动驾驶主流
  • ❌ 缺点:训练不稳定、过估计偏差、探索与利用平衡难

5. 实战场景

Atari游戏(打砖块、太空侵略者)、自动驾驶决策、机器人控制、推荐系统(序列决策)

6. 极简PyTorch代码(DQN)

import torch
import torch.nn as nn
import torch.optim as optim
import random
import numpy as np

# DQN网络
class DQN(nn.Module):
    def __init__(self, input_shape, num_actions):
        super().__init__()
        self.conv1=nn.Conv2d(input_shape[0], 32, 8, 4)
        self.relu=nn.ReLU()
        self.conv2=nn.Conv2d(32, 64, 4, 2)
        self.conv3=nn.Conv2d(64, 64, 3, 1)
        self.fc1=nn.Linear(64*7*7, 512)
        self.fc2=nn.Linear(512, num_actions)

    def forward(self, x):
        x=self.relu(self.conv1(x))
        x=self.relu(self.conv2(x))
        x=self.relu(self.conv3(x))
        x=x.flatten(1)
        x=self.relu(self.fc1(x))
        x=self.fc2(x)
        return x

# 经验回放
class ReplayBuffer:
    def __init__(self, capacity):
        self.buffer=[]
        self.capacity=capacity

    def add(self, experience):
        if len(self.buffer) >= self.capacity:
            self.buffer.pop(0)
        self.buffer.append(experience)

    def sample(self, batch_size):
        batch=random.sample(self.buffer, batch_size)
        states, actions, rewards, next_states, dones=zip(*batch)
        return (np.array(states), np.array(actions), np.array(rewards),
                np.array(next_states), np.array(dones))

九、八大网络核心对比与学习路径

1. 核心对比表

网络 数据类型 核心能力 计算方式 记忆能力 典型应用
CNN 图像/网格 空间特征提取 并行(局部) 无记忆 图像分类、检测
RNN 序列 短期序列记忆 串行 短期 简单时序预测
LSTM 长序列 长期序列记忆 串行 长期 机器翻译、语音
Transformer 序列/多模态 全局序列建模 并行(全局) 全局 BERT、GPT、大模型
GAN 任意数据 数据生成 对抗训练 无记忆 图像生成、风格迁移
GNN 图结构 拓扑关系建模 消息传递 节点记忆 推荐、药物研发
DQN 图像/状态 强化学习决策 试错学习 经验回放 游戏、自动驾驶

2. 从入门到精通学习路径

阶段1:基础(1-2周)
  • 神经网络基础:感知机、激活函数(ReLU/Sigmoid)、损失函数、反向传播
  • 框架入门:PyTorch/TensorFlow基础(张量、自动微分、模型搭建)
阶段2:核心网络(4-6周)
  1. CNN:图像分类实战(MNIST/CIFAR10)、ResNet微调
  2. RNN/LSTM:文本分类、时序预测实战
  3. Transformer:自注意力理解、BERT/GPT基础使用
  4. GAN:DCGAN生成图像、CycleGAN风格迁移
  5. GNN:GCN节点分类、推荐系统实战
  6. DQN:Atari游戏训练、强化学习基础
阶段3:进阶融合(2-4周)
  • 多模型融合:CNN+LSTM(视频分类)、GNN+Transformer(知识图谱)
  • 大模型实战:微调LLaMA、Prompt工程
  • 部署优化:模型压缩、ONNX导出、端侧部署

十、总结与进阶方向

八大神经网络是深度学习的基石:CNN看图像、RNN/LSTM处理序列、Transformer统领大模型、GAN生成数据、GNN建模图结构、DQN做决策。从基础原理到极简代码,本文一站式覆盖核心知识点,帮你快速建立深度学习知识体系。

进阶方向

  • 多模态融合:CLIP、Flux、Sora(文本/图像/视频/音频)
  • 自监督学习:MAE、SimCLR(无需标注,预训练模型)
  • 轻量化模型:MobileNet、SqueezeNet(端侧部署)
  • 前沿架构:扩散模型(Diffusion)、ViT、Mamba(Transformer替代)
Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐