K-Search: CUDA 커널 전문 지식을 Apple Silicon MLX로 이전하기

TL;DR

연구자들은 진화적 커널 최적화 프레임워크인 K-Search에 구조화된 변환 레이어를 추가하여 수십 년에 걸친 CUDA 커널 전문 지식을 Apple Silicon의 MLX 프레임워크에 적용했습니다. 이 접근 방식은 AI가 고성능 GPU 커널을 자동으로 생성하도록 하며, 네이티브 MLX Attention 커널 속도의 0.97배와 커뮤니티 구현에 비해 Mamba SSM 커널의 프리필 속도를 최대 20배까지 달성합니다.

크로스 플랫폼 커널 최적화의 도전 과제

효율적인 GPU 커널을 작성하려면 수년간의 전문 지식이 필요하며, 이러한 최적화를 한 하드웨어 벤더에서 다른 벤더로 이전하려면 보통 처음부터 다시 찾아야 합니다. CUDA 생태계는 attention 및 상태 공간 모델(SSM)과 같은 핵심 연산에 대해 광범위하게 손수 튜닝된 구현을 보유하고 있지만, Apple Silicon과 같은 최신 생태계는 이러한 최적화된 커널 깊이가 부족해 MLX 프레임워크의 능력에도 불구하고 상당한 성능이 활용되지 못하는 경우가 많습니다.

K-Search: 진화적 커널 최적화

K-Search는 반복 루프를 사용해 GPU 커널을 최적화하는 진화적 프레임워크입니다. 이 과정은 세 가지 주요 단계로 구성됩니다:

  1. Action Selection: LLM(특히 본 연구에서는 Gemini 3.5 Pro Preview)을 "GPU 커널 성능 엔지니어"로 활용하여 커널의 분류, 데이터 레이아웃 및 잠재적 병목을 분석하고, 검색 트리("world model")에서 최적화 작업을 제안합니다.
  2. Local Refinement: 코드 작성 모델이 선택된 작업을 기반으로 후보 구현을 생성하고, 이를 실제 하드웨어에서 컴파일 및 벤치마크합니다.
  3. World Model Update: LLM이 결과를 분석하여 새로운 작업을 삽입하거나 우선순위 점수를 업데이트하고, 실패한 경로를 가지치기함으로써 검색 트리를 업데이트합니다.

이 검색은 "Spec"—하드웨어 규칙과 수학적 제약을 포함한 도메인 특화 문서—에 의해 기반을 두어, 잘못된 프리미티브 생성을 방지합니다.

CUDA-to-MLX 변환 레이어

NVIDIA와 Apple Silicon 아키텍처 간의 격차를 메우기 위해, 연구자들은 CUDA 개념 지식을 MLX/Metal 전략으로 변환하는 변환 레이어를 개발했습니다. 이 레이어는 다음으로 구성됩니다:

  • Concept Mapping Tables: CUDA 프리미티브를 Metal 대응물로 매핑하는 용어집으로, 하드웨어별 제약을 포함합니다(예: __shared__ 메모리를 Metal threadgroup 메모리로 매핑하면서 Apple Silicon의 32 KB 제한과 NVIDIA의 48 KB 제한을 고려).
  • MLX-Specific Hints: 직접적인 CUDA 대응이 없는 패턴에 대한 안내로, 레지스터 기반 행 감소에 simd_shuffle_xor를 사용하거나 Apple의 빠른 fast::exp2() 하드웨어 명령을 활용하기 위해 $e^x$를 $2^{x \log_2 e}$ 로 바꾸는 "exp2 트릭" 등을 포함합니다.
  • Reusable Assertions: 전문가 커널 동작을 코드 그대로 복사하는 대신, 진화적 검색이 유지해야 할 속성으로 재구성합니다.

성능 벤치마크

Attention 커널 결과

변환 레이어의 전체 컨텍스트를 진화적 검색에 제공함으로써, 연구자들은 전문가 수준에 근접한 성능을 달성했습니다. 진화된 커널은 다음과 같은 고급 전략을 독립적으로 발견하고 구현했습니다:

  • Threadgroup 메모리 타일링
  • 온라인 소프트맥스
  • 메모리 접근을 위한 K-전치
  • exp2 트릭

이 결과는 순수 진화만으로 0.26배였던 성능이 Apple의 최첨단 네이티브 attention 커널 속도의 0.97배로 급상승했습니다.

Mamba SSM 커널 결과

K-Search를 Mamba 상태 공간 모델(SSM) 커널에 적용하여 일반화 능력을 테스트했습니다. M1 Max(64 GB)에서 mamba-370m f16을 사용했을 때, 진화된 mlx-mamba 커널은 커뮤니티 mlx-lm 구현에 비해 프리필 처리량이 크게 증가했습니다:

지표 mlx-mamba (ours) mlx-lm (community) mamba.py
Decode 152 tok/s 116 tok/s 40 tok/s
Prefill L=512 5,751 tok/s 329 tok/s 1,089 tok/s
Prefill L=1024 6,010 tok/s 327 tok/s 1,127 tok/s
Prefill L=2048 6,612 tok/s 1,092 tok/s 1,092 tok/s
Prefill L=4096 6,743 tok/s 339 tok/s 1,042 tok/s

Key Insight: ~20배 프리필 속도 향상의 원인은 병렬(프리픽스) 스캔 구현에 있습니다. 커뮤니티 mlx-lm 구현이 토큰을 순차적으로 처리하는 반면, 진화된 커널은 결합 연산을 연관적으로 사용해 $O(\log N)$ 단계로 시퀀스를 평가하여 프리필 단계에서 Apple Silicon GPU 처리량을 최대한 활용합니다.

향후 방향

연구자들은 IBM Spyre AIU를 포함한 추가 하드웨어 아키텍처를 지원하도록 이 작업을 확장하고, 융합 MoE 라우팅 및 페이지드 어텐션과 같은 더 복잡한 커널을 개발하고 있습니다. 주요 발견은 AI 기반 커널 생성에서 병목 현상이 LLM의 코딩 능력이 아니라 모델에 제공되는 아키텍처 컨텍스트와 제약 조건의 품질이라는 점입니다.

Sources