PyTorch 프로파일링 (파트 2): nn.Linear에서 융합 MLP까지
PyTorch 프로파일링 (파트 2): nn.Linear에서 융합 MLP까지
Hugging Face는 PyTorch에서 다층 퍼셉트론(MLP)을 최적화하는 과정을 상세히 설명하며, 표준 nn.Linear 레이어에서 융합 커널로 전환하여 CPU 오버헤드와 GPU 메모리 트래픽을 줄이는 방법을 보여줍니다. 주요 요점은 torch.compile이 포인트와이즈 연산을 단일 Triton 커널로 융합해 비용이 큰 HBM 왕복을 피할 수 있지만, 손수 튜닝한 커널은 컴파일 지연 및 형태 특화 오버헤드 없이도 유사한 성능 향상을 제공할 수 있다는 것입니다.
nn.Linear와 커널 에필로그의 메커니즘
PyTorch에서 nn.Linear는 행렬 곱셈과 덧셈을 감싸는 래퍼입니다. bias=True일 때 연산은 y = x @ w.T + b와 같이 실행됩니다.
에필로그를 통한 바이어스 폴딩
프로파일링 결과 바이어스 추가를 위한 별도의 aten::add 커널이 존재하지 않음을 확인할 수 있습니다. 대신 바이어스 추가는 epilogue를 사용해 행렬 곱셈 커널에 접합됩니다. 에필로그는 결과를 고대역폭 메모리(HBM)로 다시 쓰기 직전에 GEMM(General Matrix Multiply) 커널이 수행하는 작은 연산입니다. 바이어스 추가를 matmul 커널의 쓰기 단계에 통합함으로써 PyTorch는 HBM에 대한 두 번째 로드와 쓰기를 피하고 메모리 트래픽을 감소시킵니다.
CPU 디스패치와 전치 뷰
aten::t(전치) 연산이 eager CPU 디스패치 체인에서 aten::addmm 앞에 나타납니다. 그러나 이는 GPU 커널을 실행하지 않습니다. aten::t는 단순히 CPU에서 텐서 메타데이터(형태와 stride)를 재작성하여 동일한 원시 데이터에 대한 새로운 뷰를 생성합니다.
torch.compile을 단일 nn.Linear 레이어에 적용해도 사용되는 GPU 커널은 변하지 않으며—동일한 cuBLAS GEMM 커널이 실행됩니다. 대신 컴파일 시점에 결과 stride를 하드코딩함으로써 전치 뷰를 디스패치하는 CPU 오버헤드를 제거합니다.
GeGLU MLP 프로파일링
GeGLU MLP는 세 개의 선형 프로젝션(gate_proj, up_proj, down_proj)과 포인트와이즈 활성화 순서(GeLU 뒤에 곱셈)를 포함합니다.
Eager 실행 성능
Eager 모드에서 GeGLU MLP의 순전파는 다섯 개의 서로 다른 GPU 커널을 실행합니다: 세 개의 GEMM과 두 개의 포인트와이즈 커널(GeLU와 곱셈). GeLU 연산으로 생성된 중간 텐서는 HBM에 기록된 뒤 곱셈 커널에 의해 다시 읽혀야 하므로 메모리 병목 현상이 발생합니다.
GEMM 커널 변동
같은 FLOP 수를 갖더라도 모든 GEMM 커널이 동일한 것은 아닙니다. 예를 들어 down_proj는 입력 형태에 따라 cuBLAS가 다른 타일 크기(예: 128x256 vs 128x128)와 파이프라인 깊이를 선택해 데이터 재사용을 최적화하기 때문에 gate_proj보다 빠를 수 있습니다.
torch.compile와 융합 커널을 통한 최적화
torch.compile의 영향
torch.compile은 GeLU, 곱셈, reshape 연산을 하나의 융합 Triton 커널(예: triton_poi_fused__unsafe_view_gelu_mul_0)로 합쳐 MLP를 최적화합니다.
이 융합은 중간 값을 HBM에 쓰는 대신 레지스터에 유지함으로써 큰 성능 향상을 제공합니다. 그러나 이 최적화에는 비용이 따릅니다: 컴파일러가 TorchDynamo 가드와 프롤로그 오버헤드와 같은 "pre‑ops"를 도입하고, 특정 입력 형태에 맞게 커널을 특화하기 때문입니다. 입력 형태가 바뀌면 모델은 다시 트레이싱하고 재컴파일해야 합니다.
kernels 라이브러리를 이용한 손수 튜닝된 커널
컴파일 지연 및 형태 특화를 피하기 위해 Hugging Face kernels 라이브러리를 통해 손수 튜닝된 커널을 사용할 수 있습니다. 예를 들어 LigerGEGLUMLP 레이어는 GeLU와 곱셈의 융합을 내장한 사전 구축된 Triton 커널을 사용합니다.
컴파일된 커널과 손수 튜닝된 커널 비교:
| Feature | torch.compile (Inductor) |
Hand‑Tuned (Liger) |
|---|---|---|
| Fusion | 포인트와이즈 연산의 동적 융합 | 내장된 융합 |
| CPU Overhead | Dynamo 가드와 프롤로그 | 컴파일 pre‑ops 없음 |
| Shape Handling | 정적 형태에 특화(하나의 형태에 가장 빠름) | 일반적인 런치 파라미터(형태 변화에 강인) |
| Deployment | 로컬 컴파일/Triton 필요 | HF Hub를 통한 사전 구축 바이너리 |
컴파일된 커널은 극도로 특화된 특정 정적 형태에 대해 약간 더 빠를 수 있지만, 손수 튜닝된 커널은 더 견고하며 PyTorch 컴파일 파이프라인과 관련된 위험과 지연을 피합니다.
최적화 단계 요약
| Setup | GPU Change | CPU Change |
|---|---|---|
Eager nn.Linear |
바이어스 추가가 GEMM 에필로그에 접합 | 기본 디스패치 |
Compiled nn.Linear |
변화 없음(동일 cuBLAS 커널) | aten::t 뷰 bookkeeping 제거 |
| Eager MLP | 5개의 커널; 중간 결과가 HBM에 기록 | 기본 디스패치 |
| Compiled MLP | GeLU + mul이 하나의 Triton 커널로 융합 | Dynamo/guard pre‑ops 추가 |
| Liger MLP | 컴파일된 것과 동일한 융합; 튜닝된 런치 파라미터 | 컴파일 pre‑ops 없음; 재컴파일 위험 없음 |
요약: Hugging Face는 GPU 커널 융합 분석, torch.compile의 영향, 그리고 kernels 라이브러리를 통한 손수 튜닝된 커널 사용을 통해 PyTorch에서 다층 퍼셉트론(MLP)을 최적화하는 방법을 설명합니다.
제목: PyTorch 프로파일링 (파트 2): nn.Linear에서 융합 MLP까지