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
- 專案
- 專案
- 專案
- 專案