当前位置: 首页 > news >正文

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

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

在深度学习领域,卷积神经网络(CNN)长期占据主导地位,但它在处理空间层次结构时存在固有缺陷。2017年,深度学习先驱Geoffrey Hinton提出胶囊网络(Capsule Network),通过向量神经元和动态路由机制,显著提升了模型对物体空间关系的理解能力。本文将带您从零实现一个完整的CapsNet模型,包含MNIST数据处理、动态路由算法实现等核心环节,并提供可直接运行的TensorFlow 2.x代码。

1. 环境准备与数据加载

1.1 安装依赖库

确保使用Python 3.7+环境,并安装以下依赖:

pip install tensorflow==2.8.0 matplotlib numpy

1.2 MNIST数据处理

MNIST数据集包含60,000张28x28的手写数字图像。我们使用TensorFlow内置接口加载数据,并进行标准化处理:

import tensorflow as tf # 加载MNIST数据集 (x_train, y_train), (x_test, y_test) = tf.keras.datasets.mnist.load_data() # 数据预处理 def preprocess(images, labels): images = tf.expand_dims(images, -1) # 增加通道维度 images = tf.cast(images, tf.float32) / 255.0 # 归一化 return images, labels # 创建数据管道 train_ds = tf.data.Dataset.from_tensor_slices((x_train, y_train)) train_ds = train_ds.map(preprocess).shuffle(10000).batch(128)

2. 胶囊网络核心组件实现

2.1 PrimaryCaps层构建

PrimaryCaps层将传统标量神经元输出转换为向量形式的胶囊:

class PrimaryCaps(tf.keras.layers.Layer): def __init__(self, caps_dim=8, n_caps=32, kernel_size=9, strides=2): super().__init__() self.caps_dim = caps_dim self.n_caps = n_caps self.conv = tf.keras.layers.Conv2D( filters=caps_dim * n_caps, kernel_size=kernel_size, strides=strides, activation='relu' ) def call(self, inputs): # [batch, 20, 20, 256] -> [batch, 6, 6, 256] outputs = self.conv(inputs) # 重塑为胶囊格式 [batch, 6, 6, 32, 8] outputs = tf.reshape(outputs, [ tf.shape(outputs)[0], -1, self.n_caps, self.caps_dim ]) # 应用squash激活函数 return self.squash(outputs) def squash(self, vectors): norm = tf.norm(vectors, axis=-1, keepdims=True) return (norm / (1 + norm**2)) * vectors

2.2 动态路由算法实现

动态路由是胶囊网络的核心机制,决定低层胶囊如何将信息传递给高层胶囊:

def routing(u_hat, b_ij, iterations=3): """ u_hat: 低层胶囊预测向量 [batch, 1152, 10, 16] b_ij: 初始对数先验 [batch, 1152, 10, 1] """ for i in range(iterations): # 计算耦合系数c_ij c_ij = tf.nn.softmax(b_ij, axis=2) # 计算高层胶囊输入s_j s_j = tf.reduce_sum(c_ij * u_hat, axis=1, keepdims=True) # 应用squash激活函数 v_j = squash(s_j) # 更新对数先验b_ij if i < iterations - 1: agreement = tf.reduce_sum(u_hat * v_j, axis=-1, keepdims=True) b_ij += agreement return tf.squeeze(v_j, axis=1)

3. 完整CapsNet模型构建

3.1 编码器结构

编码器由卷积层、PrimaryCaps层和DigitCaps层组成:

class CapsNet(tf.keras.Model): def __init__(self): super().__init__() # 初始卷积层 self.conv1 = tf.keras.layers.Conv2D( filters=256, kernel_size=9, strides=1, activation='relu', input_shape=(28,28,1) ) # PrimaryCaps层 self.primary_caps = PrimaryCaps() # DigitCaps层参数 self.digit_caps_dim = 16 self.n_digit_caps = 10 self.W = tf.Variable( initial_value=tf.random_normal_initializer(stddev=0.1)( shape=[1, 1152, self.n_digit_caps, self.digit_caps_dim, 8] ), trainable=True ) def call(self, inputs): # 通过卷积层 [batch, 28,28,1] -> [batch,20,20,256] x = self.conv1(inputs) # PrimaryCaps层 [batch,20,20,256] -> [batch,1152,8] u = self.primary_caps(x) # 计算预测向量u_hat [batch,1152,10,16] u = tf.expand_dims(u, axis=2) # [batch,1152,1,8] u = tf.expand_dims(u, axis=3) # [batch,1152,1,1,8] u_hat = tf.reduce_sum(self.W * u, axis=-1) # 动态路由 b_ij = tf.zeros([tf.shape(inputs)[0], 1152, self.n_digit_caps, 1]) v_j = routing(u_hat, b_ij) return v_j

3.2 解码器设计与重建损失

解码器用于正则化训练过程,通过胶囊向量重建原始图像:

class Decoder(tf.keras.layers.Layer): def __init__(self): super().__init__() self.dense1 = tf.keras.layers.Dense(512, activation='relu') self.dense2 = tf.keras.layers.Dense(1024, activation='relu') self.dense3 = tf.keras.layers.Dense(784, activation='sigmoid') def call(self, inputs, y_true): # 仅使用正确类别的胶囊向量 mask = tf.one_hot(y_true, depth=10) masked = tf.reduce_sum(inputs * mask[:, None], axis=1) # 通过全连接层重建图像 x = self.dense1(masked) x = self.dense2(x) x = self.dense3(x) return tf.reshape(x, [-1, 28, 28, 1])

4. 模型训练与评估

4.1 自定义损失函数

