首页 / PyTorch 入门教程 / 模型剪枝与稀疏化

PyTorch 入门教程

模型剪枝与稀疏化

本教程共 60 篇 · 第 55 篇 · 更新于 2026-08-17 · 约 4 分钟阅读

PyTorch剪枝稀疏化pruning2:4稀疏torch.nn.utils.prune

本节目标:理解剪枝的分类和原理,会用 torch.nn.utils.prune 剪掉不重要的权重,知道什么样的稀疏才能真正加速。

剪枝是什么

深度学习模型普遍「过度参数化」:参数多到明显冗余,很多权重数值接近 0,去掉它们对结果影响很小。剪枝(Pruning)就是把这些不重要的权重置零,模型变稀疏(sparse),体积变小、推理变快。

想象一棵果树,枝丫太多反而不结果。剪掉弱枝,养分集中到主干,果子反而更甜。模型剪完枝通常还要微调一下,道理一样。

剪枝还有一个研究背景:彩票假设(lottery ticket hypothesis)。它发现大模型里藏着一批「中奖」的小子网络,单独训练它们能接近甚至超过大模型的精度。剪枝某种程度上就是在找这类子网络,这也是剪枝在学术上长期热门的原因。

剪枝的时机有讲究。先正常训练到收敛,再剪,然后微调恢复,这是标准流程。训练中途就剪也可以,但梯度会变吵,收敛不稳定,新手不建议碰。

剪枝分两大类:

  • 非结构化剪枝(unstructured):单个权重随便挑着剪,置零的位置是散点。稀疏度可以很高,但普通矩阵运算库不会跳过零,实际加速有限。
  • 结构化剪枝(structured):整行、整列或整个通道一起剪。剪完的形状规整,硬件和库能真正省计算,但精度损失通常更大。

两种流派对应两种目标:想要体积小、研究稀疏性,非结构化够用;想要推理真的变快,结构化才实在。实际部署时两者常常混用,比如卷积按通道剪、全连接按权重剪。

核心 API:torch.nn.utils.prune

torch.nn.utils.prune 提供现成的剪枝方法。最常用的是 l1_unstructured:按权重绝对值从大到小排,把最小的 amount 比例置零。

import torch
import torch.nn as nn
import torch.nn.utils.prune as prune

model = nn.Sequential(nn.Linear(16, 16), nn.ReLU(), nn.Linear(16, 4))
layer = model[0]

prune.l1_unstructured(layer, name="weight", amount=0.3)

amount 传 0 到 1 之间的小数是比例,传正整数是剪掉的具体个数。剪完看一眼:

print(layer.weight_mask)      # 0/1 掩码,1 保留 0 剪掉
print(layer.weight_orig.shape)  # 原始权重还在,另行保存
sparsity = (layer.weight == 0).float().mean()
print(f"稀疏度: {sparsity:.2%}")

这里藏着剪枝的核心机制。剪完后原来的 weight 参数被替换成了三样东西:

  1. weight_orig:未剪枝的原始参数,继续参与训练;
  2. weight_mask:0/1 掩码,以缓冲区(buffer)形式存在;
  3. 一个 forward_pre_hook 钩子:每次前向之前,用 weight = weight_orig * weight_mask 实时算出剪枝后的权重。

好处是剪枝可逆:掩码还在,随时可以调整或恢复。remove 可以把重参数化固化掉,weight 直接变成剪完的版本,掩码和钩子消失。注意这不会撤销剪枝,只是把结果写死。

这套机制还有个好处:剪枝后的模型照常保存加载。state_dict 里同时存着 weight_origweight_masktorch.save 一条龙带走,下次 load_state_dict 回来剪枝状态原样恢复,掩码和钩子都在。

prune.remove(layer, "weight")

结构化剪枝用 ln_structured。比如按第 0 维(卷积的输出通道)的 L2 范数剪一半:

prune.ln_structured(layer, name="weight", amount=0.5, n=2, dim=0)

同一个参数可以反复剪,每次的掩码会叠加,剪枝历史存在 PruningContainer 里。

想实现自己的剪枝策略也不难:继承 BasePruningMethod,实现 compute_mask 方法,声明 PRUNING_TYPEunstructuredstructuredglobal)。比如「隔一个剪一个」的玩法:

