Torch-Pruning
基于 DepGraph 的通用结构化剪枝框架,一行代码压缩任意 PyTorch 深度学习模型(CNN
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
基于 DepGraph 的通用结构化剪枝框架,一行代码压缩任意 PyTorch 深度学习模型(CNN
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
想象你有一台性能强大的游戏电脑,跑得动最新的大型3A游戏。但当你打开一个办公软件时,却发现它奇卡无比——问题不在硬件,而是软件本身太臃肿了。深度学习模型同样如此:大语言模型、视觉Transformer、扩散模型……这些模型参数动辄数十亿、数百亿,在消费级GPU上部署几乎是mission impossible。
Torch-Pruning 就是来解决这个问题的。它是一个专门给 PyTorch 深度学习模型「减重」的Python工具库,核心能力是结构化剪枝(Structural Pruning)——把模型里「不重要的部分」系统性地删除,让模型体积更小、推理更快,同时精度损失控制在可接受范围内。
这个项目可不是简单的代码堆砌,它的背后是新加坡国立大学 Learning and Vision Lab 的研究团队,在2023年的计算机视觉顶会 CVPR(IEEE/CVF Conference on Computer Vision and Pattern Recognition) 上发表了论文 DepGraph: Towards Any Structural Pruning,引用量持续增长。
论文第一作者龚凡凡(Gongfan Fang),在GitHub上维护着这个活跃度极高的项目。值得注意的是,项目还持续跟进学术前沿:2024年 NeurIPS Spotlight 的 MaskLLM、2024年 ECCV 的 Isomorphic Pruning 等最新工作都与之相关联。
深度学习模型是由层层神经网络堆叠而成的,这带来一个核心挑战:当你删除某一层的某些参数时,这些改动会像多米诺骨牌一样传导到后续层。
举个例子:假设你要删除卷积层的某个输出通道,这个通道会在下一层产生对应的输入通道——如果你删了A但没删B,模型的输出维度就对不上了。传统方法需要研究人员手动追踪每一层的依赖关系,非常繁琐。
DepGraph 算法的作用就是自动构建整个模型的依赖图(Dependency Graph),将所有「牵一发动全身」的耦合层归为一组,让剪枝过程全自动完成。你只需要告诉它「我要删掉这层的第2、6、9号通道」,DepGraph 自动算出所有需要一起删除的相关层。
Torch-Pruning 真正厉害的地方在于它的通用性。项目从2019年起步,从最初的简单CNN,一路演进到支持:
这种广覆盖度意味着,无论你手头是什么模型,Torch-Pruning 很可能已经能处理。
torch_pruning/
├── dependency/ # 依赖图构建(DepGraph 核心)
├── pruner/ # 剪枝算法实现
│ ├── algorithms/ # 具体剪枝策略(BN、Taylor、随机等)
│ ├── importance.py # 重要性评估
│ └── function.py # 底层剪枝操作
├── ops.py # 预定义各类 PyTorch 层的剪枝支持
└── utils/ # 辅助工具
这种模块化设计让研究人员和工程师可以灵活替换剪枝策略(Importance Criteria)和剪枝粒度,而无需改动模型代码本身。
import torch
from torchvision.models import resnet18
import torch_pruning as tp
model = resnet18(pretrained=True).eval()
DG = tp.DependencyGraph().build_dependency(
model, example_inputs=torch.randn(1,3,224,224)
)
group = DG.get_pruning_group(
model.conv1, tp.prune_conv_out_channels, idxs=[2, 6, 9]
)
if DG.check_pruning_group(group):
group.prune()
只需要四行核心代码,就能完成一个 ResNet18 卷积层的通道剪枝,并自动处理所有连锁依赖。
使用 Torch-Pruning 有几个值得了解的边界条件:
1. 需要开启自动梯度(AutoGrad)
DepGraph 在构建依赖图时需要 PyTorch 的自动微分机制参与分析模型结构,因此不能用 torch.no_grad() 包裹模型推理。如果你的代码中使用了 torch.no_grad() 来节省显存,需要在依赖图构建阶段暂时去掉它。
2. 训练后使用为主 当前版本对稀疏训练(Sparse Training)支持有限,大多数使用场景是训练完成后做模型压缩。对于「边训练边剪枝」的动态稀疏训练场景,建议关注 MaskLLM 等最新研究工作。
3. 保存方式特殊
剪枝后的模型结构发生了变化,不能用 model.state_dict() 方式保存,必须直接保存整个模型:torch.save(model, 'model.pth'),加载时也需要 torch.load(..., weights_only=False)。
在模型压缩领域,剪枝(Pruning)、量化(Quantization)、知识蒸馏(Distillation)是三大核心技术。Torch-Pruning 处于剪枝赛道,但它的定位更偏向于工具基础设施而非单一剪枝算法——它解决了「如何安全地给任意架构剪枝」这个底层问题,让研究者可以专注实现新的剪枝策略,而不用从零写依赖分析。
近年来随着大模型(LLM)的兴起,项目热度明显上升:GitHub Stars 已突破 3300,仓库活跃度(339个 open issues)中可以看出社区参与度相当高。项目持续更新对 Llama、Qwen、DeepSeek 等新模型的适配,说明维护者仍在积极跟进前沿需求。
Torch-Pruning 是一个技术含量高、实用性强、社区活跃的学术型开源项目。它将 CVPR 2023 的前沿研究落地为可用工具,解决了深度学习模型压缩中「结构化剪枝」的核心工程难题。对于 AI 研究者,它提供了可复现的实验框架;对于 AI 工程师,它提供了生产可用的模型压缩流水线。唯一需要注意的是,它是一个纯 Python 库,没有图形界面,需要一定的 PyTorch 基础才能上手。