胶囊网络使用边缘损失和重建损失的组合:

class CapsuleLoss(tf.keras.losses.Loss): def __init__(self, m_plus=0.9, m_minus=0.1, lambda_=0.5): super().__init__() self.m_plus = m_plus self.m_minus = m_minus self.lambda_ = lambda_ def call(self, y_true, y_pred): # 计算边缘损失 L = y_true * tf.square(tf.maximum(0., self.m_plus - y_pred)) + \ self.lambda_ * (1 - y_true) * tf.square(tf.maximum(0., y_pred - self.m_minus)) return tf.reduce_mean(tf.reduce_sum(L, axis=1))

4.2 训练流程实现

配置自定义训练循环以支持复杂损失计算:

# 初始化模型和优化器 model = CapsNet() decoder = Decoder() optimizer = tf.keras.optimizers.Adam(0.001) loss_fn = CapsuleLoss() @tf.function def train_step(images, labels): with tf.GradientTape() as tape: # 前向传播 caps_output = model(images) # 计算边缘损失 y_true = tf.one_hot(labels, depth=10) caps_loss = loss_fn(y_true, tf.norm(caps_output, axis=-1)) # 计算重建损失 reconstructed = decoder(caps_output, labels) recon_loss = tf.reduce_mean( tf.square(images - reconstructed) ) # 总损失 total_loss = caps_loss + 0.0005 * recon_loss # 反向传播 grads = tape.gradient(total_loss, model.trainable_variables + decoder.trainable_variables) optimizer.apply_gradients(zip(grads, model.trainable_variables + decoder.trainable_variables)) return caps_loss, recon_loss

4.3 可视化训练结果

训练过程中可以定期输出重建图像,直观评估模型表现:

import matplotlib.pyplot as plt def plot_reconstructions(model, decoder, test_images, test_labels, n_samples=5): # 获取模型预测 caps_output = model.predict(test_images[:n_samples]) reconstructions = decoder(caps_output, test_labels[:n_samples]).numpy() # 绘制对比图 plt.figure(figsize=(n_samples * 2, 4)) for i in range(n_samples): # 原始图像 plt.subplot(2, n_samples, i + 1) plt.imshow(test_images[i].squeeze(), cmap='gray') plt.axis('off') # 重建图像 plt.subplot(2, n_samples, i + 1 + n_samples) plt.imshow(reconstructions[i].squeeze(), cmap='gray') plt.axis('off') plt.show()

在实际项目中,我发现动态路由的迭代次数对模型性能影响显著。经过多次实验,3次迭代通常能在训练效率和模型精度间取得良好平衡。另一个关键点是重建损失的权重系数——过大会干扰胶囊学习空间特征,过小则失去正则化效果。建议从0.0005开始逐步调整。

http://www.cnnetsun.cn/news/1527080.html

相关文章:

  • 保姆级教程:用UniApp + DevEco Studio 4.0 从零打包上架一个鸿蒙应用(附全流程截图)
  • 别再只抄代码了!手把手教你给若依(RuoYi)系统加个带权限的自定义接口(附完整前后端配置)
  • Linux内核构建系统:Makefile与Kconfig解析
  • 避坑指南:Double DQN和Dueling DQN在TensorFlow 2.x中的5个常见实现错误
  • 解析 C++ 中的‘生存期保护’:利用生命周期注解规避 99% 的悬挂指针风险
  • AI学习课堂网站丨OPENMAIC丨清华团队开源项目
  • Semilimes SDK:面向MCU的轻量级安全物联网通信框架
  • 云上实战说 | TapNow x Google Cloud 带您体验从灵感到资产的秒级转化
  • 单片机存储器系统架构与工作原理详解
  • OpenClaw日程管理方案:Qwen3.5-9B解析邮件生成待办清单
  • Livox_ros_driver vs driver2:消息类型详解与ROS生态兼容性避坑指南
  • S32K FTM模块实战:从基础配置到电机控制应用
  • OpenClaw多终端控制方案:百川2-13B模型+飞书+网页端协同操作
  • 从零构建微程序控制模型机:运算器与存储器的协同实战
  • 安卓应用集成 FirebaseAuth 实现 Google 登录的完整指南
  • 2026搜索量暴涨!这几款配音软件火到刷屏
  • DeepChem:当AI遇见分子科学,如何重塑药物研发的底层逻辑
  • 医疗陪护管理系统:信息化管理在医院的应用
  • 2026年谷歌商店,谷歌三件套,Google play闪退,从根源排查到品牌适配解决方案
  • 新书速览|Excel+DeepSeek会计与财务高效办公
  • Display Driver Uninstaller深度清理实战指南
  • 嵌入式系统if/else代码优化与设计模式应用
  • 保姆级教程:在Ubuntu 20.04上从零搭建PX4无人机仿真环境(含ROS Noetic和QGC)
  • M5Stack U126 RTC驱动库:PCF8563T嵌入式实时时钟深度解析
  • 不用命令行!Win11任务栏图标消失的图形化解决方案(Explorer重启神器推荐)
  • OpenClaw技能扩展:GLM-4.7-Flash赋能文件整理自动化
  • 告别旧版Vitis HLS!2023.2 Unified IDE保姆级环境配置(含OpenCV 4.4.0 + Vitis Vision库避坑指南)
  • OpenWebUI 集成 Ollama 与 DeepSeek:打造私有化AI助手的全流程实践
  • 多解释器内存隔离实测报告:对比threading/process/subinterpreter三模型,RAM占用降低67%,GC停顿减少91%
  • OpenClaw调试技巧:百川2-13B量化模型任务失败排查手册