AI News HubLIVE
サイト内リライト2 分で読了

JAXBench:自律的なTPUカーネル最適化のベンチマーク

JAXBenchは、Google Cloud TPU向けのAI生成カーネル最適化を評価するTPUネイティブベンチマークスイートで、50のJAXワークロードを含む。研究により、TPUドキュメントによる条件付けが正解率を大幅に向上させ、探索構造がさらなる高速化をもたらすことが示されたが、高品質なTPUカーネル最適化には依然として課題がある。

ソースarXiv AI著者: Arya Tschand, Charles Hong, Julian Walker, Nina Cai, Shangkun Wang, Suvinay Subramanian, Sundar Dev, Vijay Janapa Reddi, Amir Yazdanbakhsh, Sethu Sankaran

JAXBenchは、Google Cloud TPU上でのAI生成カーネル最適化を評価するための新しいTPUネイティブベンチマークスイートです。このスイートは50のJAXワークロードで構成されており、これらは実際に関連性が高く、最適化の余地があります。そのうち17の本番ML演算子は、Llama-3.1、DeepSeek-V3、Mixtral、Mamba-2、AlphaFold2などのMaxTextライブラリのアーキテクチャから抽出され、残りの33演算子はKernelBenchから変換され、正確性が検証され、高いTPU v6e MXU利用率を達成する新しい問題サイズが設定されています。17の本番演算子のうち8つには、Tokamaxライブラリから手動最適化されたPallasカーネルが付属し、ブロックサイズが調整されて専門家の上限ベースラインを確立しています。

研究チームは、JAXBench用の候補Pallasカーネルを生成する4つのフィードバックベース手法を評価しました。Gemini 3 Flashを使用した完全なスイート全体で、ターゲット固有のコンテキストがモデル規模よりも重要であることがわかりました。特に、ドキュメントが乏しいDSLであるPallasでは顕著です。厳選されたTPUドキュメントによる条件付けにより、サンプルあたりの正解率が5.8%から37.3%に向上し、50のベンチマークのうち48を解決し、幾何平均で1.28倍の高速化を達成しました。正解率が達成されると、探索構造が大きな利益をもたらし、AutocompのビームサーチパイプラインはXLAに対して1.36倍の幾何平均高速化を達成しました。8つの手調整カーネルでは、AutocompはXLAに対して1.60倍の幾何平均高速化を達成し、Tokamaxの2.08倍という上限の大部分を回復しましたが、特殊なページ化アテンションやラグドアテンション演算子では依然として劣っています。

高品質なTPUカーネル最適化は依然として困難な課題ですが、研究チームはJAXBenchベンチマーク、評価ハーネス、ベースライン結果を公開し、オープンソースコミュニティの貢献を支援しています。これらのツールは、将来の研究に共有の最適化舞台を提供し、TPUカーネルの自動最適化の進展を加速することが期待されます。研究では、ドキュメントが乏しいDSLにおいて、対象固有のコンテキストを提供することがモデル規模の拡大よりも重要であることが示され、この発見は他の類似分野の自動コード生成にも指針を与えるものです。