JAX에서 가장 중요한 규칙 하나를 꼽으면 함수가 순수해야 한다는 거야. 순수 함수에는 두 가지 조건이 있어.
- 같은 입력은 언제나 같은 출력을 만든다. 결과가 결정적이어야 해.
- 부수 효과가 없다. 외부 상태를 읽거나 쓰지 않으며, 결과는 반환값으로만 전달해.
다음 예제를 보면 차이가 분명해.
# PURE
def add(a, b):
return a + b
def normalize(x):
return (x - x.mean()) / x.std()
# IMPURE
counter = 0
def step():
global counter
counter += 1 # ❌ side effect
return counter
cache = {}
def lookup(k):
if k in cache: # ❌ external state 읽기
return cache[k]
cache[k] = expensive(k) # ❌ external state 쓰기
return cache[k]
JAX가 순수성을 요구하는 이유는 jit, grad, vmap, pmap 같은 변환이 함수를 추적해 작동하기 때문이야. 추적은 함수를 한 번 실행하면서 어떤 연산이 어떤 순서로 일어나는지 기록하는 과정이야. JAX는 그 기록을 컴파일하고 미분하고 벡터화해.
함수 안에 부수 효과가 있다면 그 동작은 추적할 때 한 번만 일어날 수 있어. 컴파일된 함수를 다시 호출할 때는 반복되지 않아.
import jax
call_count = 0
@jax.jit
def f(x):
global call_count
call_count += 1 # 이게 정확히 한 번만 실행됨 (trace 때)
return x * 2
f(1.0) # call_count == 1 (trace + 실행)
f(2.0) # call_count == 1 (trace 안 함, cached)
f(3.0) # call_count == 1 (cached)
print(call_count) # 1, not 3
이 문제는 오류 없이 잘못된 동작을 만들 수 있어서 특히 찾기 어려워. 함수를 세 번 호출했는데 카운터가 한 번만 늘었다면, 원인은 함수가 순수하지 않기 때문이야.
🌿 함수형 사고로 전환하기
JAX가 Python의 동적이고 변경 가능한 성격을 제한하는 데는 이유가 있어. XLA가 함수를 미리 컴파일하려면 계산이 결정적이고 부수 효과가 없어야 해. "함수는 입력을 출력으로 대응시키는 수학적 사상"이라는 정의로 돌아가는 셈이지. 처음에는 답답할 수 있지만, 익숙해지면 상태의 흐름이 드러나 더 깨끗한 모델을 만들 수 있어.
상태가 필요하다면 숨기지 말고 함수의 입력과 출력으로 명시해.
# PyTorch 식 (JAX 아님)
class Counter:
def __init__(self): self.n = 0
def step(self): self.n += 1; return self.n
# JAX 식
def step(state):
return state + 1, state + 1
state = 0
state, output = step(state) # 항상 explicit
state, output = step(state)
상태를 함수의 입력과 출력으로 분리하는 패턴은 JAX 곳곳에서 반복돼. Optax와 Flax NNX, 학습 루프도 같은 구조를 사용해.