jax-ml/jax
Composable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more
해결하는 문제
JAX는 고성능 수치 계산 및 대규모 기계 학습을 위한 시스템을 제공합니다. 개발자는 NumPy와 유사한 코드를 작성할 수 있으며, 자동 미분, 속도 향상을 위한 컴파일, GPU 및 TPU와 같은 여러 하드웨어 가속기로의 확장이 자동으로 이루어집니다.
작동 방식
JAX는 XLA(가속된 선형 대수)를 사용하여 Python 및 NumPy 함수를 최적화된 머신 코드로 컴파일합니다. 이는 조합 가능한 함수 변환 시스템으로 작동합니다:
jax.grad: 자동 미분(역방향 모드 및 순방향 모드 모두 지원)을 사용하여 Python 및 NumPy 함수의 기울기를 계산합니다.jax.jit: XLA를 사용하여 함수를 끝에서 끝까지 컴파일하여 실행 속도를 향상시킵니다.jax.vmap: 함수를 자동으로 벡터화하여 배열 축을 따라 매핑하여 수동 배치 루프를 제거합니다.
대상 사용자
고성능 수치 계산 및 대규모 기계 학습에 종사하는 연구자 및 개발자에게 설계되었으며, 효율적인 기울기 계산과 수천 개의 장치에 걸쳐 계산을 확장할 수 있는 능력이 필요한 사용자에게 적합합니다.
주요 특징
- 조합 가능한 변환:
grad,jit,vmap를 임의의 순서로 조합하여 매우 최적화된 함수를 생성할 수 있습니다. - 하드웨어 가속: XLA를 통해 NVIDIA GPU, Google TPU 및 기타 가속기와의 네이티브 지원을 제공합니다.
- 유연한 미분: 루프, 분기, 재귀, 클로저를 임의의 차수까지 미분할 수 있습니다.
- 확장 옵션: 컴파일러 기반 자동 병렬화, 명시적 샤딩, 수동 디바이스 프로그래밍의 세 가지 확장 모드를 제공합니다.
관련
- 프로젝트
- 프로젝트
- 프로젝트
- 프로젝트
- 프로젝트