프로젝트 소개

JAX는 Python 및 NumPy 프로그램의 구성 가능한 함수 변환 시스템을 제공합니다. 이는 컴파일을 위해 XLA(Accelerated Linear Algebra)를 사용하여 GPU 및 TPU를 포함한 하드웨어 가속기 전반에서 수치 계산을 확장하도록 설계되었습니다. 주요 기능은 다음과 같습니다: - 자동 미분: `jax.grad`를 사용하여 Python 루프, 분기, 재귀 및 클로저를 통한 역방향 및 순방향 미분을 지원하며, 고계 도함수를 구할 수 있습니다. - 적시(JIT) 컴파일: `jax.jit`를 사용하여 순수 함수를 XLA를 통해 컴파일함으로써 가속기에서의 실행 성능을 향상시킵니다. - 자동 벡터화: `jax.vmap`을 사용하여 루프를 기본 연산으로 밀어 넣어 배열 축을 따라 함수를 매핑하며, 수동으로 배치 차원을 관리할 필요를 없애줍니다. - 확장 및 병렬화: 컴파일러 기반 자동 병렬화, 명시적 샤딩, 명시적 콜렉티브를 이용한 수동 장치별 프로그래밍 등 다양한 확장 모드를 지원합니다. JAX는 Linux, macOS, Windows를 포함한 여러 플랫폼과 호환되며 CPU, NVIDIA GPU, Google TPU, AMD GPU를 지원하고 Apple 및 Intel GPU에 대한 실험적 지원을 제공합니다.