Torch-Pruning

一键剪枝任意 PyTorch 模型,自动处理复杂依赖,压缩提速不降准。

Torch-Pruning 是一个基于 PyTorch 的结构化剪枝工具库,核心实现 CVPR 2023 论文 DepGraph(依赖图)算法,旨在解决传统剪枝方法难以处理复杂、非顺序网络结构的问题。它通过自动分析模型层间的依赖关系,构建统一的依赖图,从而实现对任意架构(包括残差连接、注意力机制、多分支结构等)的安全剪枝,无需手动设计剪枝规则。该库支持对视觉模型(如 ResNet、ViT)、大语言模型(如 LLaMA)及各类基础模型进行通道剪枝、模块剪枝等操作,提供高层次的剪枝接口和底层图分析能力,帮助用户在不显著损失精度的前提下大幅压缩模型体积、降低推理延迟。其核心能力包括自动依赖建模、多种剪枝粒度支持、与 PyTorch 生态无缝集成,以及丰富的示例和基准测试。

开源 free 开源模型
访问官网 ↗ GitHub ↗ 文档 ↗
GitHub 星标 ★ 3358
维护状态 低维护
是否开源 是
定价模式 free

项目数据

分类开源模型
开发团队VainF
所属国家
官网地址
定价模式free
价格说明开源项目,MIT许可证,完全免费,可自行部署使用。
访问状态
是否开源是
开源协议MIT
主要语言Python
技术栈/模型efficient-deep-learning,llm,model-compression,pruning,transformers,vision
GitHub 星标★ 3358
30天Star增速
HF 下载量
上线时间2019-12-15 00:00:00
最近更新2026-09-23 00:00:00
维护状态低维护
中文支持
访问方式
移动端支持
综合评分
收录时间2026-08-09
浏览次数4

使用教程

难度:入门 约 15 分钟 部署方式:库/依赖 6 步

环境要求

  • Python 环境(建议 3.8 及以上)
  • PyTorch(README 测试覆盖 1.x 与 2.x,徽章标注 PyTorch>=2.0)
  • NumPy(Torch-Pruning 仅依赖 PyTorch 和 Numpy)
  • 可用的 pip 包管理工具

安装与启动步骤

  1. 1安装 Torch-Pruning

    官方推荐方式,自动安装最新版;升级已有版本同样使用该命令。

    pip install torch-pruning --upgrade
  2. 2源码可编辑安装

    需要改源码或跟最新代码时使用,克隆仓库后以 -e 模式安装到当前环境。

    git clone https://github.com/VainF/Torch-Pruning.git
    cd Torch-Pruning && pip install -e .
  3. 3加载模型与示例输入

    以 torchvision 的 resnet18 为例准备模型和输入张量,作为剪枝的依赖图分析入口。

    import torch
    from torchvision.models import resnet18
    import torch_pruning as tp
    
    model = resnet18(pretrained=True)
    example_inputs = torch.randn(1, 3, 224, 224)
  4. 4设置忽略层并建剪枝器

    分类器层不参与剪枝,需加入 ignored_layers;BasePruner 无需稀疏训练即可使用。

    ignored_layers = []
    for m in model.modules():
        if isinstance(m, torch.nn.Linear) and m.out_features == 1000:
            ignored_layers.append(m)
    
    pruner = tp.pruner.BasePruner(
        model,
        example_inputs,
        importance=imp,
        pruning_ratio=0.5,
        ignored_layers=ignored_layers,
        round_to=8,
    )
  5. 5执行剪枝并统计

    调用 pruner.step() 完成剪枝,再用 count_ops_and_params 对比剪枝前后的 MACs 与参数量。

    base_macs, base_nparams = tp.utils.count_ops_and_params(model, example_inputs)
    tp.utils.print_tool.before_pruning(model)
    pruner.step()
    tp.utils.print_tool.after_pruning(model)
    macs, nparams = tp.utils.count_ops_and_params(model, example_inputs)
    print(f"MACs: {base_macs/1e9} G -> {macs/1e9} G, #Params: {base_nparams/1e6} M -> {nparams/1e6} M")
  6. 6用自有代码微调

    README 未提供训练脚本,剪枝后需用你自己的训练代码微调剩余权重以恢复精度。

