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.

研究の動機と目的

  • モバイル、IoT、AR/VRシステムなどリソース制約のあるデバイスへの大規模で過パrameter化されたディーブラーニングモデルのデプロイという、増大する課題に対処すること。
  • PyTorch内での共有でオープンソースのインターフェースを提供することで、モデルプルーニングの実装における技術的障壁を低減すること。
  • 共通のAPIを通じて研究者が新しいプルーニング技術の実験と貢献を容易に行えるようにすること。
  • 一貫性があり、モジュラーで拡張可能な設計原理に従い、トレーニング時およびトレーニング後プルーニングをサポートすること。
  • オンデバイス推論による効率性の向上、エネルギー消費の削減、プライバシーの強化を実現するためのモデル圧縮を促進すること。

提案手法

  • すべてのプルーニング技術に共通するインターフェースを定義する抽象基底クラス`BasePruningMethod`を導入し、`compute_mask`の実装を要求する。
  • 再パrametrizationを採用:プルーニングされたパラメータを、元のテンソルを`name_orig`として、マスクを`name_mask`としてモジュールバッファに格納することで、マスク付きバージョンに置き換える。
  • 前向き伝搬の前に実行するフックを用い、前向き伝搬中に元のテンソルとマスクを動的に乗算することで、計算グラフの整合性を保つ。
  • `L1Unstructured`、`RandomUnstructured`、`LnStructured`などの専用クラスを通じて、構造的および非構造的プルーニングをサポートし、プルーニング量と次元を設定可能にしている。
  • `PruningContainer`を用いて、同じパラメータに対して複数回のプルーニング操作を追跡することで、反復的プルーニングを可能にする。
  • `prune.global_unstructured`などのユーティリティ関数を提供し、ネットワーク全体にわたるグローバルなマグニチュードベースのプルーニングを実行。すべてのパラメータをプールして比較する。

実験結果

リサーチクエスチョン

  • RQ1PyTorch内に、多様なプルーニング戦略をサポートする統合的で拡張可能かつ使いやすいプルーニングインターフェースを設計するにはどうすればよいか?
  • RQ2PyTorchのautogradおよびシリアル化ワークフローにシームレスに統合できる安全で元に戻せたり、組み合わせ可能なプルーニング操作を実現するためのアーキテクチャパターンは何か?
  • RQ3レイヤー単位および反復的プルーニングと互換性を保ちつつ、ネットワーク全体にわたるグローバルプルーニングを効率的に実装するにはどうすればよいか?
  • RQ4プルーディングされたモデルが、データ損失なしに永続的に保存可能または元に戻せるようにする仕組みは何か?
  • RQ5PyTorchのモジュールシステムの内部知識を深く理解しなくても、研究者が新しいプルーニング手法を簡単に実装・貢献できるようにするAPIの構造は何か?

主な発見

  • `torch.nn.utils.prune`モジュールは、PyTorchにおいて構造的および非構造的プルーニングを適用するための一貫性があり、オープンソースのインターフェースを最小限のコード変更で提供している。
  • プルーニング操作はPyTorchのautogradシステムと完全に互換性があり、トレーニング前、トレーニング中、トレーニング後にも適用可能で、結果はモデルの`state_dict`に保持される。
  • `global_unstructured`を介してネットワーク全体にわたるグローバルプルーニングが可能であり、すべてのレイヤーにわたって接続の下位20%をマグニチュードベースでプルーニングできる。
  • 同じパラメータに対して`PruningContainer`を用いて反復的にプルーニングが可能であり、例えば3つのエントリをプルーニングした後に残りのチャンネルの50%をプルーニングするなど、段階的な圧縮戦略を実現できる。
  • ハードプルーニング(バイナリマスク)とソフトプルーニングの両方をサポートし、`prune.remove`を用いることで、プルーニングを永続的に解除し、元のパラメータ名に戻すことができる。
  • 設計により、プルーディングされたモデルのシリアライズおよびデシリアライズがシームレスに可能となり、標準的なPyTorchのモデル保存・読み込みワークフローと互換性がある。

より良い研究を、今すぐ始めましょう

論文の読解から最終レビューまで、研究時間を劇的に削減しましょう。

クレジットカード登録不要

このレビューはAIが作成し、人間の編集者が確認しました。