실전에서 자주 쓰는 grad 변형 세 가지를 살펴보자.
1. argnums, 어느 인자에 대한 미분이냐
def f(a, b, c):
return a ** 2 + b * c
# default — argnums=0
df_da = jax.grad(f)(2.0, 3.0, 4.0) # 4.0
# 특정 인자
df_db = jax.grad(f, argnums=1)(2.0, 3.0, 4.0) # 4.0
df_dc = jax.grad(f, argnums=2)(2.0, 3.0, 4.0) # 3.0
# 여러 인자 동시
g_a, g_b = jax.grad(f, argnums=(0, 1))(2.0, 3.0, 4.0)
2. value_and_grad, 값과 그래디언트 같이
학습 루프에서는 손실 값과 그래디언트가 모두 필요해. jax.grad만 부르면 손실 값을 알기 위해 순전파를 한 번 더 돌려야 해서 낭비야. value_and_grad는 한 번의 순전파와 역전파로 둘 다 얻어.
def loss_fn(params, x, y):
pred = jnp.dot(x, params)
return jnp.mean((pred - y) ** 2)
# 비효율
loss = loss_fn(params, x, y)
grads = jax.grad(loss_fn)(params, x, y) # forward 두 번
# 효율
loss, grads = jax.value_and_grad(loss_fn)(params, x, y)
학습 코드에서는 거의 항상 value_and_grad를 써. JAX의 가장 흔한 무료 성능 향상.
3. has_aux, 추가 정보 함께 반환
학습할 때는 손실 외에 지표(accuracy, perplexity 등)도 함께 계산하고 싶을 수 있어. 그런데 grad는 스칼라 출력만 받음. 해결, has_aux=True로 두 번째 출력은 "auxiliary, 그래디언트 계산 안 함":
def loss_and_metrics(params, x, y):
pred = jnp.dot(x, params)
loss = jnp.mean((pred - y) ** 2)
metrics = {
"mae": jnp.mean(jnp.abs(pred - y)),
"max_err": jnp.max(jnp.abs(pred - y)),
}
return loss, metrics # tuple
(loss, metrics), grads = jax.value_and_grad(
loss_and_metrics, has_aux=True
)(params, x, y)
print(f"loss: {loss}, mae: {metrics['mae']}")
has_aux=True로, 첫 번째 출력에 대해서만 미분, 두 번째는 그대로 통과해.
전형적 학습 스텝 패턴:
@jax.jit
def train_step(state, batch):
'''state: {"params": ..., "step": ...}, batch: (x, y)'''
def loss_and_metrics(params):
x, y = batch
pred = model_apply(params, x)
loss = compute_loss(pred, y)
metrics = {"acc": accuracy(pred, y)}
return loss, metrics
(loss, metrics), grads = jax.value_and_grad(
loss_and_metrics, has_aux=True
)(state["params"])
new_params = jax.tree.map(
lambda p, g: p - 0.001 * g, state["params"], grads
)
new_state = {"params": new_params, "step": state["step"] + 1}
return new_state, loss, metrics
💡 항상 value_and_grad
학습 코드에서는 jax.grad만 따로 쓰는 일이 드물어. 손실과 그래디언트를 함께 구하는 jax.value_and_grad를 기본으로 기억하고, 지표까지 반환할 때는 has_aux=True를 사용해. 두 패턴이 90%의 학습 코드를 포괄해.