Quickstart: How to think in JAX

Quickstart: How to think in JAX

JAX는 NumPy 스타일의 배열 지향 수치 연산 라이브러리인데, 자동 미분과 JIT 컴파일이 더해져 고성능 머신러닝 연구에 쓰여요. 기존 NumPy 코드에 익숙하다면 진입 장벽이 크지 않아요. JAX의 핵심 기능을 빠르게 짚어 보면서 어떤 특징이 있는지 살펴볼게요.

출처: Quickstart: How to think in JAX

JAX가 제공하는 것

JAX는 CPU·GPU·TPU를 로컬이든 분산이든 같은 인터페이스로 다루게 해 줘요.

  • JAX는 CPU·GPU·TPU에서 실행되는 계산에 NumPy 같은 통일된 인터페이스를 제공해요. 로컬이든 분산이든 마찬가지예요.
  • Open XLA(오픈소스 머신러닝 컴파일러 생태계) 기반의 Just-In-Time(JIT) 컴파일이 내장돼 있어요.
  • JAX 함수는 자동 미분 변환을 통해 그라디언트를 효율적으로 계산할 수 있어요.
  • JAX 함수는 자동으로 벡터화되어 입력 배치를 나타내는 배열 위에 효율적으로 매핑될 수 있어요.

설치

JAX는 CPU용으로 Linux, Windows, macOS에서 PyPI에서 바로 설치할 수 있어요.

pip install jax

NVIDIA GPU용이라면:

pip install -U "jax[cuda13]"

플랫폼별 상세 설치 방법은 Installation 문서를 참고해요.

JAX vs NumPy

핵심 개념을 정리하면 이래요.

  • JAX는 편의를 위해 NumPy 스타일 인터페이스를 제공해요.
  • 덕 타이핑(duck-typing) 덕분에 JAX 배열은 NumPy 배열의 드롭인 대체물로 자주 쓰일 수 있어요.
  • NumPy 배열과 달리 JAX 배열은 **항상 불변(immutable)**이에요.

NumPy는 수치 데이터를 다루는 잘 알려진 강력한 API예요. 편의를 위해 JAX는 NumPy API를 밀접하게 따르는 jax.numpy를 제공해요. 보통 jnp 별칭으로 import 해요:

import jax.numpy as jnp

이 import만으로 평범한 NumPy 프로그램처럼 JAX를 쓸 수 있어요. NumPy 스타일 배열 생성 함수, 파이썬 함수와 연산자, 배열 속성·메서드를 모두 사용할 수 있어요:

import matplotlib.pyplot as plt

x_jnp = jnp.linspace(0, 10, 1000)
y_jnp = 2 * jnp.sin(x_jnp) * jnp.cos(x_jnp)
plt.plot(x_jnp, y_jnp);

npjnp로 바꾼 것만 빼면 코드 블록은 NumPy에서 기대하는 것과 완전히 동일하고, 결과도 같아요. JAX 배열은 플로팅 같은 곳에서 NumPy 배열 자리에 그대로 쓰일 수 있어요.

배열 자체는 서로 다른 파이썬 타입으로 구현돼요:

import numpy as np
import jax.numpy as jnp

x_np = np.linspace(0, 10, 1000)
x_jnp = jnp.linspace(0, 10, 1000)

type(x_np)
numpy.ndarray

type(x_jnp)
jaxlib._jax.ArrayImpl

파이썬의 덕 타이핑 덕분에 JAX 배열과 NumPy 배열은 많은 곳에서 서로 바꿔 쓸 수 있어요. 그런데 JAX와 NumPy 배열 사이에는 중요한 차이 하나가 있어요. 바로 JAX 배열은 불변이라, 한 번 만들면 내용을 바꿀 수 없다는 거예요.

NumPy에서 배열을 변경하는 예시를 보면:

# NumPy: mutable arrays

x = np.arange(10)
x[0] = 10
print(x)

[10 1 2 3 4 5 6 7 8 9]

이렇게 원소를 바로 바꿔 쓸 수 있죠. JAX에서는 이런 in-place 변경이 허용되지 않아요. 대신 함수형 방식으로 새 값을 만들어야 해요.

더 알아보기