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}}}

이 실험을 실행하면, 학습된 학생 ResNet-18이 티처의 부드러운 로짓(KD 손실)과 표준 크로스 엔트로피 손실을 함께 사용하여 학습됩니다.

더 알아보기

  • 완전한 API 문서: https://yoshitomo-matsubara.net/torchdistill/
  • 데모 노트북(예: 중간 표현 추출)은 demo/ 폴더에 있으며, Google Colab에서 직접 열 수 있습니다.
  • 벤치마크 및 샘플 결과는 프로젝트 사이트와 examples/ 디렉터리에 나와 있습니다.

인용 torchdistill을 논문에서 사용할 경우, 2021년 torchdistill 워크숍 논문과 2023년 torchdistill meets Hugging Face 논문의 두 논문을 인용해 주세요. README에 BibTeX 항목이 제공됩니다.


결론 – torchdistill은 재현 가능한 지식 증류 연구를 위한 진정한, 활발히 유지 관리되는 라이브러리입니다. 반복적인 학습 코드를 추상화하고, 다양한 KD 방법을 지원하며, 비전과 NLP 모델 모두와 호환되어, 저수준의 PyTorch 루프에 빠지지 않고 티처-스터디 학습을 실험하고 싶은 누구에게나 유용합니다.

관련

  • 프로젝트
  • 프로젝트
  • Dispatch
  • 프로젝트
  • 프로젝트