class EveryOther(prune.BasePruningMethod):
    PRUNING_TYPE = "unstructured"

    def compute_mask(self, t, default_mask):
        mask = default_mask.clone()
        mask.view(-1)[::2] = 0
        return mask

EveryOther.apply(layer, name="weight")

compute_mask 收两个参数:原始张量和已有掩码。返回的新掩码会和旧掩码叠加,这是迭代剪枝能保持历史的原因。

全局剪枝:从整个模型的角度

上面都是「局部剪枝」:每层各剪各的。更聪明的是全局剪枝(global pruning):把整个模型所有权重放在一起排大小,统一剪掉最小的 20%。结果各层的稀疏度不同,冗余多的层多剪点,整体效果通常更好。

prune.global_unstructured(
    [(model[0], "weight"), (model[2], "weight")],
    pruning_method=prune.L1Unstructured,
    amount=0.2,
)

想剪遍全模型,用 named_modules 遍历,按层类型区别对待:

for name, module in model.named_modules():
    if isinstance(module, nn.Conv2d):
        prune.l1_unstructured(module, name="weight", amount=0.2)
    elif isinstance(module, nn.Linear):
        prune.l1_unstructured(module, name="weight", amount=0.4)
Note

剪完务必微调几轮(在掩码保持不动的前提下继续训练),让剩下的权重弥补损失,这叫「剪枝-微调」范式。不微调直接部署,精度往往惨不忍睹。另外,剪枝后的权重保存在 state_dict 里(weight_orig + weight_mask),正常 torch.save 就能完整存下来。

半结构化稀疏:2:4 的巧思

散点式的非结构化稀疏没法加速,NVIDIA 想了个折中:2:4 半结构化稀疏(semi-structured sparsity)。每 4 个连续元素里恰好剪 2 个,模式固定,硬件就能用专用内核跳过零计算。Ampere 架构(RTX 30 系、A100)起有原生支持。

它的妙处是平衡:稀疏度固定 50%,精度损失却很小——NVIDIA 的测试里 ResNet-50、BERT-Large 剪到 2:4 后精度几乎不变。配合 torch.compile,实测 BERT 推理能到 2 倍加速。

from torch.sparse import to_sparse_semi_structured

# 先把权重掩成 2:4 模式,再转成稀疏格式
# 需要 Ampere+ GPU,且权重为 FP16
linear.weight = torch.nn.Parameter(
    to_sparse_semi_structured(linear.weight)
)

PyTorch 提供了配套工具:torch.ao.pruning.WeightNormSparsifier 负责把权重剪成 2:4 模式,to_sparse_semi_structured 负责压缩加速,整个流程在框架内闭环。

有个形状约束要记住:2:4 稀疏要求张量的最后一维能被 4 整除,BERT 这类模型的线性层基本都满足,但任务头等小层不一定,通常跳过不剪。

稀疏张量与 MaskedTensor

稀疏数据还有专门的存储格式。COO 格式只存「坐标 + 数值」对;CSR 格式用压缩行索引进一步省内存。对高稀疏度的张量,这两种格式能大幅压缩存储。PyTorch 里可用 tensor.to_sparse_csr() 等 API 转换。

它们的差别在存储效率:COO 是通用的,任何稀疏模式都能存;CSR 按行压缩,行内连续的非零越多越划算,适合按行剪枝出来的模式。两种格式目前都是 beta 状态,接口可能微调,学习阶段知道概念就够。

MaskedTensor 是另一个思路:数据和掩码绑在一起,运算时自动跳过被掩码的位置。比如求均值时只统计有效元素,不会把 0 混进去。它现在是原型(prototype)状态,适合研究场景,生产环境谨慎使用。

Tip

给初学者的路径建议:先学会 torch.nn.utils.prune 的局部和全局剪枝,跑通「剪枝-微调」流程;需要真加速时再研究 2:4 半结构化稀疏。散点剪枝省体积可以,指望它加速要谨慎——没有硬件配合的零,只是心理安慰。

最后说一句剪枝和量化的关系。两者都是模型压缩手段,可以叠加:先剪枝去掉冗余权重,再量化把剩下的压成 INT8,模型能同时享受两轮瘦身。实际部署时这俩加 torch.compile 是黄金组合,效果 1+1+1 > 3。