从零构建胶囊网络:一份面向实践者的TensorFlow 2.x深度指南
如果你已经对卷积神经网络(CNN)的经典架构了如指掌,却在某些图像识别任务中遇到了瓶颈——比如,模型偶尔会将一个随意摆放的眼睛、鼻子和嘴巴误判为人脸——那么,是时候将目光投向一个更具“空间意识”的模型了。胶囊网络(Capsule Network, CapsNet)正是为了解决这类问题而生。它由深度学习先驱Geoffrey Hinton提出,其核心思想是让网络不仅识别特征的存在,更能理解特征之间的相对位置、旋转、缩放等姿态关系。对于希望将前沿理论转化为实际代码的开发者而言,理解概念只是第一步,亲手搭建并运行一个CapsNet,才能真正体会其精妙之处。本文将以最经典的MNIST手写数字识别为战场,带你用TensorFlow 2.x从零开始,一步步构建、训练并理解一个完整的胶囊网络。我们将避开繁复的理论推导,聚焦于可运行的代码、清晰的层结构实现,以及那个关键的“动态路由”算法是如何在代码中跃动的。
1. 环境准备与核心概念再认识
在动手写代码之前,确保你的开发环境已经就绪。我们需要TensorFlow 2.x,它直观的Keras API会让我们的构建过程顺畅许多。同时,建议使用Python 3.7以上版本,并配备一块支持CUDA的GPU以加速训练,尽管MNIST数据集较小,CPU也能完成。
pip install tensorflow==2.10.0 numpy matplotlib
胶囊网络的核心单元是“胶囊”(Capsule),它与传统神经元有本质区别。为了让你在编码时心中有图,我们快速回顾几个关键点:
- 向量输出:传统神经元输出一个标量(表示特征激活的强度),而胶囊输出一个向量。这个向量的模长代表了某个实体(如数字“7”)存在的概率,其方向则编码了该实体的姿态信息(如倾斜角度、粗细等)。
- 动态路由:这是CapsNet的灵魂。低层胶囊的预测向量需要被“路由”到高层胶囊。路由权重不是通过反向传播静态学习的,而是在前向传播过程中,通过一个迭代的“协议”动态确定的。简单说,低层胶囊的预测如果与某个高层胶囊的当前输出很“一致”(向量点积大),那么它们之间的连接权重就会增强。
- 等变性:这是CapsNet追求的目标。意味着当输入图像中的对象发生某种变换(如平移、旋转)时,胶囊输出的向量也会发生相应的、可预测的变化(如方向改变),而不仅仅是模长(存在概率)保持不变。
理解了这些,我们就可以开始搭建网络的第一块砖瓦了。
2. 构建胶囊网络的核心层
我们的CapsNet结构遵循经典论文设计,主要包括一个标准的卷积层、一个PrimaryCapsule层和一个DigitCapsule层。我们将以面向对象的方式,自定义Keras层来实现后两个胶囊层。
2.1 PrimaryCapsule层实现
PrimaryCapsule层是第一个胶囊层,它接收卷积层的输出,并将其转换为多个初级胶囊。我们可以将其理解为一组并行的卷积操作,但每个卷积核会产生一个向量而非标量。
import tensorflow as tf
from tensorflow.keras import layers, Model
import numpy as np
class PrimaryCaps(layers.Layer):
"""
主胶囊层。
将卷积特征图转换为多个胶囊向量。
"""
def __init__(self, num_capsules, dim_capsule, kernel_size, strides, padding='valid', **kwargs):
super(PrimaryCaps, self).__init__(**kwargs)
self.num_capsules = num_capsules # 初级胶囊的数量,例如32
self.dim_capsule = dim_capsule # 每个胶囊的维度,例如8
# 使用一个卷积层来生成所有胶囊的激活向量
# 输出通道数 = num_capsules * dim_capsule
self.conv = layers.Conv2D(filters=num_capsules * dim_capsule,
kernel_size=kernel_size,
strides=strides,
padding=padding,
activation='relu')
def call(self, inputs):
# inputs shape: [batch, height, width, channels]
output = self.conv(inputs) # shape: [batch, h', w', num_capsules*dim_capsule]
batch_size = tf.shape(output)[0]
# 重塑为 [batch, h'*w'*num_capsules, dim_capsule]
# 即,将空间位置和胶囊数量合并为“初级胶囊”的总数
output_reshaped = tf.reshape(output, (batch_size, -1, self.dim_capsule))
# 应用 squash 激活函数,确保向量模长在0-1之间
output_squashed = self.squash(output_reshaped)
return output_squashed
def squash(self, vectors, axis=-1, epsilon=1e-7):
"""
Squash 激活函数。
保持向量方向,但将其模长压缩到 (0,1) 区间。
"""
squared_norm = tf.reduce_sum(tf.square(vectors), axis=axis, keepdims=True)
scale = squared_norm / (1 + squared_norm) / tf.sqrt(squared_norm + epsilon)
return scale * vectors
注意:
squash函数是非线性的关键。它避免了使用ReLU或Sigmoid,而是通过一个与模长相关的因子来缩放整个向量,从而在不改变方向的前提下,将模长规范到合适的范围。
2.2 DigitCapsule层与动态路由算法
这是网络最核心的一层,包含了动态路由算法。该层接收PrimaryCapsule层输出的所有向量,并输出10个胶囊(对应0-9十个数字),每个胶囊是一个16维向量。
class DigitCaps(layers.Layer):
"""
数字胶囊层。
通过动态路由算法,将初级胶囊的输出路由到数字胶囊。
"""
def __init__(self, num_capsules, dim_capsule, routing_iterations=3, **kwargs):
super(DigitCaps, self).__init__(**kwargs)
self.num_capsules = num_capsules # 数字胶囊数量,10
self.dim_capsule = dim_capsule # 数字胶囊维度,16
self.routing_iterations = routing_iterations # 路由迭代次数,通常3次足够
def build(self, input_shape):
# input_shape: [batch, num_primary_capsules, dim_primary_capsule]
self.num_primary_capsules = input_shape[1]
self.dim_primary_capsule = input_shape[2]
# 定义变换矩阵 W,用于将初级胶囊的8维输出映射到数字胶囊的16维空间
# W shape: [num_primary_capsules, num_capsules, dim_capsule, dim_primary_capsule]
self.W = self.add_weight(shape=[self.num_primary_capsules,
self.num_capsules,
self.dim_capsule,
self.dim_primary_capsule],
initializer='glorot_uniform',
trainable=True,
name='transformation_matrix')
super(DigitCaps, self).build(input_shape)
def call(self, inputs):
# inputs shape: [batch, num_primary_capsules, dim_primary_capsule]
batch_size = tf.shape(inputs)[0]
# 扩展 inputs 维度以进行矩阵乘法
# inputs_expanded shape: [batch, num_primary_capsules, 1, 1, dim_primary_capsule]
inputs_expanded = tf.expand_dims(tf.expand_dims(inputs, axis=2), axis=2)
inputs_tiled = tf.tile(inputs_expanded, [1, 1, self.num_capsules, 1, 1])
# 计算预测向量 u_hat = W * u
# W 需要扩展 batch 维度
W_expanded = tf.expand_dims(self.W, axis=0) # [1, num_primary, num_caps, dim_cap, dim_primary]
W_tiled = tf.tile(W_expanded, [batch_size, 1, 1, 1, 1])
# u_hat shape: [batch, num_primary_capsules, num_capsules, dim_capsule]
u_hat = tf.squeeze(tf.matmul(W_tiled, inputs_tiled), axis=-1)
# 初始化路由对数 b 为0
b = tf.zeros(shape=[batch_size, self.num_primary_capsules, self.num_capsules])
# 动态路由迭代过程
for i in range(self.routing_iterations):
# 通过 softmax 将路由对数 b 转换为耦合系数 c(每个初级胶囊对高层胶囊的权重)
c = tf.nn.softmax(b, axis=-1) # shape: [batch, num_primary, num_caps]
c_expanded = tf.expand_dims(c, axis=-1) # [batch, num_primary, num_caps, 1]
# 计算高层胶囊的加权输入和 s = sum(c * u_hat)
s = tf.reduce_sum(tf.multiply(c_expanded, u_hat), axis=1) # [batch, num_caps, dim_cap]
# 对 s 应用 squash 函数得到高层胶囊输出 v
v = self.squash(s) # [batch, num_caps, dim_cap]
# 更新路由对数 b = b + u_hat · v
if i < self.routing_iterations - 1:
v_expanded = tf.expand_dims(v, axis=1) # [batch, 1, num_caps, dim_cap]
agreement = tf.reduce_sum(u_hat * v_expanded, axis=-1) # [batch, num_primary, num_caps]
b += agreement
# 最终输出 v shape: [batch, num_capsules=10, dim_capsule=16]
return v
def squash(self, vectors, axis=-1, epsilon=1e-7):
squared_norm = tf.reduce_sum(tf.square(vectors), axis=axis, keepdims=True)
scale = squared_norm / (1 + squared_norm) / tf.sqrt(squared_norm + epsilon)
return scale * vectors
提示:动态路由的过程可以类比于一个“共识形成”机制。初级胶囊不断向高层胶囊“喊话”(发送预测向量u_hat),高层胶囊根据当前所有“喊话”形成一个初步意见(向量v)。初级胶囊发现自己的“喊话”与某个高层胶囊的“意见”越一致(点积越大),下次就会把更多的“音量”(耦合系数c)分配给那个高层胶囊。经过几轮迭代,共识达成,路由关系得以确定。
3. 组装完整的CapsNet模型与损失函数
有了核心层,我们现在来组装完整的编码器(Encoder)和解码器(Decoder),并定义胶囊网络特有的边际损失(Margin Loss)。
3.1 模型组装
编码器包括一个卷积层、PrimaryCaps层和DigitCaps层。解码器是一个简单的全连接网络,用于从数字胶囊向量重建输入图像,起到正则化作用。
def build_capsnet(input_shape=(28, 28, 1), num_classes=10, routing_iterations=3):
"""
构建完整的胶囊网络模型。
"""
# 输入层
input_image = layers.Input(shape=input_shape)
# ---------- 编码器 ----------
# 第一层:标准卷积层
conv = layers.Conv2D(filters=256, kernel_size=9, strides=1, padding='valid', activation='relu')(input_image)
# 第二层:PrimaryCaps层
primary_caps = PrimaryCaps(num_capsules=32, dim_capsule=8, kernel_size=9, strides=2, padding='valid')(conv)
# 第三层:DigitCaps层
digit_caps = DigitCaps(num_capsules=num_classes, dim_capsule=16, routing_iterations=routing_iterations)(primary_caps)
# ---------- 解码器 ----------
# 将正确的数字胶囊向量输入解码器(训练时使用掩码)
masked_capsule = layers.Lambda(lambda x: self.mask(x))(digit_caps) # 这里需要定义mask函数
# 三个全连接层重建图像
fc1 = layers.Dense(512, activation='relu')(masked_capsule)
fc2 = layers.Dense(1024, activation='relu')(fc1)
reconstructed = layers.Dense(np.prod(input_shape), activation='sigmoid')(fc2)
reconstructed_reshaped = layers.Reshape(input_shape)(reconstructed)
# 创建模型
train_model = Model(inputs=input_image, outputs=[digit_caps, reconstructed_reshaped])
return train_model
# 辅助函数:在训练时,只将正确标签对应的胶囊向量传递给解码器
def mask(digit_caps):
# 这里需要结合标签y来操作,具体在损失函数部分体现
pass
3.2 损失函数设计
胶囊网络的损失由两部分组成:用于分类的边际损失和用于正则化的重建损失。
边际损失鼓励正确类别的胶囊模长大(接近1),而不正确类别的模长小(接近0),并设置上下边际。
def margin_loss(y_true, y_pred):
"""
y_true: one-hot标签, shape=[batch, 10]
y_pred: DigitCaps输出, shape=[batch, 10, 16]
"""
# 计算每个数字胶囊向量的模长
norm = tf.sqrt(tf.reduce_sum(tf.square(y_pred), axis=-1)) # shape=[batch, 10]
# 计算边际损失
L = y_true * tf.square(tf.maximum(0., 0.9 - norm)) + \
0.5 * (1 - y_true) * tf.square(tf.maximum(0., norm - 0.1))
return tf.reduce_mean(tf.reduce_sum(L, axis=1))
重建损失就是简单的均方误差(MSE),衡量原始输入图像与解码器重建图像之间的差异。
reconstruction_loss = tf.keras.losses.MeanSquaredError()
在训练时,总损失是边际损失加上一个缩放系数(如0.0005)乘以重建损失。这个系数控制着正则化的强度。
4. 训练、评估与结果分析
现在,让我们在MNIST数据集上训练这个模型,并观察其表现。
4.1 数据准备与训练流程
# 加载MNIST数据
(x_train, y_train), (x_test, y_test) = tf.keras.datasets.mnist.load_data()
# 归一化并增加通道维度
x_train = x_train.reshape(-1, 28, 28, 1).astype('float32') / 255.0
x_test = x_test.reshape(-1, 28, 28, 1).astype('float32') / 255.0
# 转换为one-hot标签
y_train_onehot = tf.keras.utils.to_categorical(y_train, 10)
y_test_onehot = tf.keras.utils.to_categorical(y_test, 10)
# 构建模型
model = build_capsnet()
model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=0.001),
loss=[margin_loss, 'mse'],
loss_weights=[1., 0.0005],
metrics={'digit_caps': 'accuracy'})
# 训练模型
history = model.fit(x_train, [y_train_onehot, x_train],
batch_size=128,
epochs=20,
validation_data=(x_test, [y_test_onehot, x_test]))
4.2 性能评估与可视化
训练完成后,我们不仅关心测试准确率,更想直观感受胶囊网络的“内部理解”。
分类准确率:在MNIST上,一个结构合理的CapsNet通常能达到99.5%以上的测试准确率,与优秀的CNN模型相当。但它的价值不止于此。
重建可视化:这是CapsNet最有趣的部分。我们可以故意扰动数字胶囊向量的某个维度(比如代表笔划粗细或角度的维度),然后让解码器重建图像,观察变化。
import matplotlib.pyplot as plt
def reconstruct_with_perturbation(model, sample_image, true_label, dimension_to_perturb):
"""
扰动指定维度并重建图像。
"""
# 获取样本的胶囊输出
digit_caps_output = model.layers[-3].output # 假设DigitCaps是倒数第三层
partial_model = Model(inputs=model.input, outputs=digit_caps_output)
capsule_vector = partial_model.predict(sample_image[np.newaxis, ...])[0]
original_vector = capsule_vector[true_label].copy()
perturbations = np.linspace(-0.5, 0.5, 5)
fig, axes = plt.subplots(1, len(perturbations), figsize=(15, 3))
for i, pert in enumerate(perturbations):
perturbed_vector = original_vector.copy()
perturbed_vector[dimension_to_perturb] += pert
# 将扰动后的向量输入解码器部分进行重建(需要构建解码器子模型)
# ... 解码器重建代码 ...
# axes[i].imshow(reconstructed, cmap='gray')
# axes[i].set_title(f'Pert: {pert:.1f}')
axes[i].axis('off')
plt.show()
通过这种可视化,你可能会发现某些维度系统地控制着笔划的宽度、数字的倾斜度或局部变形,这直接印证了胶囊“向量方向编码姿态信息”的设计初衷。
4.3 与CNN的对比思考
为了更清晰地理解CapsNet的特性,我们可以从几个维度将其与经典CNN进行对比:
| 特性维度 | 经典CNN (如LeNet-5, VGG) | 胶囊网络 (CapsNet) |
|---|---|---|
| 基本输出单元 | 标量(激活值) | 向量(模长+方向) |
| 空间信息处理 | 依赖池化层实现近似不变性,会丢失位置信息 | 通过向量方向显式编码姿态,追求等变性 |
| 部件-整体关系 | 隐式学习,高层特征对部件空间排列不敏感 | 通过动态路由显式建模,对空间结构鲁棒 |
| 对抗样本鲁棒性 | 相对脆弱,微小扰动可能导致误判 | 理论上对某些空间变换的扰动更鲁棒 |
| 参数量与计算量 | 相对高效,优化成熟 | 参数量通常更大,动态路由增加计算成本 |
| 可解释性 | 特征图可视化,但高层语义模糊 | 胶囊向量维度可能对应具体姿态参数,重建可视化提供直观解释 |
注意:表格中“对抗样本鲁棒性”一项,CapsNet在理论上有优势,但在实际复杂的对抗攻击面前,其优势并不绝对,仍是活跃的研究领域。
在实际项目中,是否选择CapsNet需要权衡。对于MNIST这类简单、规范的数据集,CNN已足够好且更快。但当任务涉及精细的空间关系推理(如医学影像中器官的相对位置、工业质检中零件的组装姿态),或者你对模型的内部可解释性有更高要求时,投入时间研究和应用胶囊网络可能会带来惊喜。我自己的经验是,在一个需要识别轻微旋转和重叠字符的验证码项目上,在CNN基础上引入胶囊思想改造最后一层,显著降低了误报率。
构建胶囊网络的过程,更像是在编写一个具有“协商”机制的微型社会程序,而非简单的分层变换。尽管其目前在大规模、复杂数据集上的应用和效率仍面临挑战,但作为一种颠覆性的思考方向,亲手实现它无疑能极大地拓宽你对深度学习模型设计的认知边界。代码仓库中的完整脚本包含了数据加载、训练循环、可视化等所有细节,你可以直接运行并开始你的探索。
&spm=1001.2101.3001.5002&articleId=154229411&d=1&t=3&u=d0344a5f1cc44071b69cf7e9d2894dea)
368

被折叠的 条评论
为什么被折叠?



