PyTorch 2 Export QAT 실전 튜토리얼

PyTorch 2 Export QAT 실전 튜토리얼

torch.export 기반 그래프 모드에서 QAT를 수행하는 방법을 공식 튜토리얼로 실습할 수 있어요. PTQ 플로우와 흡사해서 배우기 쉬워요.

흐름

from torchao.quantization.pt2e.quantize_pt2e import (
    prepare_qat_pt2e, convert_pt2e,
)
from executorch.backends.xnnpack.quantizer.xnnpack_quantizer import (
    get_symmetric_quantization_config, XNNPACKQuantizer,
)

# 1) 프로그램 캡처 (torch.export)
m = torch.export.export(m, example_inputs).module()

# 2) 양자화 인식 훈련 준비
quantizer = XNNPACKQuantizer().set_global(get_symmetric_quantization_config())
m = prepare_qat_pt2e(m, quantizer)

# 3) (훈련 생략)

# 4) 변환
m = convert_pt2e(m)
torchao.quantization.pt2e.move_exported_model_to_eval(m)

알아둘 점

  • prepare_qat_pt2e 는 적절한 위치에 fake quantize를 넣고 Conv2d+BN 같은 QAT 퓨전을 수행해요.
  • 캡처 후 model.eval()/model.train() 대신 move_exported_model_to_eval()/to_train() 을 써요.
  • 몇 에폭 훈련 후엔 observer를 끄고 BN 통계 업데이트를 중지하면 정확도가 더 좋아져요.

더 알아보기