深度学习 深度学习把非线性层堆叠起来,构建层次化的表示,自动把原始输入转换成有用的特征。本文件涵盖 MLP、激活函数、反向传播、CNN、RNN、LSTM、注意力、Transformer、GAN、VAE、扩散模型以及归一化技术。 是什么让一个网络变"深"?浅层网络只有一个隐藏层;深层网络有很多层。深度让网络能够构建层次化的表示:靠前的层学到简单特征(边缘、色调),靠后的层把它们组合成复杂概念(人脸、句子)。正是这种组合性赋予了深度学习力量。 最简单的深度网络是多层感知机(multi-layer perceptron, MLP),也叫全连接网络或稠密网络。
深度学习把非线性层堆叠起来,构建层次化的表示,自动把原始输入转换成有用的特征。本文件涵盖 MLP、激活函数、反向传播、CNN、RNN、LSTM、注意力、Transformer、GAN、VAE、扩散模型以及归一化技术。
是什么让一个网络变"深"?浅层网络只有一个隐藏层;深层网络有很多层。深度让网络能够构建层次化的表示:靠前的层学到简单特征(边缘、色调),靠后的层把它们组合成复杂概念(人脸、句子)。正是这种组合性赋予了深度学习力量。
最简单的深度网络是多层感知机(multi-layer perceptron, MLP),也叫全连接网络或稠密网络。每一层计算:
这里 W 是权重矩阵(第 2 章),b 是偏置向量,\sigma 是非线性激活函数。一层的输出成为下一层的输入。要是没有这个非线性,堆叠层就毫无意义:W_2(W_1 x) = (W_2 W_1)x,不过是又一个线性变换。这正是第 2 章里矩阵乘法的塌缩。
**激活函数(activation function)**引入了让深度有意义的非线性。
ReLU(Rectified Linear Unit,修正线性单元):\text{ReLU}(x) = \max(0, x)。它是用得最广的激活函数。计算快、对正输入不饱和,还能产生稀疏的激活(许多神经元恰好输出零)。缺点是:输入为负的神经元永远输出零,要是它们永久卡在那里,就会"死掉"、停止学习。
Sigmoid:\sigma(x) = \frac{1}{1+e^{-x}},把输入压到 (0, 1)。在二分类的输出层很有用,但在隐藏层里会有麻烦,因为输入远离零时梯度会消失(曲线几乎是平的)。
Tanh:\tanh(x) = \frac{e^x - e^{-x}}{e^x + e^{-x}},压到 (-1, 1)。零中心化(不像 sigmoid),有利于梯度流动,但在两端仍然会受到梯度消失之苦。
GELU(Gaussian Error Linear Unit,高斯误差线性单元):\text{GELU}(x) = x \cdot \Phi(x),其中 \Phi 是标准正态分布的 CDF。它是 ReLU 的平滑近似,允许少量负值通过。GELU 是 GPT 和 BERT 的默认激活。
Swish:\text{Swish}(x) = x \cdot \sigma(x),又一个平滑门控。实践中与 GELU 类似。
一个有 d_{\text{in}} 个输入、d_{\text{out}} 个输出的稠密层有 d_{\text{in}} \times d_{\text{out}} + d_{\text{out}} 个参数(权重加偏置)。矩阵乘 Wx 就是第 2 章的矩阵-向量乘法。在批处理设定下,输入是形状为 (B, d_{\text{in}}) 的矩阵 X,输出是形状为 (B, d_{\text{out}}) 的 XW^T + b。
**通用近似定理(universal approximation theorem)**指出,一个有足够多神经元的单隐藏层网络,可以以任意精度逼近紧致域上的任何连续函数。听起来好像深度无关紧要,但关键在于"足够多神经元"。实践中,深层网络能用指数级更少的参数表示同样的函数。深度带来的是效率,而不仅仅是表达能力。
随着网络变深,会出现两种梯度病。梯度消失(vanishing gradients):当梯度穿过许多层(经由链式法则,第 3 章)时,它会被许多因子相乘。如果这些因子持续小于 1(sigmoid 和 tanh 饱和时就会这样),梯度就指数级地萎缩到零,靠前的层几乎学不到东西。梯度爆炸(exploding gradients):如果因子持续大于 1,梯度就指数级膨胀,导致数值溢出和训练不稳。
解决梯度消失/爆炸的办法:
**权重初始化(weight initialisation)**很重要,因为它决定了训练伊始激活和梯度的尺度。权重太大,激活会爆炸;太小,又会消失。
Xavier(Glorot)初始化从一个方差为 \frac{2}{d_{\text{in}} + d_{\text{out}}} 的分布中取权重。在假设激活是线性或 tanh 的情况下,它能让各层激活的方差大致保持不变。
He(Kaiming)初始化使用方差 \frac{2}{d_{\text{in}}},这是为 ReLU 激活校准的(因为 ReLU 把一半激活置零,你需要双倍的方差来补偿)。
**归一化层(normalisation layer)**通过保证每层的输入有一致的统计量(大致零均值、单位方差)来稳定训练。
**批量归一化(Batch Normalisation, BatchNorm)**沿批次维度做归一化:对每个通道/特征,在小批量的所有样本上算均值和方差,然后归一化。它加上可学习的缩放(\gamma)和平移(\beta)参数,使网络能在需要时撤销归一化:
BatchNorm 有个毛病:它依赖批次大小。批次很小时统计量很嘈杂。推理时你用的是运行平均值而不是批次统计量,这就造成了训练/测试之间的不一致。
**层归一化(Layer Normalisation, LayerNorm)**对每个样本沿特征维度做归一化。它不依赖于批次中的其他样本,因此成为 Transformer 和循环网络的标准选择。
**实例归一化(Instance Normalisation)**对每个样本、每个通道独立地沿空间维度归一化。它在风格迁移里很流行。
**组归一化(Group Normalisation)**把通道分成若干组,在每组内部归一化。它是 LayerNorm 和 InstanceNorm 之间的折中。
Dropout 是一种正则化技术,训练时随机把 p 比例的神经元置零。这迫使网络不依赖任何单个神经元,从而鼓励冗余的表示。测试时所有神经元都激活。**倒置 dropout(inverted dropout)**在训练时把激活乘以 \frac{1}{1-p},这样测试时就无需再缩放。这是标准实现。
**卷积神经网络(Convolutional Neural Network, CNN)**利用空间结构。它不像稠密层那样把每个输入连到每个输出,而是让一个小滤波器(卷积核)在输入上滑动,在每个位置算一个内积。同一组滤波器权重在所有位置共享,这极大地减少了参数,并天然带来平移不变性。
对于二维输入和大小为 k \times k 的卷积核 K,卷积运算是:
输出尺寸取决于三个超参数。**步长(stride)**控制卷积核每次移动多少像素(步长 2 把空间维度减半)。**填充(padding)**在输入边缘补零("same" 填充保持空间尺寸,"valid" 填充不保持)。输出尺寸公式:\text{out} = \lfloor (\text{in} - k + 2p) / s \rfloor + 1。
**池化(pooling)**层对特征图下采样。最大池化取每个窗口里的最大值;平均池化取均值。池化在保留最重要信息的同时减少空间维度。
**空洞卷积(dilated convolution)**在卷积核元素之间插入间隔,在不增加参数的情况下扩大感受野。膨胀率为 2 意味着 3x3 的卷积核覆盖的是 5x5 的区域。
1x1 卷积是用 1x1 卷积核做的卷积。它不看空间邻居,而是在通道之间混合信息。可以把它看作在每个空间位置应用一个稠密层。它被用来廉价地改变通道数。
跳跃连接(skip connection)(残差连接)让输入绕过一个或多个层:\text{output} = F(x) + x。这一层只需学习残差 F(x) = \text{output} - x,当最优变换接近恒等时就容易得多。ResNet(残差网络)用这一招堆叠了超过 100 层,解决了更深的网络反而比浅网络表现更差的退化问题。
CNN 构建起特征层次。靠前的层检测边缘和纹理。中间层把它们组合成部件(眼睛、车轮)。靠后的层识别整个物体。每一层的感受野(它能"看到"的输入区域)随深度增大。
**嵌入(embedding)**把离散的词元(单词、字符、物品 ID)映射到稠密向量。嵌入层就是一张查找表:一个形状为(词表大小, 嵌入维度)的矩阵 E。查找词元 i 就是选 E 的第 i 行。这等价于乘上一个 one-hot 向量,不过是矩阵-向量乘法(第 2 章)的一个特例。嵌入在训练中学习,所以相似的词元最终会得到相似的向量。
**分词(tokenisation)**是把原始文本转换成词元序列的过程。词级分词按空格切分,但处理不了没见过的词。子词分词(subword tokenisation)(BPE、WordPiece、SentencePiece)把文本拆成高频的子词单元,在词表大小和覆盖范围之间取得平衡。"unhappiness"这个词可能变成 ["un", "happiness"] 或 ["un", "happ", "iness"]。
**循环神经网络(Recurrent Neural Network, RNN)**逐个元素地处理序列,维护一个把信息向后传递的隐藏状态:
隐藏状态 h_t 是网络截至时刻 t 所见一切的一个压缩摘要。同样的权重 W_h 和 W_x 在所有时间步上共享(权重共享,就像 CNN 在空间上共享权重一样)。
朴素 RNN 在长序列上举步维艰,因为梯度消失:从第 t 步到第 t-k 步的梯度信号要经过 k 次与 W_h 的相乘,会指数级地萎缩(或爆炸)。
LSTM(Long Short-Term Memory,长短期记忆)通过引入一个独立的细胞状态 c_t 来解决这个问题,它几乎不受干扰地随时间流动。三个门控制什么信息进入、离开和保留:
遗忘门决定从细胞状态里擦除什么:f_t = \sigma(W_f [h_{t-1}, x_t] + b_f)
输入门决定写入什么新信息:i_t = \sigma(W_i [h_{t-1}, x_t] + b_i),候选值为 \tilde{c}_t = \tanh(W_c [h_{t-1}, x_t] + b_c)
细胞状态更新:c_t = f_t \odot c_{t-1} + i_t \odot \tilde{c}_t
输出门决定暴露什么:o_t = \sigma(W_o [h_{t-1}, x_t] + b_o),且 h_t = o_t \odot \tanh(c_t)
细胞状态就像一条传送带:信息可以跨越许多时间步几乎不变地流动(遗忘门保持接近 1),这就解决了长程依赖的梯度消失问题。
GRU(Gated Recurrent Unit,门控循环单元)通过把细胞状态和隐藏状态合二为一,并用两个门代替三个,简化了 LSTM:一个更新门(合并了遗忘和输入)和一个重置门。GRU 参数更少,性能常常与 LSTM 相当。
RNN(包括 LSTM)的根本局限是顺序处理:你必须先处理词元 1,再处理词元 2,再处理词元 3。这阻碍了并行化,并造成信息瓶颈,因为所有上下文都必须挤过那个固定大小的隐藏状态。
**注意力(attention)**同时解决了这两个问题。注意力不再把整个输入压缩成一个固定向量,而是让模型回看所有输入位置,决定哪些与当前输出相关。
现代的表述使用查询、键和值(query, key, value,即 Q、K、V)。把它想象成图书馆检索:你有一个查询(你在找什么)、一些键(每本书上的标签)和一些值(书的实际内容)。你把查询和所有键做比较,弄清楚该检索哪些值。
缩放点积注意力(scaled dot-product attention):
QK^T 计算每个查询和每个键之间的相似度。这是一个矩阵乘法(第 2 章),其元素是内积,衡量的是余弦相似度(第 1 章)。除以 \sqrt{d_k} 防止内积变得过大(否则 softmax 会饱和,产生接近 one-hot 的分布,导致梯度消失)。softmax 把相似度转成概率分布。乘以 V 得到值的加权组合。
**多头注意力(multi-head attention)**并行运行 h 个注意力运算,每个都用 Q、K、V 不同的学习到的投影。这让模型能同时关注来自不同表示子空间的信息。一个头可能关注句法关系,另一个关注语义关系。各头的输出被拼接后再投影:
每个位置得到一个独特的向量,模型可以用它来区分位置。现代模型常常改用可学习的位置嵌入或相对位置编码(RoPE、ALiBi)。
Transformer 并行处理所有词元(自注意力矩阵 QK^T 一次矩阵乘法算完),这使它在现代硬件上比 RNN 训练快得多。代价是自注意力在序列长度上是 O(n^2)(每个词元都要关注其他所有词元),而 RNN 是 O(n)。这也是为什么长上下文模型需要特殊的注意力变体(稀疏注意力、线性注意力、flash attention)。
**视觉 Transformer(Vision Transformer, ViT)**把 Transformer 用到图像上:把图像切成固定大小的小块(比如 16x16),把每块展平成向量,再把这些小块当作一个词元序列。一个可学习的 [CLS] 词元被加在最前面,它最终的表示用于分类。尽管没有卷积的归纳偏置,ViT 在足够数据上训练后能追平甚至超过 CNN。
MLP-Mixer 是一个更简单的架构,它用 MLP 同时取代了注意力和卷积。它在"混合词元"的 MLP(跨空间位置应用)和"混合通道"的 MLP(跨特征应用)之间交替。它的表现颇具竞争力,这说明现代架构的关键洞见并非注意力本身,而是在词元和特征之间高效地混合信息。
**自编码器(autoencoder)**通过训练一个网络去重建自己的输入,来学习压缩表示。编码器把输入映射到一个低维的瓶颈(隐编码),解码器再把它映射回来:
瓶颈迫使网络学习最重要的特征。自编码器用于降维、去噪(在带噪声的输入上训练,重建干净的输出)和异常检测(重建误差高意味着输入反常)。
**变分自编码器(Variational Autoencoder, VAE)加了一个概率化的转折。编码器不再编码到单个点 z,而是输出一个分布的参数(一个高斯的均值 \mu 和方差 \sigma^2)。隐编码从这个分布中采样:z = \mu + \sigma \odot \epsilon,其中 \epsilon \sim \mathcal{N}(0, I)。这个重参数化技巧(reparameterisation trick)**让采样变得可微,从而梯度能够流过。
VAE 的损失有两项:
import jax import jax.numpy as jnp import matplotlib.pyplot as plt from sklearn.datasets import make_circles # 数据 X, y = make_circles(n_samples=500, noise=0.1, factor=0.5, random_state=42) X, y = jnp.array(X), jnp.array(y, dtype=jnp.float32) # 初始化一个 2 层 MLP:2 -> 16 -> 16 -> 1 def init_params(key): k1, k2, k3 = jax.random.split(key, 3) return { 'W1': jax.random.normal(k1, (2, 16)) * 0.5, 'b1': jnp.zeros(16), 'W2': jax.random.normal(k2, (16, 16)) * 0.5, 'b2': jnp.zeros(16), 'W3': jax.random.normal(k3, (16, 1)) * 0.5, 'b3': jnp.zeros(1), } def forward(params, x): h = jnp.maximum(0, x @ params['W1'] + params['b1']) # ReLU h = jnp.maximum(0, h @ params['W2'] + params['b2']) # ReLU logit = (h @ params['W3'] + params['b3']).squeeze() return jax.nn.sigmoid(logit) def loss_fn(params, X, y): pred = forward(params, X) return -jnp.mean(y * jnp.log(pred + 1e-7) + (1 - y) * jnp.log(1 - pred + 1e-7)) grad_fn = jax.jit(jax.grad(loss_fn)) params = init_params(jax.random.PRNGKey(0)) lr = 0.1 for step in range(2000): grads = grad_fn(params, X, y) params = {k: params[k] - lr * grads[k] for k in params} # 画决策边界 xx, yy = jnp.meshgrid(jnp.linspace(-2, 2, 200), jnp.linspace(-2, 2, 200)) grid = jnp.column_stack([xx.ravel(), yy.ravel()]) zz = forward(params, grid).reshape(xx.shape) plt.figure(figsize=(7, 6)) plt.contourf(xx, yy, zz, levels=[0, 0.5, 1], alpha=0.3, colors=['#e74c3c', '#3498db']) plt.scatter(X[y==0,0], X[y==0,1], c='#e74c3c', s=10, label='Class 0') plt.scatter(X[y==1,0], X[y==1,1], c='#3498db', s=10, label='Class 1') plt.title("MLP Decision Boundary on Concentric Circles") plt.legend(); plt.grid(alpha=0.3); plt.show() acc = jnp.mean((forward(params, X) > 0.5) == y) print(f"Accuracy: {acc:.2%}")
jnp.convolve 做对比。import jax.numpy as jnp import matplotlib.pyplot as plt def conv1d(signal, kernel): """从零实现的一维卷积(valid 模式)。""" n, k = len(signal), len(kernel) output = jnp.zeros(n - k + 1) for i in range(n - k + 1): output = output.at[i].set(jnp.sum(signal[i:i+k] * kernel)) return output # 构造一个阶跃信号 t = jnp.linspace(0, 4, 200) signal = jnp.where(t < 1, 0.0, jnp.where(t < 2, 1.0, jnp.where(t < 3, 0.5, 1.5))) # 边缘检测卷积核 edge_kernel = jnp.array([-1.0, 0.0, 1.0]) # 我们的实现 vs 内置实现 our_output = conv1d(signal, edge_kernel) jnp_output = jnp.convolve(signal, edge_kernel, mode='valid') fig, axes = plt.subplots(3, 1, figsize=(10, 6), sharex=True) axes[0].plot(t, signal, color='#3498db', linewidth=1.5) axes[0].set_title("Original Signal"); axes[0].set_ylabel("Value") axes[1].plot(t[:len(our_output)], our_output, color='#e74c3c', linewidth=1.5) axes[1].set_title("After Edge Detection (our conv1d)"); axes[1].set_ylabel("Value") axes[2].plot(t[:len(jnp_output)], jnp_output, color='#27ae60', linewidth=1.5, linestyle='--') axes[2].set_title("After Edge Detection (jnp.convolve)"); axes[2].set_ylabel("Value") axes[2].set_xlabel("t") plt.tight_layout(); plt.show() print(f"Outputs match: {jnp.allclose(our_output, jnp_output)}")
import jax import jax.numpy as jnp import matplotlib.pyplot as plt def scaled_dot_product_attention(Q, K, V): """缩放点积注意力。""" d_k = Q.shape[-1] scores = Q @ K.T / jnp.sqrt(d_k) weights = jax.nn.softmax(scores, axis=-1) output = weights @ V return output, weights # 示例:4 个词元,嵌入维度 8 key = jax.random.PRNGKey(42) k1, k2, k3 = jax.random.split(key, 3) seq_len, d_model = 4, 8 Q = jax.random.normal(k1, (seq_len, d_model)) K = jax.random.normal(k2, (seq_len, d_model)) V = jax.random.normal(k3, (seq_len, d_model)) output, weights = scaled_dot_product_attention(Q, K, V) print(f"Q shape: {Q.shape}") print(f"Attention weights shape: {weights.shape}") print(f"Output shape: {output.shape}") print(f"\nAttention weights (rows sum to 1):") print(weights) print(f"Row sums: {weights.sum(axis=-1)}") # 可视化注意力 fig, ax = plt.subplots(figsize=(5, 4)) im = ax.imshow(weights, cmap='Blues', vmin=0, vmax=1) ax.set_xlabel("Key position"); ax.set_ylabel("Query position") ax.set_title("Attention Weights") tokens = ['tok 0', 'tok 1', 'tok 2', 'tok 3'] ax.set_xticks(range(4)); ax.set_xticklabels(tokens) ax.set_yticks(range(4)); ax.set_yticklabels(tokens) for i in range(4): for j in range(4): ax.text(j, i, f"{weights[i,j]:.2f}", ha='center', va='center', fontsize=10) plt.colorbar(im); plt.tight_layout(); plt.show()
import jax import jax.numpy as jnp import matplotlib.pyplot as plt from sklearn.datasets import make_moons # 数据 X, _ = make_moons(n_samples=500, noise=0.05, random_state=42) X = jnp.array(X) # 自编码器:2 -> 8 -> 1 -> 8 -> 2 def init_ae(key): k1, k2, k3, k4 = jax.random.split(key, 4) return { 'enc_W1': jax.random.normal(k1, (2, 8)) * 0.5, 'enc_b1': jnp.zeros(8), 'enc_W2': jax.random.normal(k2, (8, 1)) * 0.5, 'enc_b2': jnp.zeros(1), 'dec_W1': jax.random.normal(k3, (1, 8)) * 0.5, 'dec_b1': jnp.zeros(8), 'dec_W2': jax.random.normal(k4, (8, 2)) * 0.5, 'dec_b2': jnp.zeros(2), } def encode(p, x): h = jnp.tanh(x @ p['enc_W1'] + p['enc_b1']) return h @ p['enc_W2'] + p['enc_b2'] def decode(p, z): h = jnp.tanh(z @ p['dec_W1'] + p['dec_b1']) return h @ p['dec_W2'] + p['dec_b2'] def ae_loss(p, X): z = encode(p, X) X_hat = decode(p, z) return jnp.mean((X - X_hat) ** 2) grad_fn = jax.jit(jax.grad(ae_loss)) params = init_ae(jax.random.PRNGKey(0)) lr = 0.01 for step in range(3000): grads = grad_fn(params, X) params = {k: params[k] - lr * grads[k] for k in params} z = encode(params, X) X_hat = decode(params, z) fig, axes = plt.subplots(1, 2, figsize=(12, 5)) axes[0].scatter(X[:,0], X[:,1], c=z.squeeze(), cmap='viridis', s=10) axes[0].set_title("Original Data (coloured by latent code)") axes[1].scatter(X_hat[:,0], X_hat[:,1], c=z.squeeze(), cmap='viridis', s=10) axes[1].set_title("Reconstruction from 1D bottleneck") for ax in axes: ax.set_aspect('equal'); ax.grid(alpha=0.3) plt.tight_layout(); plt.show() print(f"Reconstruction MSE: {ae_loss(params, X):.4f}")