본문 바로가기
C.W.K.
Stream
Lesson 01 of 05 · published

순수 함수란 무엇이고 JAX는 왜 요구할까

~9 min · purity, jax, tutorial

Level 0호기심
0 XP0/73 lessons0/17 achievements
0/100 XP to next level100 XP to go0% complete

JAX에서 가장 중요한 규칙 하나를 꼽으면 함수가 순수해야 한다는 거야. 순수 함수에는 두 가지 조건이 있어.

  1. 같은 입력은 언제나 같은 출력을 만든다. 결과가 결정적이어야 해.
  2. 부수 효과가 없다. 외부 상태를 읽거나 쓰지 않으며, 결과는 반환값으로만 전달해.

다음 예제를 보면 차이가 분명해.

# 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, 학습 루프도 같은 구조를 사용해.

Code

import jax.numpy as jnp

# PURE: output depends only on input, no side effects
def pure_fn(x):
    return jnp.sum(x ** 2)

# IMPURE: depends on global state
scale = 2.0
def impure_global(x):
    return jnp.sum(x ** 2) * scale  # Reads global variable!

# IMPURE: has side effects
results = []
def impure_sideeffect(x):
    result = jnp.sum(x ** 2)
    results.append(result)  # Side effect: modifies external list!
    return result

# IMPURE: mutates input
def impure_mutation(x):
    x[0] = 0  # Side effect: modifies input! (also fails in JAX)
    return jnp.sum(x)
import jax
import jax.numpy as jnp

# Demonstration: global state is captured at trace time
multiplier = 2.0

@jax.jit
def buggy_multiply(x):
    return x * multiplier

print(buggy_multiply(jnp.array(3.0)))  # 6.0

multiplier = 10.0  # Change the global
print(buggy_multiply(jnp.array(3.0)))  # Still 6.0! JIT cached the old value
@jax.jit
def correct_multiply(x, multiplier):
    return x * multiplier

print(correct_multiply(jnp.array(3.0), 2.0))   # 6.0
print(correct_multiply(jnp.array(3.0), 10.0))  # 30.0 — correct!

External links

Exercise

순수하지 않은 함수 세 개를 작성해. (1) 전역 값 읽기, (2) 인자로 받은 리스트 변경, (3) print 호출을 각각 포함하게 해. 각 함수를 jit으로 감싼 뒤 서로 다른 입력으로 두 번 호출하고 결과를 정리해. JAX는 언제나 오류를 내는 것이 아니라 이전 값을 캐시한 채 실행하기도 해. 그 미묘함이 순수성이 중요한 이유야.

Progress

Progress is local-only — sign in to sync across devices.
이 페이지에서 버그를 발견하셨거나 피드백이 있으세요?문제 신고

댓글 0

🔔 답글 알림 (로그인 필요)
로그인댓글을 남기려면 로그인해 주세요.

아직 댓글이 없어요. 첫 댓글을 남겨보세요.