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
  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト