胶囊网络实战:用TensorFlow 2.x从零搭建CapsNet(附MNIST代码)

从零构建胶囊网络:一份面向实践者的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基础上引入胶囊思想改造最后一层,显著降低了误报率。

构建胶囊网络的过程,更像是在编写一个具有“协商”机制的微型社会程序,而非简单的分层变换。尽管其目前在大规模、复杂数据集上的应用和效率仍面临挑战,但作为一种颠覆性的思考方向,亲手实现它无疑能极大地拓宽你对深度学习模型设计的认知边界。代码仓库中的完整脚本包含了数据加载、训练循环、可视化等所有细节,你可以直接运行并开始你的探索。

评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值