关键配置

配置项必填说明示例
pruning_ratio是默认通道/维度剪枝比例0.5
ignored_layers否不参与剪枝的层列表,如最终分类器ignored_layers
round_to否通道数对齐倍数,便于实际加速,建议 4 或 88
global_pruning否开启全局剪枝,在所有层上统一排名True
isomorphic否开启同构剪枝,缓解全局剪枝过度剪裁问题True
pruning_ratio_dict否按模块或模块元组自定义剪枝比例{(model.layer1, model.layer2): 0.4, model.layer3: 0.2}

如何确认成功

运行后打印的 MACs 与 #Params 相比剪枝前明显下降,且模型仍可正常前向推理。

常见问题

Q:为什么 pruning_ratio=0.5 时参数量并没减少一半?

A:该比例指通道数比例,输入输出维度都会被移除,实际参数剪枝率约为 1-(1-p)^2;想去掉 50% 参数可用 pruning_ratio=0.30。

Q:为什么最后一层分类器要放进 ignored_layers?

A:README 明确提示不要剪枝最后的分类器,否则输出维度与类别数不匹配,会导致模型失效。

Q:如何做全局剪枝或同构剪枝?

A:在 BasePruner 中设置 global_pruning=True 和/或 isomorphic=True 即可;全局剪枝可能过度剪裁某些层,同构剪枝用于缓解该问题。

Q:想逐组查看和控制剪枝过程怎么办?

A:使用 pruner.step(interactive=True) 取回所有 group,再调用 group.prune() 逐个处理,注意 group 必须顺序处理不能存成列表。

注意事项

  • Torch-Pruning 仅依赖 PyTorch 和 Numpy,无需额外重型依赖。
  • 通道维度建议对齐到 4x 或 8x,实际推理加速更明显。
  • 剪枝后必须自行微调,README 未提供训练脚本。
  • 交互式剪枝时 group 必须顺序处理,不要保存为列表后再统一执行。

核心亮点

  • 自动建模层间依赖,支持残差、注意力等复杂结构,无需手工指定剪枝规则
  • 覆盖视觉模型和大语言模型,提供 LLaMA、ViT 等前沿架构的剪枝示例
  • API 设计简洁,与 PyTorch 原生集成,可快速嵌入现有训练/推理流程

不足之处

  • 文档相对简略,高级用法需阅读源码或论文
  • 对超大规模模型(如百亿参数)的剪枝效率有待验证

适用场景

  • 在部署前压缩视觉模型(如 ResNet、ViT)以适配边缘设备
  • 对大语言模型进行结构化剪枝,减少显存占用和推理成本
  • 研究结构化剪枝算法,作为 DepGraph 方法的复现和扩展基础

替代项目

NNI、PaddleSlim、torch.nn.utils.prune

项目介绍

Torch-Pruning 是模型训练领域的开源项目,由 VainF 开发,2019 年首次发布。

在全站 13,090 个收录项目中,它的 GitHub 星标数(3,358)位列前 30%,在模型训练分类的 522 个项目里位列前 13%。

项目已超过三个月没有代码更新,维护节奏明显放缓,最近一次代码更新于 2026-09-23。MIT许可证,完全免费,可自行部署使用。

它主要面向的使用场景是:在部署前压缩视觉模型(如ResNet、ViT)以适配边缘设备。同类可对比的替代方案包括 NNI、PaddleSlim、torch.nn.utils.prune。

上一篇:lorax

下一篇:alpa

同类项目推荐

Long-RL 开源

让强化学习轻松驾驭超长序列,训练更稳更快

Long-RL: Scaling RL to Long Sequences (NeurIPS 2025)

★ 729 2026-08-09
TPA 开源

把注意力复杂度从平方降到线性,长序列不再卡顿

[NeurIPS 2025 Spotlight] TPA: Tensor ProducT ATTenTion Transformer (https://arxiv.or···

★ 463 2026-08-09
minimind 开源

两小时造出你的小模型,练手入门不求人

Train a 64M-parameter LLM from scratch in just 2h!

★ 62178 2026-08-09
pruna 开源

一键优化 AI 模型,推理提速又省资源。

Pruna is a model optimization framework built for developers, enabling you to delive···

★ 1286 2026-08-13