yoshitomo-matsubara/torchdistill
A coding-free framework built on PyTorch for reproducible deep learning studies. PyTorch Ecosystem. 🏆26 knowledge distillation methods presented at TPAMI, CVPR, ICLR, ECCV, NeurIPS, ICCV, AAAI, etc are implemented so far. 🎁 Trained models, training logs and configurations are available for ensuring the reproducibiliy and benchmark.
torchdistill – 知識蒸留のための設定駆動型フレームワーク
何であるか
- PyTorch を基盤とするオープンソースの Python ライブラリで、カスタムのトレーニングループを書かずに 知識蒸留 (ティーチャー・スタディー学習)実験を実行できます。
- すべてのコンポーネント – モデル、データセット、オプティマイザー、損失関数、および蒸留損失自体 – は宣言的 YAML ファイルで記述されます。ライブラリはそのファイルを読み込み、オブジェクトを構築し、実験を実行します。
主な機能(README に記載)
| 機能 | 重要性 |
|---|---|
| モジュール式蒸留手法 – FitNets、Attention Transfer、Relational KD、Variational Information Distillation など、最先端の KD 技術を実装しており、数行の設定で試すことができます。 | |
フォワード・フックマネージャー – モデルの forward メソッドを変更せずに、任意のレイヤーからの中間活性化を取得できます。蒸留(ティーチャー・スタディー特徴マッチング)やモデル解析に有用です。 |
|
「コードなし」実験 – 単一の YAML ファイルを編集するだけで、データセット、モデル、トレーニングハイパーパラメータを含む全体のパイプラインを定義できます。README には、YAML から torchvision.datasets.CIFAR10 オブジェクトを完全に作成する CIFAR-10 の設定例も掲載されています。 |
|
| 広範なタスクカバレッジ – 画像分類、物体検出、セマンティックセグメンテーション、および NLP(Hugging-Face Transformers を通じた GLUE タスク)のための例題スクリプトが提供されています。 | |
| 事前学習済みモデル – CIFAR-10/100 用の再実装モデルや、Hugging-Face Model Hub にホストされたトランスフォーマーチェックポイントへのリンクを含んでいます。 | |
| PyTorch エコシステムメンバ – PyTorch エコシステムに公式に掲載されており、同じパッケージ化およびドキュメント規約に従っています。 | |
インストールが簡単 – pip install torchdistill(または pipenv を通じて)。 |
使い方
- ティーチャーモデルと/またはスタディーモデル、データセット、オプティマイザー、適用する蒸留損失を宣言する YAML ファイルを書きます。また、特徴抽出に使用するレイヤーも指定できます。
- 提供された CLI を実行する(またはライブラリをインポート) – フレームワークがオブジェクトを構築し、フォワードフックを登録してトレーニングを開始します。
- 必要に応じて、
ForwardHookManager.pop_io_dict()を使って保存された中間テンソルを確認し、デバッグや研究分析に活用できます。
典型的なワークフロー例(CIFAR-10)
models:
teacher_model:
key: 'resnet34'
kwargs:
pretrained: true
student_model:
key: 'resnet18'
kwargs:
pretrained: false
datasets:
cifar10/train: !import_call
key: 'torchvision.datasets.CIFAR10'
init:
kwargs:
root: '~/datasets/cifar10'
train: true
download: true
transform: !import_call
key: 'torchvision.transforms.Compose'
init:
kwargs:
transforms:
- !import_call {key: 'torchvision.transforms.RandomCrop', init: {kwargs: {size: 32, padding: 4}}}
- !import_call {key: 'torchvision.transforms.RandomHorizontalFlip', init: {kwargs: {p: 0.5}}}
- !import_call {key: 'torchvision.transforms.ToTensor'}
- !import_call {key: 'torchvision.transforms.Normalize', init: {kwargs: {mean: [0.49,0.48,0.44], std: [0.24,0.24,0.26]}}}
training:
epochs: 200
optimizer: !import_call {key: 'torch.optim.SGD', init: {kwargs: {lr: 0.1, momentum: 0.9, weight_decay: 5e-4}}}
distillation:
loss: !import_call {key: 'torchdistill.loss.kd.KDLoss', init: {kwargs: {temperature: 4.0, alpha: 0.7}}}
この実験を実行すると、教師モデルのソフト化されたロジット(KD損失)と標準の交差エントロピー損失を用いて、学生の ResNet-18 がトレーニングされます。
さらに学ぶには
- 完全な API ドキュメント: https://yoshitomo-matsubara.net/torchdistill/
- デモノートブック(例:中間表現の抽出)は
demo/フォルダにあり、Google Colab で直接開けます。 - ベンチマークとサンプル結果はプロジェクトサイトおよび
examples/ディレクトリに掲載されています。
引用 torchdistill を論文で使用する場合、2021 年の torchdistill ウォークショップ論文と 2023 年の torchdistill meets Hugging Face 論文の 2 つの論文を引用してください。README に BibTeX エントリが提供されています。
結論 – torchdistill は再現可能な知識蒸留研究のための本格的で、積極的にメンテナンスされているライブラリです。ボイラープレートのトレーニングコードを抽象化し、幅広い KD メソッドをサポートし、ビジョンおよび NLP モデルの両方と連携できるため、低レベルの PyTorch ループに深入りせずにティーチャー・スタディー学習を実験したい人にとって非常に有用です。
関連
- プロジェクト
- プロジェクト
- Dispatch
- プロジェクト
- プロジェクト