jax.numpy의 약속은 단순해. NumPy의 API를 거의 그대로 제공하되 연산은 XLA 위에서 실행한다는 거야. 그래서 NumPy 코드에서 JAX로 옮기는 첫 단계는 import 한 줄을 바꾸는 것부터 시작해.
import numpy as np
import jax.numpy as jnp
a = np.array([1, 2, 3])
b = jnp.array([1, 2, 3])
print(np.sum(a)) # 6
print(jnp.sum(b)) # 6
# 거의 모든 게 그대로 — 함수 이름, 시그니처, 동작
np.zeros((3, 3)); jnp.zeros((3, 3))
np.linspace(0,1,5); jnp.linspace(0,1,5)
np.dot(a, a); jnp.dot(b, b)
함수 이름과 시그니처는 비슷하지만 실행 모델에는 중요한 차이가 있어.
- 불변 배열:
jnp.array는 직접 변경할 수 없어.a[0] = 5대신a = a.at[0].set(5)를 사용해. - 자동 장치 배치: 배열을 만들면 사용 가능한 가속기에 배치된다.
- float32 기본 dtype: NumPy의 기본값인 float64와 다르다.
- 별도의 난수 API:
jax.random과 명시적인 키를 사용한다. Track 8에서 자세히 배운다.
⚠️ 첫 번째 함정
NumPy 코드를 그대로 옮기면 인덱스 할당에서 거의 틀림없이 막혀. .at[].set()을 이용한 함수형 갱신 패턴을 익혀야 해.