때로는 자동 미분이 만든 그래디언트를 그대로 사용하기 어려워. 너무 느리거나 수치적으로 불안정할 수도 있고, 사용자 정의 의미를 부여해야 할 수도 있어. jax.custom_vjp / jax.custom_jvp로 직접 정의해.
예: numerically stable softmax-cross-entropy
import jax
import jax.numpy as jnp
# 단순 구현 — log(softmax(x)) 가 underflow 위험
def naive_loss(logits, target):
p = jax.nn.softmax(logits)
return -jnp.sum(target * jnp.log(p + 1e-10))
# JAX 의 jax.nn.log_softmax — 안정
def stable_loss(logits, target):
log_p = jax.nn.log_softmax(logits)
return -jnp.sum(target * log_p)
여기까지는 표준 함수로 충분해. 사용자 정의 그래디언트가 필요한 대표 사례로 양자화의 straight-through estimator(STE)를 살펴보자.
@jax.custom_vjp
def ste_round(x):
'''forward 에선 round, backward 에선 identity'''
return jnp.round(x)
def ste_round_fwd(x):
return jnp.round(x), x # residual = x (저장)
def ste_round_bwd(x_residual, grad_out):
return (grad_out,) # grad 가 그대로 통과
ste_round.defvjp(ste_round_fwd, ste_round_bwd)
# 사용
def loss(x):
return jnp.sum(ste_round(x) ** 2)
grad = jax.grad(loss)(jnp.array([1.5, 2.7, 3.1]))
print(grad) # gradient 가 흐름 (round 가 미분 0 인데도)
round는 거의 모든 점에서 미분값이 0이야. 그러나 STE는 identity인 것처럼 그래디언트를 통과시켜 학습할 수 있게 해.
custom_jvp, 순방향 모드
@jax.custom_jvp
def square(x):
return x ** 2
@square.defjvp
def square_jvp(primals, tangents):
x, = primals
dx, = tangents
primal_out = x ** 2
tangent_out = 2 * x * dx # 우리가 정의한 forward derivative
return primal_out, tangent_out
대부분의 사용 사례에는 custom_vjp(reverse mode)를 써. custom_jvp는 순방향 모드가 더 효율적인 경우에 알맞아.
실전 예: 외부 solver의 그래디언트
@jax.custom_vjp
def ode_solve(initial_state, t_final, params):
'''black-box solver 호출'''
return some_ode_solver(initial_state, t_final, params)
def ode_solve_fwd(initial_state, t_final, params):
final_state = some_ode_solver(initial_state, t_final, params)
return final_state, (initial_state, t_final, params, final_state)
def ode_solve_bwd(residuals, grad_final):
'''adjoint method 로 gradient 계산'''
initial_state, t_final, params, final_state = residuals
# 별도 ODE 를 풀어 gradient 얻음
grad_initial, grad_params = solve_adjoint(...)
return (grad_initial, None, grad_params)
ode_solve.defvjp(ode_solve_fwd, ode_solve_bwd)
과학 계산에서 솔버 자체를 미분 가능하게 만드는 표준 패턴이야. Track 13의 Diffrax가 이 과정을 깔끔하게 감싸 줘.
🔬 언제 custom 그래디언트?
(1) 수치 안정성, log-sum-exp, log(softmax) 같은 수치 안정 표현해. (2) 미분 불가능한 round, argmax, sample 연산에는 STE 패턴을 적용할 수 있어. STE 패턴. (3) 외부 solver, ODE, fixed-point iteration, 계층으로 사용하는 옵티마이저. adjoint 메서드. (4) 새로운 정규화, weight constraint, manifold projection. 일반적으로 코드의 99%는 표준 자동 미분이면 충분해. 분명한 이유가 있을 때만 사용자 정의 그래디언트를 써.
이 패턴은 JAX의 고급 영역이야. 하지만 한 번 익히면 연구 코드의 표현력이 크게 넓어져.