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 を通じて)。

使い方

  1. ティーチャーモデルと/またはスタディーモデル、データセット、オプティマイザー、適用する蒸留損失を宣言する YAML ファイルを書きます。また、特徴抽出に使用するレイヤーも指定できます。
  2. 提供された CLI を実行する(またはライブラリをインポート) – フレームワークがオブジェクトを構築し、フォワードフックを登録してトレーニングを開始します。
  3. 必要に応じて、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
  • プロジェクト
  • プロジェクト