Skip to main content
QUICK REVIEW

[论文解读] Streamlining Tensor and Network Pruning in PyTorch

M. Paganini, Jessica Zosa Forde|arXiv (Cornell University)|Apr 28, 2020
Computational Physics and Python Applications参考文献 11被引用 7
一句话总结

本文介绍了 PyTorch 的 `torch.nn.utils.prune` 模块,这是一个统一的、开源的接口,用于在训练、推理或训练后对神经网络层应用结构化和非结构化剪枝。它使研究人员和实践者能够通过极少的代码更改减少模型大小和计算量,使用一致的 API 支持迭代剪枝、全局重要性比较,以及剪枝模型的便捷序列化。

ABSTRACT

In order to contrast the explosion in size of state-of-the-art machine learning models that can be attributed to the empirical advantages of over-parametrization, and due to the necessity of deploying fast, sustainable, and private on-device models on resource-constrained devices, the community has focused on techniques such as pruning, quantization, and distillation as central strategies for model compression. Towards the goal of facilitating the adoption of a common interface for neural network pruning in PyTorch, this contribution describes the recent addition of the PyTorch torch.nn.utils.prune module, which provides shared, open source pruning functionalities to lower the technical implementation barrier to reducing model size and capacity before, during, and/or after training. We present the module's user interface, elucidate implementation details, illustrate example usage, and suggest ways to extend the contributed functionalities to new pruning methods.

研究动机与目标

  • 为在移动设备、物联网和 AR/VR 等资源受限设备上部署大型、过参数化的深度学习模型所面临的日益严峻挑战提供解决方案。
  • 通过在 PyTorch 内提供一个共享的开源接口,降低实现模型剪枝的技术门槛。
  • 使研究人员能够通过统一的 API 轻松实验并贡献新的剪枝技术。
  • 支持训练时和训练后剪枝,采用一致、模块化且可扩展的设计原则。
  • 通过设备端推理实现模型压缩,提升效率、降低能耗并增强隐私保护。

提出的方法

  • 引入 `BasePruningMethod` 作为抽象基类,定义所有剪枝技术的共享接口,要求实现 `compute_mask` 方法。
  • 采用重参数化:通过将原始张量存储为 `name_orig` 和掩码存储为 `name_mask` 在模块缓冲区中,用掩码版本替换被剪枝的参数。
  • 通过前向预钩子应用剪枝,在前向传播期间动态地将原始张量与掩码相乘,从而保持计算图的完整性。
  • 通过专用类如 `L1Unstructured`、`RandomUnstructured` 和 `LnStructured` 支持结构化和非结构化剪枝,可配置剪枝量和维度。
  • 使用 `PruningContainer` 支持对同一参数进行多次剪枝操作的迭代剪枝,实现对多个剪枝操作的追踪。
  • 提供实用函数如 `prune.global_unstructured`,用于在整个网络范围内执行基于全局重要性的剪枝,将所有参数池化后进行比较。

实验结果

研究问题

  • RQ1如何在 PyTorch 内设计一个统一、可扩展且用户友好的剪枝接口,以支持多种剪枝策略?
  • RQ2哪些架构模式能够实现安全、可逆且可组合的剪枝操作,使其无缝集成到 PyTorch 的自动微分和序列化工作流中?
  • RQ3如何高效地实现对整个模型的全局剪枝,同时保持与逐层剪枝和迭代剪枝的兼容性?
  • RQ4哪些机制能确保剪枝后的模型保持可序列化,并能永久保存或无数据丢失地恢复?
  • RQ5如何设计 API 使得研究人员能够轻松实现并贡献新的剪枝方法,而无需深入了解 PyTorch 模块系统的内部细节?

主要发现

  • `torch.nn.utils.prune` 模块成功提供了一致的、开源的接口,可在 PyTorch 中以极少的代码更改实现结构化和非结构化剪枝。
  • 剪枝操作与 PyTorch 的自动微分系统完全兼容,可应用于训练前、训练中或训练后,结果保留在模型的 `state_dict` 中。
  • 通过 `global_unstructured` 支持在整个网络范围内进行全局剪枝,可实现对所有层中连接按重要性排序后剪除最差的 20%。
  • 通过 `PruningContainer` 可对同一参数迭代应用剪枝,支持渐进式压缩策略,例如先剪除 3 个条目,再剪除剩余通道的 50%。
  • 该模块同时支持硬剪枝(二值掩码)和软剪枝,可通过 `prune.remove` 永久移除剪枝,将剪枝后的张量恢复为原始参数名称。
  • 该设计支持剪枝模型的无缝序列化与反序列化,确保与标准 PyTorch 模型保存和加载工作流的兼容性。

更好的研究,从现在开始

从阅读论文到最终审阅,大幅缩短您的研究时间。

无需绑定信用卡

本解读由 AI 生成,并经人工编辑审核。