JAX 오류 메시지가 처음엔 외계어처럼 보이는데, 패턴 익히면 빠르게 해석할 수 있어.
1. ConcretizationTypeError
jax.errors.ConcretizationTypeError: Abstract tracer value encountered
where concrete value is expected: Traced<ShapedArray(float32[])>...
의미: 추적된 값을 Python의 구체적인 값처럼 다루려 함 (예: if x > 0:, .item(), int(x)).
@jax.jit
def f(x):
if x > 0: # ❌ Python if on Tracer
return x
else:
return -x
해결:
@jax.jit
def f(x):
return jnp.where(x > 0, x, -x) # ✅
# 또는
@jax.jit
def f(x):
return jax.lax.cond(x > 0, lambda: x, lambda: -x)
2. TracerArrayConversionError
TracerArrayConversionError: The numpy.ndarray conversion method
was called on Traced<...>
의미: 추적된 값을 NumPy 함수에 넣음. NumPy가 실제 배열을 원하는데 Tracer를 받은 거야.
@jax.jit
def f(x):
return np.sin(x) # ❌ np 대신 jnp
해결: jnp로.
3. 캐시된 컴파일, 부수 효과가 다시 실행되지 않는 증상
오류는 나지 않아. 결과가 이상해.
global_counter = 0
@jax.jit
def f(x):
global global_counter
global_counter += 1
return x * 2
f(1.0); f(2.0); f(3.0)
print(global_counter) # 1 (3 아님!)
해결: 상태를 함수 인자로 노출. def f(x, counter): return x*2, counter+1.
4. NonHashableStaticArgumentsError
TypeError: unhashable type: 'list'
의미: static_argnames로 표시된 인자가 hashable이어야 함 (list는 안 되고 tuple만 돼).
@partial(jax.jit, static_argnames=("shape",))
def f(x, shape):
return jnp.zeros(shape) + x
f(1.0, [3, 3]) # ❌ list
f(1.0, (3, 3)) # ✅ tuple
5. ShapeMismatchError
jax 에서 두 호출의 shape 이 달라서 recompile 됨
의미: 같은 jit 함수를 다른 shape으로 호출해 → 매번 새로 컴파일해. 의도한 동작이면 괜찮아, 아니면 static_argnames로 분리하거나 패딩하면 돼.
🔍 디버깅 순서
(1) 오류 메시지의 첫 줄만 봐도 범주를 잡을 수 있어. ConcretizationTypeError = 추적된 값으로 제어 흐름을 결정했다는 뜻. TracerArrayConversion = 잘못된 이름 공간. (2) jit 빼고 즉시 실행 모드로 돌려 봐, 그러면 일반 Python 오류처럼 동작해서 원인 찾기 쉬워. (3) jax.disable_jit() 컨텍스트 매니저로 감싸도 돼.
실용 팁: 새 함수를 처음 짤 땐 jit 안 붙이고 즉시 실행 모드로 한 번 돌려 정상 동작을 확인한 뒤 jit을 추가해. 컴파일 오류는 까다로워서, 한 번에 다 잡으려고 하면 시간 낭비야.