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
- 项目
- 项目
- 项目
- 项目