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 통계 업데이트를 중지하면 정확도가 더 좋아져요.