卷积神经网络 卷积神经网络(CNN)直接从像素数据中学习空间特征层级,用梯度优化得到的滤波器取代手工设计的那一套。本文件涵盖卷积机制、池化、步长、空洞卷积、感受野,以及定义了图像分类的里程碑式架构(LeNet、AlexNet、VGG、ResNet、Inception、EfficientNet)。 在文件 01 里,我们手工设计了用于边缘检测、模糊和角点检测的滤波器。一个很自然的问题是:我们能不能从数据里学到最优的滤波器?这正是卷积神经网络(CNN)所做的事情。 CNN 不再靠手工挑选滤波器权重,而是通过梯度下降(第 6 章)来学,从而发现对当前任务直接有用的特征。 在第 6 章里,我们介绍过卷积运算、CNN 的基础以及"学习滤波器"这个想法。
卷积神经网络(CNN)直接从像素数据中学习空间特征层级,用梯度优化得到的滤波器取代手工设计的那一套。本文件涵盖卷积机制、池化、步长、空洞卷积、感受野,以及定义了图像分类的里程碑式架构(LeNet、AlexNet、VGG、ResNet、Inception、EfficientNet)。
在文件 01 里,我们手工设计了用于边缘检测、模糊和角点检测的滤波器。一个很自然的问题是:我们能不能从数据里学到最优的滤波器?这正是卷积神经网络(CNN)所做的事情。
CNN 不再靠手工挑选滤波器权重,而是通过梯度下降(第 6 章)来学,从而发现对当前任务直接有用的特征。
在第 6 章里,我们介绍过卷积运算、CNN 的基础以及"学习滤波器"这个想法。这里我们更深入地讨论让 CNN 在十多年里一直主导计算机视觉的那些架构创新。
回顾核心的卷积运算:一个大小为 k \times k 的滤波器 K 在输入特征图上滑动,在每个位置算一次点积(第 6 章)。输出尺寸由三个超参数控制:
卷积后的输出空间尺寸:
其中 \text{in} 是输入尺寸,k 是核大小,p 是填充,s 是步长。这个公式对高度和宽度分别独立适用。
一个神经元的**感受野(receptive field)**是原始输入中能影响它的值的区域。
感受野随每层增长:每个卷积层大致增加 k - 1 像素(有步长或空洞时增加更多)。
**池化(pooling)**层在保留最重要信息的同时缩小空间维度。
**全局平均池化(Global Average Pooling,GAP)**把每个通道的整个空间范围平均成一个数,得到一个长度等于通道数的向量。GAP 取代了许多现代架构末尾的全连接层,大幅减少参数量,同时起到结构性正则化的作用。
**批归一化(Batch Normalisation,BatchNorm)**把每个小批量内的激活归一化为零均值、单位方差,然后再做一个可学习的缩放和平移(第 6 章)。在 CNN 中,BatchNorm 是按通道做的:对每个通道独立地在批量和空间维度上计算统计量。它能让训练更稳定、允许更高的学习率,还起到了轻度的正则化作用。
Dropout(第 6 章)在训练时随机把神经元置零。
在 CNN 里,**空间 Dropout(Dropout2D)**丢弃的是整张特征图通道,而不是单个像素,这样更有效,因为特征图里相邻像素高度相关。
**数据增强(data augmentation)**通过在训练时对每张图像施加随机变换来人为地扩大训练集:水平翻转、随机裁剪、旋转、颜色抖动(调整亮度、对比度、饱和度、色相)以及 cutout(遮挡随机矩形区域)。网络看到的每张图都有很多种形态,迫使它学到变换不变的特征,而不是记住特定的像素模式。
进阶的增强策略包括 Mixup(把两张图像及其标签混合:\tilde{x} = \lambda x_i + (1-\lambda) x_j,\tilde{y} = \lambda y_i + (1-\lambda) y_j)、CutMix(把一张图的矩形小块贴到另一张图上,并按面积比例混合标签)以及 RandAugment(从一个固定集合里随机采样一串增强操作,只用一个强度参数控制)。
CNN 架构的历史,就是一路变得更深、更高效的故事,每一代都解决了限制上一代的问题。
LeNet-5(LeCun 等,1998)是最早的 CNN,为手写数字识别设计。两个卷积层接三个全连接层,使用平均池化和 tanh 激活。它证明了学习到的滤波器优于手工设计的特征,但按今天的标准它非常小(6 万参数)。
AlexNet(Krizhevsky 等,2012)以巨大优势赢得 ImageNet 比赛,点燃了深度学习革命。关键创新包括:ReLU 激活(取代了容易梯度消失的 tanh)、用于正则化的 dropout、数据增强以及在 GPU 上训练。五个卷积层、三个全连接层、共 6000 万参数。
VGG(Simonyan 和 Zisserman,2014)证明了只用 3x3 滤波器深堆叠比大滤波器效果更好。两个堆叠的 3x3 滤波器和单个 5x5 滤波器感受野相同,但参数更少(2 \times 3^2 = 18 对比 5^2 = 25),还多了一次非线性。VGG-16(16 层)和 VGG-19(19 层)至今仍被广泛用作特征提取器。它的架构非常简单:卷积块的通道数递增(64、128、256、512),每块后面跟最大池化。
Inception 模块能同时捕捉多个尺度的特征。1x1 滤波器捕捉逐点模式,3x3 捕捉局部纹理,5x5 捕捉更大的结构。拼接把它们全部融合成一个丰富的表示。
ResNet(He 等,2016)解决了退化问题(degradation problem):更深的网络反而比浅网络表现更差,这不是因为过拟合,而是因为更难优化。解决方案是跳跃连接(skip connection,残差连接):
当输入和输出维度不一致(因为步长或通道数变化)时,用一个**投影捷径(projection shortcut)**对 x 做 1x1 卷积来匹配维度:\text{output} = F(x) + W_s x。
瓶颈块(bottleneck block)(用于 ResNet-50 及更深的版本)用三个卷积:1x1 降通道、3x3 做空间处理、1x1 再把通道升回去。这比两个 3x3 卷积更省,能堆出更深的网络。
DenseNet(Huang 等,2017)把跳跃连接的想法推得更远:在密集块内,每一层都和之后的每一层相连。第 l 层把前面所有层的特征图作为输入:x_l = H_l([x_0, x_1, \ldots, x_{l-1}]),其中 [\cdot] 表示沿通道维度拼接。这鼓励特征复用、增强梯度流动,还能减少总参数量。
高效架构面向移动设备和边缘硬件的部署,这些场景对算力、内存和能耗都很敏感。
MobileNet(Howard 等,2017)用**深度可分离卷积(depthwise separable convolutions)**替代标准卷积,把运算分解成两步:
一个 C_{\text{in}} 输入通道、C_{\text{out}} 输出通道的标准 k \times k 卷积在每个空间位置上要做 k^2 \cdot C_{\text{in}} \cdot C_{\text{out}} 次乘法。深度可分离卷积只要 k^2 \cdot C_{\text{in}} + C_{\text{in}} \cdot C_{\text{out}},大约减少了 k^2 倍。对 3x3 滤波器来说,大约便宜 9 倍。
MobileNet-V2 引入了倒残差块(inverted residual block):先用 1x1 卷积升通道,在升维后的空间里做深度卷积,再用 1x1 卷积降回去。跳跃连接放在窄(瓶颈)层之间,与 ResNet 的模式正好相反。扩展比通常是 6。
EfficientNet(Tan 和 Le,2019)引入了复合缩放(compound scaling):与其只独立地缩放深度、宽度或分辨率,不如按固定比例同时缩放这三个维度。给定一个缩放系数 \phi:
ShuffleNet 用**分组卷积(group convolutions)加通道重排(channel shuffle)**来降低 1x1 卷积的成本(在 MobileNet 风格架构里 1x1 卷积占大头)。分组卷积把通道分成若干组,每组内独立做卷积,但这会阻碍跨组的信息流动。重排操作在组之间重新排列通道,几乎不增加开销就恢复了信息混合。
**迁移学习(transfer learning)**是把在一个任务上训练好的模型拿去适应另一个任务的做法。在计算机视觉里,这几乎总是意味着从一个在 ImageNet(140 万张图、1000 类)上预训练的模型出发,再适应到某个特定领域的数据集(医疗影像、卫星图、制造缺陷等)。
特征提取(feature extraction):冻结所有卷积层,移除最后的分类头,只在上面训练一个新的头。被冻结的层充当一个通用的特征提取器。当目标领域和 ImageNet 相似、目标数据集又很小时,这种做法效果很好。
微调(fine-tuning):解冻部分或全部卷积层,用较小的学习率训练。预训练权重作为起点,而不是固定的特征。微调通常从只解冻靠后的层(它们捕捉高层、任务特定的特征)开始,必要时也解冻更早的层。
迁移学习之所以有效,是因为 CNN 的浅层学到的是通用特征(边缘、纹理、颜色),这些在各任务间都有用;而深层学到的是任务特定的特征。一个训练来给动物分类的网络,对给建筑分类同样有可用的边缘检测器。
可视化 CNN 能揭示网络学到了什么,帮助调试异常行为。
**激活图(activation map,特征图)**展示每个滤波器对某张输入图像的输出。浅层激活看起来像边缘图;深层的激活越来越抽象、空间上越来越粗糙。
Grad-CAM(Gradient-weighted Class Activation Mapping,Selvaraju 等,2017)高亮对模型预测最重要的输入图像区域。它的做法是:
**特征反演(feature inversion)**通过优化一张随机图像使其匹配目标特征,从特征表示中重建输入图像(对像素值做梯度下降)。这能揭示网络在每一层保留了多少信息。浅层能重建出近乎完美的图像;深层的重建可辨认但扭曲,说明精细的空间细节丢失了,但语义内容被保留下来。
Deep Dream 和**神经风格迁移(neural style transfer)**是特征可视化的创意应用。Deep Dream 最大化某一层神经元的激活,产生超现实的、图案被放大的图像。神经风格迁移优化一张目标图,使其同时匹配一张图的内容特征(来自深层)和另一张图的风格特征(滤波器激活的 Gram 矩阵,捕捉的是纹理统计量)。
import jax import jax.numpy as jnp import jax.lax as lax import matplotlib.pyplot as plt def conv2d(x, kernel, stride=1): """单个输入、单个滤波器的简单二维卷积。""" return lax.conv(x[None, None], kernel[None, None], (stride, stride), 'SAME')[0, 0] def max_pool(x, size=2): """2x2 最大池化。""" H, W = x.shape x = x[:H//size*size, :W//size*size] return x.reshape(H//size, size, W//size, size).max(axis=(1, 3)) def init_cnn(key): k1, k2, k3 = jax.random.split(key, 3) return { 'conv1': jax.random.normal(k1, (5, 5)) * 0.3, 'conv2': jax.random.normal(k2, (3, 3)) * 0.3, 'fc_w': jax.random.normal(k3, (64, 1)) * 0.1, 'fc_b': jnp.zeros(1), } def forward_cnn(params, img): # Conv1 -> ReLU -> Pool h = jnp.maximum(0, conv2d(img, params['conv1'])) h = max_pool(h) # Conv2 -> ReLU -> Pool h = jnp.maximum(0, conv2d(h, params['conv2'])) h = max_pool(h) # 展平并分类 flat = h.ravel() # 填充或截断到固定大小 flat = jnp.pad(flat, (0, max(0, 64 - len(flat))))[:64] logit = (flat @ params['fc_w'] + params['fc_b']).squeeze() return jax.nn.sigmoid(logit) # 生成合成数据:0 类 = 低频图案,1 类 = 高频图案 def make_data(key, n=200): images, labels = [], [] for i in range(n): k1, key = jax.random.split(key) x, y = jnp.meshgrid(jnp.linspace(0, 4*jnp.pi, 32), jnp.linspace(0, 4*jnp.pi, 32)) if i < n // 2: img = jnp.sin(x) + jax.random.normal(k1, (32, 32)) * 0.1 labels.append(0) else: img = jnp.sin(4 * x) * jnp.sin(4 * y) + jax.random.normal(k1, (32, 32)) * 0.1 labels.append(1) images.append(img) return images, jnp.array(labels, dtype=jnp.float32) key = jax.random.PRNGKey(42) images, labels = make_data(key) params = init_cnn(jax.random.PRNGKey(0)) def loss_fn(params, img, label): pred = forward_cnn(params, img) return -(label * jnp.log(pred + 1e-7) + (1 - label) * jnp.log(1 - pred + 1e-7)) grad_fn = jax.grad(loss_fn) lr = 0.01 for epoch in range(5): total_loss = 0.0 for img, label in zip(images, labels): grads = grad_fn(params, img, label) params = {k: params[k] - lr * grads[k] for k in params} total_loss += loss_fn(params, img, label) print(f"Epoch {epoch}: loss = {total_loss / len(images):.4f}") # 测试准确率 preds = jnp.array([forward_cnn(params, img) > 0.5 for img in images]) acc = jnp.mean(preds == labels) print(f"Accuracy: {acc:.2%}")
import jax.numpy as jnp import matplotlib.pyplot as plt def compute_receptive_field(layers): """从一组 (核大小, 步长) 元组计算感受野大小。""" rf = 1 # 从 1 个像素开始 stride_product = 1 for k, s in layers: rf += (k - 1) * stride_product stride_product *= s return rf # 对比各种架构 configs = { 'Single 5x5': [(5, 1)], 'Two 3x3': [(3, 1), (3, 1)], 'Three 3x3': [(3, 1), (3, 1), (3, 1)], 'Single 7x7': [(7, 1)], '3x3 stride 2 + 3x3': [(3, 2), (3, 1)], } print(f"{'Config':<25} {'RF':>4} {'Params (per channel)':>20}") print('-' * 55) for name, layers in configs.items(): rf = compute_receptive_field(layers) # 参数量:每层的 k^2 之和(按每对输入输出通道计) params = sum(k * k for k, s in layers) print(f"{name:<25} {rf:>4} {params:>20}") # 可视化感受野 fig, axes = plt.subplots(1, 3, figsize=(14, 4)) for ax, (name, rf_size) in zip(axes, [('5x5 filter', 5), ('Two 3x3 filters', 5), ('Three 3x3 filters', 7)]): grid = jnp.zeros((9, 9)) c = 4 # 中心 half = rf_size // 2 grid = grid.at[c-half:c+half+1, c-half:c+half+1].set(1.0) ax.imshow(grid, cmap='Blues', vmin=0, vmax=1) ax.set_title(f'{name}\nRF = {rf_size}x{rf_size}') ax.set_xticks(range(9)); ax.set_yticks(range(9)) ax.grid(True, alpha=0.3) plt.suptitle('Receptive Field Comparison') plt.tight_layout(); plt.show()
import jax import jax.numpy as jnp import matplotlib.pyplot as plt def simple_cnn(params, img): """简单 CNN,同时返回预测和最后一个卷积层的激活。""" # 卷积层(作为 Grad-CAM 的"最后一个卷积层") H, W = img.shape k = params['conv'].shape[0] pad = k // 2 img_pad = jnp.pad(img, pad, mode='edge') activation_map = jnp.zeros((H, W)) for i in range(H): for j in range(W): activation_map = activation_map.at[i, j].set( jnp.sum(img_pad[i:i+k, j:j+k] * params['conv']) ) activation_map = jnp.maximum(0, activation_map) # ReLU # 全局平均池化 -> 全连接 -> 输出 pooled = activation_map.mean() logit = pooled * params['w'] + params['b'] return jax.nn.sigmoid(logit), activation_map # 测试图:左侧有一块亮区(类别指示) img = jnp.zeros((32, 32)) img = img.at[8:24, 4:16].set(1.0) img = img.at[5:10, 20:28].set(0.3) key = jax.random.PRNGKey(42) params = { 'conv': jax.random.normal(key, (5, 5)) * 0.3, 'w': jnp.array(2.0), 'b': jnp.array(-0.5), } # 计算 Grad-CAM def class_score(params, img): pred, _ = simple_cnn(params, img) return pred # 获取激活图和梯度 pred, act_map = simple_cnn(params, img) grad_fn = jax.grad(lambda img: simple_cnn(params, img)[0]) img_grad = grad_fn(img) # 权重 = 梯度的全局平均(简化的单通道 Grad-CAM) alpha = img_grad.mean() grad_cam = jnp.maximum(0, alpha * act_map) # ReLU grad_cam = (grad_cam - grad_cam.min()) / (grad_cam.max() - grad_cam.min() + 1e-8) fig, axes = plt.subplots(1, 3, figsize=(14, 4)) axes[0].imshow(img, cmap='gray'); axes[0].set_title('Input Image'); axes[0].axis('off') axes[1].imshow(act_map, cmap='viridis'); axes[1].set_title('Activation Map'); axes[1].axis('off') axes[2].imshow(img, cmap='gray', alpha=0.6) axes[2].imshow(grad_cam, cmap='jet', alpha=0.4) axes[2].set_title(f'Grad-CAM (pred={pred:.2f})'); axes[2].axis('off') plt.tight_layout(); plt.show()
import jax import jax.numpy as jnp def standard_conv(x, kernel): """标准卷积:(H, W, C_in) * (k, k, C_in, C_out) -> (H, W, C_out)。""" H, W, C_in = x.shape k, _, _, C_out = kernel.shape pad = k // 2 x_pad = jnp.pad(x, ((pad, pad), (pad, pad), (0, 0)), mode='constant') out = jnp.zeros((H, W, C_out)) for i in range(H): for j in range(W): patch = x_pad[i:i+k, j:j+k, :] # (k, k, C_in) for c in range(C_out): out = out.at[i, j, c].set(jnp.sum(patch * kernel[:, :, :, c])) return out def depthwise_separable_conv(x, dw_kernel, pw_kernel): """深度可分离卷积:先深度卷积 (k,k,C_in),再逐点卷积 (C_in, C_out)。""" H, W, C_in = x.shape k = dw_kernel.shape[0] pad = k // 2 x_pad = jnp.pad(x, ((pad, pad), (pad, pad), (0, 0)), mode='constant') # 深度卷积:每个通道一个滤波器 dw_out = jnp.zeros((H, W, C_in)) for i in range(H): for j in range(W): for c in range(C_in): patch = x_pad[i:i+k, j:j+k, c] dw_out = dw_out.at[i, j, c].set(jnp.sum(patch * dw_kernel[:, :, c])) # 逐点卷积:跨通道的 1x1 卷积 out = dw_out @ pw_kernel return out # 配置 H, W, C_in, C_out, k = 8, 8, 16, 32, 3 key = jax.random.PRNGKey(42) k1, k2, k3, k4 = jax.random.split(key, 4) x = jax.random.normal(k1, (H, W, C_in)) std_kernel = jax.random.normal(k2, (k, k, C_in, C_out)) * 0.1 dw_kernel = jax.random.normal(k3, (k, k, C_in)) * 0.1 pw_kernel = jax.random.normal(k4, (C_in, C_out)) * 0.1 # 对比 std_params = k * k * C_in * C_out dw_params = k * k * C_in + C_in * C_out std_flops = H * W * k * k * C_in * C_out dw_flops = H * W * (k * k * C_in + C_in * C_out) print(f"Standard conv: {std_params:>8,} params, {std_flops:>10,} FLOPs") print(f"Depthwise separable conv: {dw_params:>8,} params, {dw_flops:>10,} FLOPs") print(f"Parameter reduction: {std_params / dw_params:.1f}x") print(f"FLOP reduction: {std_flops / dw_flops:.1f}x") std_out = standard_conv(x, std_kernel) ds_out = depthwise_separable_conv(x, dw_kernel, pw_kernel) print(f"\nStandard output shape: {std_out.shape}") print(f"Depthwise sep output shape: {ds_out.shape}")