jax-ml/jax-triton
jax-triton contains integrations between JAX and OpenAI Triton
何を解決するか
TritonカーネルをJAXプログラムに統合する方法を提供し、開発者がTritonのPythonベースの言語を使って高性能なカスタムGPUカーネルを記述し、JAXの関数型エコシステム内(jax.jitコンパイル関数内を含む)でシームレスに実行できるようにします。
動作方法
このライブラリは triton_call 関数を導入し、ブリッジとして機能します。JAX配列をTritonカーネルの引数として渡すことが可能になります。また、jax.new_ref で作成されたJAX Ref オブジェクトのインプレース変更もサポートしており、新しい出力配列を割り当てることなくデータを変更できます。
対象ユーザー
カスタムTritonカーネルのパフォーマンスを必要とするが、全体のモデルアーキテクチャや高レベルのオーケストレーションにはJAXを使い続けたい開発者や研究者。
特徴
- JAX配列からTritonカーネルを呼び出すサポート。
- 効率的な実行のための
jax.jitとの互換性。 - インプレース変更のためのJAX
Refを通じた入出力パラメータのサポート。 - Gluon方言との統合。
関連
- Dispatch
- プロジェクト
- プロジェクト
- プロジェクト
- プロジェクト