torch-model-compression
清华大学开源的PyTorch模型压缩工具,通过ONNX图分析实现自动剪枝、量化和重参数化
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
清华大学开源的PyTorch模型压缩工具,通过ONNX图分析实现自动剪枝、量化和重参数化
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
深度学习模型越来越胖——ResNet-50 有 25MB,GPT-2 动辄 500MB,移动端和边缘设备根本跑不动。怎么让大模型在小设备上飞起来?THU-MIG 团队开源的 torch-model-compression 工具库给出了答案:不需要你深入理解模型内部结构,工具会自动分析、自动剪枝、自动量化,一行命令完成模型压缩。
图1:THU-MIG 团队官方头像
在实际工程中,模型部署面临三重矛盾:
精度与速度的矛盾——ResNet-50 在边缘芯片上推理一张图要几百毫秒;
内存与性能的矛盾——手机端 APP 不可能内置一个 100MB 的模型;
成本与规模的矛盾——大模型推理费用居高不下,企业迫切需要降本增效。
传统方案要求开发者必须深入理解模型结构,手动识别哪些卷积通道可以裁剪、哪些权重可以量化。这不仅门槛极高,而且极易出错——一个索引算错,整个模型就废了。清华大学 MIG 团队正是看到这一痛点,开发了这套自动化工具。
torchpruner 是模型结构分析工具,核心是 ONNXGraph 类——它将 PyTorch 模型导出为 ONNX 静态计算图,然后在这个图上做精细化操作。你可以把它理解为模型的X光机:输入一个模型,它能告诉你每个卷积层的输入输出通道数、每个算子的依赖关系。
更重要的是,torchpruner 实现了通道级剪枝分析(cut_analysis)。给定一个卷积层,工具会自动计算:如果裁掉第 0、1、2、3 个输出通道,哪些下游节点的输入维度会因此改变、哪些权重矩阵需要同步裁剪。传统做法要手动推导这些依赖关系,torchpruner 全自动完成。
import torchpruner
import torchvision
# 加载模型并建立 ONNX 静态图
model = torchvision.models.resnet50()
graph = torchpruner.ONNXGraph(model)
graph.build_graph(inputs=(torch.zeros(1,3,224,224),))
# 分析 conv1 层,裁剪后4个通道
result = conv1_module.cut_analysis(attribute_name='weight', index=[0,1,2,3], dim=0)
# 一键执行剪枝,自动处理所有维度对齐问题
model, context = torchpruner.set_cut(model, result)
支持的模型结构覆盖极广:ResNet 系列(含 ResNet56、ResNet110)、VGGNet、MobileNet、ShuffleNet、Inception、MNASNet、UNet、FCN、DeepLab V3 等。支持的算子包括 Conv/Group Conv/TransposeConv/FC、Pooling、BatchNorm、Relu/Sigmoid、concat、view、transpose,以及量化相关算子 quantize_per_tensor/dequantize_per_tensor。
torchslim 集成了多种模型压缩算法,是真正开箱即用的算法库:
重参数化系列(Reparameterization):
剪枝方法(Pruning):
量化感知训练(QAT):
QAT 模块支持感知量化训练,将浮点模型转换为 INT8/INT4 定点模型,并可直接导出为 TensorRT 可部署格式。实测在 ResNet/Unet 上效果良好,量化精度损失控制在 1% 以内。
整个工具库的技术路线非常清晰:以 ONNX 静态图作为桥梁,连接 PyTorch 的动态计算图和最终的压缩操作。
torchpruner 的 graph.py 中定义了核心的 ONNXGraph 类和 create_operator() 工厂函数——它根据 ONNX 算子类型动态注册对应的处理器(operator.py),支持 40+ 种算子的结构化分析。operator/onnx_operator.py 足足有 50KB,是整个仓库最核心的文件之一。 module_pruner/pruners.py 实现了剪枝器的注册机制,mask_utils.py 管理剪枝掩码矩阵,register.py 负责算子和剪枝策略的全局注册。
torchslim 的 slim_solver.py(15KB)定义了统一的压缩求解器接口,pruning/resrep.py 和 pruning/csgd.py 分别实现两种剪枝算法,quantizing/qat.py 实现量化训练逻辑。整体采用Solver模式——用户配置好训练 hook,Solver 自动调度训练循环与压缩操作。
安装极简:一行 python setup.py install 搞定,依赖只有 PyTorch、ONNX、ONNXRuntime、scikit-learn 和 tensorboardX。
零 Docker/无 Web UI:这是纯 Python 算法库,不是可以直接运行的 Web 服务。如果需要产品化部署,需要自己包装服务层。
GPU 不是必须:推理分析阶段完全不依赖 GPU,CPU 即可完成结构分析。但压缩训练阶段(尤其是 ResRep)强烈推荐 GPU——训练速度差异可达 10 倍以上。
Python 版本兼容性:依赖 PyTorch >= 1.7,建议配合 Python 3.6-3.9 使用,避免与最新 Python 版本的兼容性问题。
已知不支持:Faster R-CNN、Mask R-CNN 等两阶段检测器暂不支持,RNN/LSTM/GRU 和 Transformer 系列也在未来支持清单中,使用前需确认目标模型类型。
该项目代表了一个重要趋势:从手动调参到自动压缩的范式转变。随着 AI 落地加速,模型压缩已经从学术研究变成工程必需品。ONNX 作为中间表示的角色也越发重要——它让一次训练、多端部署成为可能。
该工具与 PyTorch 生态的深度整合(直接操作 PyTorch Module 对象、无需脱离训练框架)降低了工程团队的使用门槛。THU-MIG 作为清华大学的研究团队,其算法实现(如 ResRep)具有学术背书,值得信赖。
对于 AI 爱好者而言,这个工具让我们得以窥探模型压缩的内部原理——通过 torchpruner 的 cut_analysis 功能,可以直观看到裁剪一个通道会引发哪些连锁反应。对于 AI 开发者而言,torchslim 提供了开箱即用的生产级压缩工具,可直接集成到模型部署流水线中。
项目信息: