한 줄 요약
MaxKernel은 TPU 커널 최적화를 위한 3가지 에이전트 패러다임을 갖춘 시스템으로, JAXBench에서 1.58× 평균 가속을 달성한다.
핵심 기여도
- **3가지 에이전트 패러다임**: HITL, Auto, Graph-Based Autonomous Search를 통해 커널 개발의 유연성과 성능을 동시에 확보.
- **JaxBench에서 1.58× 평균 가속**: 기존 XLA 기준 대비, 50개 커널 작업에서 지오메트릭 평균 성능 향상.
- **인간 작성 커널 대비 2.32× 가속**: 8개 실제 커널에서 기존 수작업 기준을 초과.
- **오픈소스 공개**: https://github.com/AI-Hypercomputer/accelerator-agents/tree/main/MaxKernel.
핵심 아이디어
MaxKernel은 TPU 커널 개발의 복잡성을 해결하기 위해 **LLM과 실시간 컴파일러 피드백을 결합한 에이전트 기반 시스템**을 제안한다. 기존 커널 개발은 수작업으로 메모리 계층 관리, DMA 파이프라인 조율, 타일링 전략 도출 등이 필요했으며, LLM만으로는 API의 빠가르기, 메모리 제약, 불투명한 컴파일 오류로 인해 한계가 있었다. MaxKernel은 이 문제를 해결하기 위해 **3가지 패러다임**을 도입한다:
1. **HITL (Human-in-the-Loop)**: 인간 개입이 필요한 핵심 결정 시점에서 개발자와 협업.
2. **Auto (Autonomous Loop)**: 계획, 생성, 테스트, 하드웨어 프로파일링을 반복하는 자동화된 최적화 루프.
3. **Graph-Based Autonomous Search**: 커널 설계 공간을 그래프로 모델링하여 **광범위한 탐색**을 수행.
이 시스템은 **공유된 하위 에이전트 풀**을 통해 **계획, 구현, 자가 디버깅, 테스트, 하드웨어 프로파일링**을 처리하며, XProf와 같은 도구와 통합되어 실시간 피드백을 제공한다.
기술적 접근법
- **모듈 구성**: HITL, Auto, Graph-Based Autonomous Search 3가지 패러다임.
- **하위 에이전트**: 커널 설계, 구현, 테스트, 프로파일링을 담당하는 전용 에이전트 풀.
- **알고리즘**:
- Auto 에이전트는 최대 5회 반복하며, 5개 독립 실행 후 중앙값 성능을 기준으로 평가.
- Graph-Based Search는 **빔 너비 3, 최대 깊이 3, 노드당 2개 확장**의 구조를 사용.
- **데이터셋**: **JaxBench (50개 커널 작업)**, **SOTA OSS 모델**.
- **평가 지표**: **compilation rate, correctness rate, geometric mean speedup, fast-p fraction**.
주요 결과
- **JaxBench에서 1.58× 평균 가속**: XLA 기준 대비, 50개 커널 작업에서 지오메트릭 평균 성능 향상.
- **8개 실제 커널에서 2.32× 가속**: 수작업 기준 대비, MaxKernel이 상위 성능.
- **정확도**: **jnp.allclose**를 사용한 수치 비교에서 **대부분 10⁻² 허용 오차 내 정확도** 달성.
- **컴파일 성공률**: 생성된 커널 중 **대부분 성공적으로 컴파일**됨.
의의 및 한계
MaxKernel은 **LLM 기반 커널 생성과 하드웨어 최적화 간의 격차를 해소**하는 데 기여하며, **TPU 최적화 작업의 자동화와 가속화**를 가능하게 한다. 특히, **XProf 기반 실시간 피드백**과 **그래프 기반 탐색**을 통해 **로컬 최적점 탈출**과 **광범위한 설계 공간 탐색**이 가능하다는 점에서 학술적·실용적 가치가 있다.
하지만, **에이전트의 탐색 전략**은 여전히 제한적일 수 있으며, **다양한 하드웨어 아키텍처에 대한 확장성**은 추가 연구가 필요하다. 또한, **LLM의 생성 불확실성**은 여전히 존재하며, **정확한 하드웨어 프로파일링과 조합된 정책 개선**이 필요하다.
실용적 활용
MaxKernel은 **TPU 기반 딥러닝 모델의 커널 최적화**, **고성능 컴퓨팅 분야의 커스텀 가속기 개발**, **LLM 및 대규모 모델의 효율적 실행**에 적용 가능하다. 특히, **JAX/Pallas 기반 개발 환경**에서 실용적 활용이 높으며, **소프트웨어-하드웨어 통합 최적화**를 필요로 하는 산업 분야에 유용하게 사용될 수 있다.