JAXBench는 Google Cloud TPU에서 AI 생성 커널 최적화를 평가하는 최초의 TPU 네이티브 벤치마크 스위트다. Llama-3.1, DeepSeek-V3, Mixtral, Mamba-2, AlphaFold2 등에서 추출한 프로덕션 ML 연산자 17개와 KernelBench에서 이식한 33개 등 총 50개 JAX 워크로드로 구성되며, 8개는 Tokamax의 수작업 최적화 Pallas 커널로 전문가 상한 기준을 제공한다. Gemini 3 Flash 실험에서 문서가 부족한 Pallas DSL에서는 모델 규모보다 타깃 특화 컨텍스트가 더 중요했다. 선별된 TPU 문서를 제공하자 샘플별 정확도가 5.8%에서 37.3%로 뛰었고 50개 중 48개를 해결(기하평균 1.28배 가속), Autocomp 빔서치는 XLA 대비 1.36배, 수작업 튜닝 커널에서는 1.60배까지 도달했다. GPU에 편중됐던 자율 커널 최적화 연구를 TPU 생태계로 확장하는 기반이다.
- •TPU 네이티브 커널 최적화 벤치마크 JAXBench—50개 JAX 워크로드(프로덕션 17 + KernelBench 이식 33)
- •문서 부족한 Pallas DSL에서는 모델 규모보다 타깃 특화 컨텍스트가 성능 결정
- •선별된 TPU 문서 제공 시 정확도 5.8%→37.3%, 50개 중 48개 해결(1.28배 기하평균 가속)
- •Autocomp 빔서치로 XLA 대비 1.36배, 수작업 튜닝 8개 커널에서 1.60배—Tokamax 상한 2.08배에는 미달
- •벤치마크·평가 하네스·베이스라인 공개
JAXBench: Benchmarking Autonomous TPU Kernel Optimization
본문 미리보기
arXiv:2607.20466v1 Announce Type: new Abstract: Rigorous benchmarks have driven progress in autonomous GPU kernel performance optimization by establishing a shared target to hillclimb on, but no equivalent exists for TPUs. We present JAXBench, a TPU-native benchmark suite for AI-generated kernel optimization on Google Cloud TPUs. JAXBench comprises 50 JAX workloads that are both relevant and provide headroom for optimization. We extract 17 production ML operators from architectures in the publi
전체 내용이 궁금하다면?
원문을 직접 읽어보세요