torch.compile로 추론하기

torch.compile로 추론하기 (inference)

torch.compile은 PyTorch 코드를 최적화된 커널로 컴파일해 추론 속도를 크게 높여줍니다. 이 기능은 TorchDynamo가 코드를 그래프로 컴파일하고, TorchInductor가 그 그래프를 다시 최적화된 커널로 컴파일하는 방식으로 동작해요. 강력한 최적화 도구라서, 많은 경우 단 한 줄의 코드만 추가하면 됩니다.

출처: 문서

본문

모델을 torch.compile로 감싸면 컴파일된 최적화 모델이 반환됩니다.

from transformers import AutoModelForCausalLM

model = AutoModelForCausalLM.from_pretrained("google/gemma-2b", device_map="auto")
compiled_model = torch.compile(model)

[!TIP] torch.compile을 처음 호출할 때는 모델을 컴파일해야 해서 느립니다. 이후 컴파일된 모델을 호출할 때는 다시 컴파일할 필요가 없어 훨씬 빨라져요.

컴파일 과정을 조절할 수 있는 매개변수는 여러 가지가 있습니다. 그중 중요한 두 가지는 아래와 같습니다. 전체 매개변수 목록은 torch.compile 문서를 참고하세요.

모드 (Modes)

mode 매개변수는 컴파일을 위한 여러 성능 옵션을 제공합니다. 여러 모드를 시도해 보고 내 사용 사례에 가장 잘 맞는 것을 찾아 보세요.

  • default는 속도와 메모리의 균형을 잡아주는 옵션입니다.
  • reduce-overhead는 메모리를 조금 더 쓰는 대신 Python 오버헤드를 줄여, 더 빨라질 수 있어요.
  • max-autotune은 가장 빠른 속도를 제공하지만 컴파일에 더 오래 걸립니다.
from transformers import AutoModelForCausalLM

model = AutoModelForCausalLM.from_pretrained("google/gemma-2b", device_map="auto")
compiled_model = torch.compile(model, mode="reduce-overhead")

Fullgraph

Fullgraph는 성능을 극대화하기 위해 모델 전체를 단일 그래프로 컴파일하려 시도합니다. 그래프 중단(graph break)을 만나면 torch.compile은 오류를 일으키는데, 이는 모델을 단일 그래프로 컴파일할 수 없다는 뜻입니다.

from transformers import AutoModelForCausalLM

model = AutoModelForCausalLM.from_pretrained("google/gemma-2b", device_map="auto")
compiled_model = torch.compile(model, mode="reduce-overhead", fullgraph=True)

벤치마크 (Benchmarks)

아래 표는 다양한 GPU와 배치 크기에서, 여러 비전(vision) 작업에 같은 이미지를 사용했을 때 torch.compile을 켜고 끈 상태의 평균 추론 시간(밀리초)을 비교한 성능 벤치마크입니다.

아래 표에서 Subset을 선택하면 서로 다른 GPU는 물론, PyTorch nightly 2.1.0dev와 reduce-overhead 모드가 켜진 torch.compile의 벤치마크로 전환할 수 있어요.