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
OpenAI Triton 1.0 출시OpenAI는 연구자들이 깊은 CUDA 전문 지식 없이도 매우 효율적인 GPU 커널을 작성할 수 있게 해주는 오픈 소스 Python 스타일 언어이자 컴파일러인 Triton 1.0을 출시했습니다.
- 프로젝트
triton-lang/tritonCUDA보다 생산성이 높고, 다른 DSL보다 더 유연하며, 고도로 효율적인 사용자 정의 딥러닝 프리미티브를 작성할 수 있는 언어와 컴파일러.
- 프로젝트
ByteDance-Seed/Triton-distributedTriton-distributed는 ByteDance-Seed가 오픈 소스로 공개한 컴파일러로, Triton 언어를 확장하여 GPU 연산과 통신을 오버랩하기 위한 프리미티브를 제공합니다. NVIDIA 및 AMD GPU에서 고성능 분산 커널(GEMM, MoE, flash-decode, AllToAll)을 구현하여 수동으로 튜닝된 라이브러리에 필적하는 속도 향상을 제공합니다. 소스 또는 사전 빌드된 pip wheel을 통해 설치하고, @triton_dist.jit을 사용하여 커널을 작성하며, 제공된 저수준 통신 API를 사용할 수 있습니다. 이 프로젝트는 활발히 업데이트되고 있으며(최신 뉴스 2026년 9월), MIT 라이선스로 배포됩니다.
- 프로젝트
- 프로젝트
BobMcDear/attorch커스텀 딥러닝 연산을 개발하기 위해 읽기 쉽고 수정이 용이한 신경망 레이어 컬렉션을 제공하는 PyTorch nn 모듈의 Triton 기반 서브셋입니다.