torch.compile 소개 — 파이썬 코드를 최적 커널로
torch.compile 소개
torch.compile은 PyTorch 2.0부터 나온, 파이토치 코드를 더 빠르게 만드는 방법이에요. 파이썬 코드를 JIT 컴파일해 최적화된 커널로 바꿔주면서, 필요한 코드 변경은 아주 적어요. 코드를 거의 고치지 않고도 대부분의 모델에서 30%~2배 가량의 속도 향상을 기대할 수 있어요.
동작 원리
torch.compile은 파이썬 코드를 트레이스하면서 PyTorch 연산(op)을 찾아요. 추적이 어려운 코드는 그래프 브레이크(graph break)가 생기는데, 이는 에러가 아니라 '최적화 기회가 줄었다'는 의미예요.
import torch
def fn(x, y):
a = torch.sin(x).cuda()
b = torch.sin(y).cuda()
return a + b
new_fn = torch.compile(fn, backend="inductor")
기본 백엔드는 inductor로, Triton 커널을 생성해요. mode="reduce-overhead" 같은 옵션은 CUDA graphs를 활용해 파이썬 오버헤드를 더 줄여요.
왜 중요한가
- 최소 코드 변경으로 가속.
- 그래프 브레이크도 에러가 아니라 최적화 기회 손실로 처리해 안전하게 폴백.
- PyTorch 2.0 이후 표준 가속 수단.