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 객체의 인플레이스(mutate)를 지원하여, 새로운 출력 배열을 할당하지 않고도 데이터를 수정할 수 있습니다.

대상 사용자

커스텀 Triton 커널의 성능을 필요로 하지만, 전체 모델 아키텍처와 고수준 오케스트레이션에는 JAX를 유지하고 싶은 개발자 및 연구자.

주요 기능

  • JAX 배열에서 Triton 커널 호출 지원.
  • 효율적인 실행을 위한 jax.jit 호환성.
  • JAX Ref를 통한 인-아웃 파라미터 지원으로 인플레이스 변경 가능.
  • Gluon 다이얼렉트와의 통합.

관련

  • Dispatch
  • 프로젝트
  • 프로젝트
  • 프로젝트
  • 프로젝트