Hugging Face와 함께 ROCm 커널을 쉽게 구축하고 공유하기
Hugging Face는 ROCm 호환 커널의 생성 및 배포를 간소화하기 위한 가이드와 도구를 공개했습니다. kernels 라이브러리와 kernel-builder를 활용하면 개발자는 AMD 하드웨어용 고성능 GPU 연산을 구축하고 Hugging Face Hub를 통해 공유할 수 있어 재현성을 보장하고 PyTorch와의 원활한 통합을 가능하게 합니다.
RadeonFlow GEMM 커널 예시
빌드 과정을 보여주기 위해 Hugging Face는 RadeonFlow GEMM 커널을 사용합니다. 이 커널은 AMD Instinct MI300X GPU에 최적화된 고성능 FP8 블록 단위 행렬 곱셈 구현입니다.
기술 사양
- Precision: 입력에
e4m3fnuzFP8 부동소수점 형식을 사용하여 처리량을 높이고 메모리 대역폭을 줄입니다. - Accuracy: 제한된 FP8 동적 범위에도 불구하고 수치적 안정성을 유지하기 위해 블록별 스케일링 팩터(
a_scale및b_scale)를 사용합니다. - Inputs/Outputs:
a:e4m3fnuz형식의 K × Mb:e4m3fnuz형식의 K × Na_scale:fp32형식의 (K // 128) × Mb_scale:fp32형식의 (K // 128) × (N // 128)c:bf16형식의 M × N
- Recognition: 이 커널은 2025년 6월 AMD Developer Challenge 2025에서 Grand Prize를 수상했습니다.
kernel-builder로 ROCm 커널 구축하기
맞춤형 커널을 개발하면 복잡한 빌드 플래그와 ABI 문제에 직면할 수 있습니다. Hugging Face kernels 라이브러리는 구조화된 프로젝트 조직과 재현성을 위한 Nix 사용을 통해 이러한 복잡성을 추상화합니다.
프로젝트 구조
프로젝트는 빌더가 파일 유형을 식별할 수 있도록 특정 디렉터리로 구성됩니다:
build.toml: 빌드 프로세스를 조정하는 프로젝트 매니페스트.gemm/: 원시 HIP 소스 코드(.hip는 구현,.h는 헤더)를 포함합니다.flake.nix: 종속성을 고정하여 재현 가능한 빌드 환경을 보장합니다.torch-ext/: 커널을 PyTorch 연산자로 노출하기 위해 필요한 C++ 바인딩 및 Python 래퍼를 포함합니다.
구성 및 등록
build.toml: 백엔드(rocm등), 대상 아키텍처(gfx942등 MI300 시리즈용) 및 소스 파일을 정의합니다.- PyTorch Integration: 커널은
TORCH_LIBRARY_EXPAND를 사용하여 네이티브 PyTorch 연산자로 등록됩니다. 이를 통해torch.ops를 통해 커널에 접근할 수 있으며 PyTorch 프레임워크의 일급 구성 요소로 동작합니다. - Python Wrapper:
__init__.py파일은 사용자 친화적인 인터페이스를 제공하며, 기본 연산자를 호출하기 전에 텐서 생성 및 형태 검증을 처리합니다.
재현성 및 배포
Nix 기반 빌드 프로세스
빌드는 Nix를 통해 처리되어 서로 다른 머신 간에 환경이 동일하도록 보장합니다.
- Locking:
nix flake update는flake.lock파일을 생성하여 kernel-builder와 그 종속성을 고정합니다. - Caching: Hugging Face 캐시(
cachix를 통해)를 사용하여 PyTorch 버전의 비용이 많이 드는 재빌드를 방지합니다. - Multi-version Support:
nix build . -L명령은 PyTorch와 ROCm의 모든 지원 버전에 대해 커널 빌드를 자동화할 수 있습니다.
Hugging Face Hub를 통한 배포
빌드가 완료되면 커널은 kernels upload 명령이나 바이너리 파일(.so 파일)을 위한 Git Xet을 통해 Hugging Face Hub에 업로드됩니다. 이는 전통적인 설치가 필요 없게 하며, 사용자는 get_kernel을 사용해 Hub에서 직접 커널을 로드할 수 있습니다:
import torch
from kernels import get_kernel
# Load the kernel from the Hub
gemm = get_kernel("kernels-community/gemm")
# Execute the kernel
result = gemm.gemm(A_fp8, B_fp8, A_scale, B_scale, C)
관련 자료
kernelslibrary: 커널을 구축, 관리 및 로드하기 위한 핵심 라이브러리.- Kernels Community Hub: 커뮤니티가 만든 커널을 발견하고 공유할 수 있는 중앙 저장소.