实现模型剪枝技术

使用 PyTorch Lightning 构建可扩展 AI 模型

Sergiy Tkachuk

Director, GenAI Productivity

何时使用剪枝?

$$

  • 📱 适合在边缘/嵌入式设备部署模型时使用

$$

  • ➕ 可与量化结合,获得叠加效率提升

$$

  • ⚡ 当降低延迟或模型大小为优先时使用
使用 PyTorch Lightning 构建可扩展 AI 模型

什么是模型剪枝?

$$

  • 移除神经网络中不重要的连接
  • 生成便于存储与计算的稀疏模型
  • 常用方法是 L1 非结构化剪枝

剪枝示例

使用 PyTorch Lightning 构建可扩展 AI 模型

什么是模型剪枝?

import torch.nn.utils.prune as prune

prune.l1_unstructured(model.fc, name="weight",
                      amount=0.4)

print(model.fc.weight.data)
tensor([[ 0.25, -0.13,  0.05,  0.70],
        [-0.88,  0.31, -0.02,  0.44]]) # 剪枝前


tensor([[ 0.25, -0.13,  0.00,  0.70],
        [ 0.00,  0.31,  0.00,  0.44]]) # 剪枝后(40% 权重置为 0)
使用 PyTorch Lightning 构建可扩展 AI 模型

理解剪枝掩码

$$

  • 剪枝为每个目标权重张量添加二值掩码

$$

  • 掩码 = 1 --> 保留权重
  • 掩码 = 0 --> 前向计算时将权重设为 0

$$

  • 在移除掩码前,权重仍占用内存
使用 PyTorch Lightning 构建可扩展 AI 模型

使剪枝永久生效

  • 默认情况下,被剪枝的权重仍属于原始张量
  • 要最终确定剪枝,需移除重参数化
  • 将稀疏层转换为权重置零的标准层
Sequential(
  (fc): Linear(
    in_features=128, out_features=64,
    bias=True
    (weight): PrunedParam()
  )
) # 在 prune.remove 之前
import torch.nn.utils.prune as prune

prune.remove(model.fc, 'weight')

# Print model structure
print(model)

Sequential(
  (fc): Linear(in_features=128,
               out_features=64,
               bias=True)
) # 在 prune.remove 之后
使用 PyTorch Lightning 构建可扩展 AI 模型

评估剪枝影响

  • 比较原始模型与剪枝模型的性能
  • 预期准确率略降,但尺寸/内存大幅节省
  • 有助于评估部署时的取舍是否可接受

剪枝权衡图

使用 PyTorch Lightning 构建可扩展 AI 模型

Passons à la pratique !

使用 PyTorch Lightning 构建可扩展 AI 模型

Preparing Video For Download...