返回
查看原链接原链接
Bilibili17分0秒 · —

深度学习模型改进三步法 · 笔记

深度学习模型改进“三步法”:从模块导入到模型缝合实战

核心结论:对深度学习模型进行改进(如添加注意力模块)只需掌握“先导入模块、再实例化赋值、最后在 forward 中调用”的三步法,即可对 UNet、ViT、pix2pix 等主流模型完成模块缝合。


核心要点

背景与问题

  • 很多同学想做深度学习模型改进(如添加注意力机制、卷积模块等),但不知道如何将新模块添加到已有模型中。
  • 本次讲解以 pix2pix(图像迁移模型) 为主案例,并结合 UNetViT 两个典型模型进行实战演示。

模块代码来源

  • 模块代码可从 原论文 中摘取,例如:
  • S-LET:通道独立通道与空间混合注意力
  • EMA:高效多尺度注意力机制
  • 其他经典/前沿注意力机制、卷积块等
  • 讲师整理了一份模块文档,涵盖经典注意力机制、顶会顶刊前沿注意力机制、卷积块等,适用于整个 AI 领域。

三步法总览

  1. 第一步:导入模块(import 或直接复制代码到模型文件中)。
  2. 第二步:在需要插入的位置进行实例属性赋值(如 self.attn = EMA(channel)),完成模块定义。
  3. 第三步:在模型的 forward 前向传播中调用该模块,传入特征图并接收输出。

详细解析与实战案例

案例一:UNet 模型改进(视觉 CV 典型代表)

模型认知
  • UNet 是一个 U 型架构,改进思路是找到可插入的位置
  • 添加 attention(注意力机制)本质上是对特征做权重分配
  • 插入位置:通常在卷积操作之后(即激活机制之后)加一个注意力模块。
修改步骤
  • 找到模型实现文件:unet.py 或有些项目叫 model.py
  • 第一步——导入模块
  from module.ema import EMA  # 从 modules 中导入高效多尺度注意力机制
  • 第二步——实例属性赋值(在 __init__ 中):
  self.attn = EMA(channel)  # 传入通道数
  • 第三步——在 forward 中调用
  x = self.attn(x)  # 将特征图 x 传入模块进行多尺度特征聚合
验证方法
  • 运行训练脚本,通过 Ctrl+F 检索 model,确认加载的是改进后的 UNet。
  • 评估标准
  • 对比改进前后的收敛指标(如 IoU)。
  • 有提升则保留,无提升则更换模块(如换成 CBMMLA 等)。

案例二:ViT 模型改进(更复杂)

模型认知
  • ViT(Vision Transformer)在视觉、自然语言处理、多模态、生成式 AI 中应用广泛。
  • 插入位置:在图像进入 Transformer 的 encoder 特征提取之前,对特征做权重分配。
修改步骤
  • 第一步——导入模块(二选一):
  • 方法一:from module.ema import EMA
  • 方法二:直接将 EMA 模块代码复制到模型配置文件中。
  • 推荐方法一,便于后续替换多个模块进行实验。
  • 第二步——实例属性赋值(在 __init__ 中):
  self.ema = EMA(channel)  # 先给一个初步值,后续根据实际张量形状修正
  • 第三步——维度变换处理(关键难点):
  • 问题:ViT 内部的张量是 3DB, N, D),而 EMA 模块需要 4DB, C, H, W)输入,维度不匹配。
  • 解决方法:打印当前张量形状确认维度后,进行变换。
维度变换详细过程

假设当前张量形状为 B, N, D = [32, 197, 768](B=32,N=197,D=768):

  1. 去掉 class token[32, 197, 768][32, 196, 768](去掉分类信息)。
  2. 重整为图像格式
  • 对 N=196 开根号 = 14,得到 [32, 14, 14, 768](即 B, H, W, D)。
  1. 维度交换[32, 14, 14, 768][32, 768, 14, 14](即 B, D, H, W),此时 D=C=768,符合 4D 张量 B, C, H, W 要求。
  2. 输入 EMA 模块:输出仍为 4D [32, 768, 14, 14]
  3. 恢复原始 3D 形状
  • 维度交换回 [32, 14, 14, 768],展平为 [32, 196, 768]
  • torch.cat 将 class token 加回,恢复为 [32, 197, 768]
  1. 通道数修正:EMA 的通道数需改为 768(即 C=D=768),参数必须传对。
验证结果
  • 训练脚本运行后,加载预训练权重时提示“缺失 EMA 权重”,这是正常现象(预训练权重中没有该新模块),说明模块已成功插入模型。

