JAX의 기본 dtype은 float32이고 NumPy의 기본값은 float64야. 이 차이를 모르고 코드를 옮기면 같은 수식을 계산했는데도 결과가 달라질 수 있어.
import numpy as np
import jax.numpy as jnp
a_np = np.array([1.0, 2.0, 3.0])
a_jax = jnp.array([1.0, 2.0, 3.0])
print(a_np.dtype) # float64
print(a_jax.dtype) # float32 (!)
float32가 기본인 이유는 가속기의 특성과 관련이 있어. GPU와 TPU에서 float32 연산은 보통 float64보다 2~30배 빠르고, 대부분의 머신러닝 작업은 float32 정밀도로도 충분히 수렴해.
import jax
jax.config.update("jax_enable_x64", True)
import jax.numpy as jnp
a = jnp.array([1.0, 2.0])
print(a.dtype) # float64 (이제 됨)
⚠️ 정밀도 함정
NumPy에서 float64로 계산하던 코드를 JAX로 옮기면 같은 결과가 나오지 않을 수 있어. 학습이 수렴하지 않거나 NaN이 나타나면 dtype부터 확인해.