AMD MI300용 커스텀 커널 만들기
Hugging Face와 AMD는 VLLM을 사용하여 FP8으로 Llama 3.1 405B를 서비스하는 성능을 향상시키기 위해 AMD MI300X용 오픈소스 최적화 커널 세트를 개발했습니다. 잔차 연결/RMS 정규화/FP8 변환 커널을 결합한 커널, SwiGLU 활성화/FP8 변환 커널을 결합한 커널, 그리고 Skinny GEMM 커널이라는 세 가지 특정 커스텀 커널을 구현함으로써, 디코딩 단계에서(입력 크기 1, 출력 크기 128으로 측정) 지연 시간이 크게 감소했습니다.
커스텀 커널 구현 및 성능 향상
결합된 RMS 정규화 커널
RMS 정규화 커널은 잔차 연결, 행별 Root Mean Square (RMS) 정규화, 그리고 FP8 양자화를 하나의 연산으로 결합하여 디코더 블록의 시작 부분을 최적화합니다.
기술 최적화:
- 벡터화된 메모리 접근: 커널은 128비트 폭 로드를 사용하여 명령당 8개의 FP16 요소를 가져오며, 메모리 접근이 결합되고 연속적으로 이루어져 워프 효율을 최대화합니다.
- 공유 메모리 (SMEM) 활용: 반복적인 VRAM(전역 메모리) 접근을 피하기 위해, 숨겨진 상태 $x$의 수정된 버전을 공유 메모리에 저장합니다. Llama 405B의 경우 차원 $d=16384$가 각 컴퓨트 유닛당 제공되는 64KB 공유 메모리 내에 들어갑니다.
- 블록 수준 감소: 커널은 각 행마다 하나의 스레드 블록을 할당하고, RMS 정규화에 필요한 합산을 위해 스레드들을 공유 메모리로 동기화합니다.
결과: "Vectorized + SMEM" 구현은 표준 PyTorch와 VLLM의 기존 구현을 모두 능가했으며, 다양한 배치 크기에서 상당한 속도 향상을 제공했습니다.
결합된 SwiGLU 커널
SwiGLU 커널은 MLP 블록의 "Gate / Up" 투영에 대한 활성화 함수와 이후의 FP8 양자화를 결합합니다.
기술 최적화:
- 패킹된 명령: 커널은 FP16 덧셈 및 곱셈을 위해 MI300X 패킹 명령을 활용하여 명령당 작업량을 증가시킵니다.
- 빠른 수학 근사: 지연 시간을 줄이기 위해 커널은 표준
exp명령을 입력을 $\log(2)$로 스케일링하여 더 빠른exp2명령으로 교체하며, 이는 정밀도 손실이 거의 없습니다. - 패킹된 FP32에서 FP8 변환: MI300X는 FP32에서 FP8 변환만 지원하므로, 커널은 패킹된 변환 명령을 활용하여 성능을 향상시킵니다.
결과: 커스텀 SwiGLU 커널은 평균적으로 PyTorch보다 14배 이상 빠르고, VLLM 커널보다 27%에서 100% 더 빠릅니다.
Skinny GEMM 커널
표준 라이브러리의 일반 행렬 곱셈(GEMM) 커널은 행이 매우 적은 "skinny" 행렬(낮은 배치 크기로 디코딩할 때 일반적)에서는 비효율적인 경우가 많으며, 이는 타일링 기회가 제한되어 GPU 활용도가 낮아지기 때문입니다.
기술 최적화:
- Split-K 알고리즘: 커널은 공유 K 축을 따라 GEMM을 여러 개의 서브-GEMM으로 나누어 동시에 실행합니다. 이는 작업을 분산시켜 활성 컴퓨트 유닛(CU)의 수를 증가시키고, 각 CU가 K 축에 소비하는 시간을 줄입니다.
- 패딩 제거를 위한 희소성 트릭: 행 수가 최소 밀집 텐서 코어 명령 크기(예: 16)보다 작을 때 낭비되는 패딩을 피하기 위해, 커널은 4:2 구조화된 희소성 명령을 사용합니다. 밀집 8행 행렬을 16행 희소 행렬로 매핑하여
16x16x64희소 명령을 사용할 수 있게 하며, 이는 가장 작은 밀집 명령보다 두 배 깊이를 가집니다. - 워프 특화 및 비동기 실행: 낮은 연산 강도를 해결하기 위해, 커널은 워프를 "생산자"(VRAM에서 공유 메모리로 데이터를 로드하는 전용)와 "소비자"(연산 전용)로 분리합니다. 공유 메모리의 큐가 이 워프들을 비동기적으로 조정하여, 소비자가 느린 VRAM 로드를 기다리는 동안 유휴 상태가 되지 않도록 합니다.
결과: Skinny GEMM 커널은 낮은 행 수(M = 1, 8, 16)에서 PyTorch에 비해 눈에 띄는 속도 향상을 보이며, 특히 QKV 및 Gate/Up 투영에서 그렇지만 배치 크기가 32로 증가함에 따라 이득은 감소합니다.
구현 및 제공
개발된 모든 커널은 hf-rocm-kernels GitHub 저장소에서 제공되며, 여기에는 소스 코드, Python 바인딩, 벤치마크 스크립트 및 테스트 스위트가 포함됩니다. 이 커널들은 독립적으로 사용하거나 VLLM에 통합하도록 설계되었습니다. 결과를 재현하려면 Hugging Face는 개발 중에 사용된 특정 ROCm 6.3.1 컨테이너를 사용할 것을 권장합니다.