案例三:pix2pix 模型改进(U-Net 生成器)

模型定位方法
  • 找模型的方式
  1. 直接找模型定义文件(通常在 models/ 目录下)。
  2. 找不到则找 训练脚本train.py),在其中追踪模型来源。
  • 旧版 pix2pix 是基于早期 PyTorch 版本编写,写法与现代略有区别,需定位真正的模型类。
两种实现方式
  1. 旧版实现:模型基于 U-Net 128,在卷积块后面加注意力(如 S-LET 或空间注意力)。
  2. 基于 PyTorch 的主流实现:在 models/ 下的模型定义文件中找到核心模型类。
主流实现修改步骤(以残差生成器为例)
  • 检索 contrastresnet 找到当前使用的生成器类(如 ResnetGenerator)。
  • 第一步——导入模块:在文件开头 from module.ema import EMA
  • 第二步——实例属性赋值(在 __init__ 中):
  self.ema = EMA(channel)
  • 第三步——在 forward 中调用
  • 原代码:output = self.model(input)
  • 修改为:
    output = self.ema(self.model(input))
    # 或使用残差连接:
    output = self.ema(self.model(input)) + input
  • 具体选择(直接返回输出还是残差连接)需通过实验对比效果决定。
注意事项
  • 如果改 U-Net 类,同样遵循三步法。
  • 改模型的前提:必须对当前模型有充分了解,找到真正的模型文件和使用的类(class)。

方法与步骤总结

三步法通用模板

  1. 导入模块:在模型文件开头导入所需模块(或直接将模块代码复制到配置文件中)。
  2. 实例化赋值:在 __init__ 中通过 self.xxx = ModuleName(channel) 定义模块。
  3. forward 调用:在前向传播中调用 output = self.xxx(input_feature)

维度不匹配的通用处理策略

  • 打印张量形状:在 forward 中临时打印输入输出张量形状。
  • 3D ↔ 4D 变换
  • 3D (B, N, D) → 4D (B, C, H, W):去掉 class token → 重整为 B, H, W, D → 交换维度为 B, D, H, W(D=C)。
  • 4D → 3D:交换维度为 B, H, W, C → 展平为 B, H×W, C → 拼接 class token。

模块替换策略

  • 如果改进后指标无提升,可更换模块(如换成 CBM、MLA 等),通过实验对比选择最优方案。

案例与数据

各案例关键信息对比

案例模型架构插入位置特殊处理验证方式
UNetU 型架构卷积之后无需维度变换(已是 4D)训练对比 IoU 指标
ViTTransformerencoder 特征提取之前需 3D↔4D 维度变换检查预训练权重缺失提示
pix2pix残差生成器/U-Net生成器模型输出可选残差连接实验对比效果

关键参数示例

  • ViT 中典型输入:B=32, N=197, D=768,其中 N=196(patch 数)+ 1(class token),图像 patch 排列为 14×14。

限制与待确认问题

  • 模块选择依据:未给出哪种模块(EMA vs. CBM vs. MLA)效果最优的确定性结论,需通过实验验证。
  • 残差连接 vs 直接输出:pix2pix 案例中未明确哪种方式更好,需通过实验决定。
  • 旧版代码兼容性:早期 pix2pix 代码基于旧版 PyTorch 编写,与当前 PyTorch 版本存在差异,需要适应。
  • 通道数确定:ViT 案例中 EMA 模块的通道数需根据实际张量形状(D=C)确定,不能拍脑袋决定。

行动清单

  1. 确定要改进的模型,并充分理解其架构(如 UNet、ViT、pix2pix)。
  2. 从原论文或模块文档中选取合适的注意力/卷积模块代码。
  3. 找到模型定义文件和真正使用的类(通过训练脚本或检索关键词)。
  4. 按三步法操作:导入模块 → 实例化赋值 → forward 中调用。
  5. 若遇到维度不匹配,打印张量形状,进行 3D↔4D 变换处理。
  6. 运行训练脚本,观察预训练权重加载提示(缺失新模块权重 = 插入成功)。
  7. 对比改进前后收敛指标,决定保留或更换模块。

总结

  • 三步法是通用的模型改进方法论,适用于视觉、NLP、多模态等多种场景。
  • 改进的关键在于对模型架构有清晰认知,找到合理的插入位置。
  • 维度变换是常见难点,但通过打印张量形状和系统化变换即可解决。
  • 改进效果以实验对比为准,保留有提升的模块,无提升则替换其他模块继续实验。