Automatic vectorization
Automatic vectorization
배치 데이터를 처리할 때 같은 함수를 여러 입력에 반복 적용하는 일이 많은데, JAX는 이를 자동화해 줘요. jax.vmap() 변환은 함수의 벡터화된 구현을 자동으로 만들어 줘서, 수동으로 배치 로직을 다시 짜는 수고를 덜어요. jit과 마찬가지로 함수를 트레이싱하는 방식으로 동작해요.
수동 벡터화
벡터화된 구현을 만드는 건 특히 어렵진 않지만, 함수가 인덱스·축·입력의 다른 부분을 다루는 방식을 바꿔야 해요. 함수의 복잡도가 커질수록 이런 재구현은 지저분하고 오류가 나기 쉬워져요. 다행히 JAX는 다른 방법을 제공해요.
자동 벡터화
JAX에서 jax.vmap() 변환은 함수의 벡터화된 구현을 자동으로 생성해요:
auto_batch_convolve = jax.vmap(convolve)
auto_batch_convolve(xs, ws)
Array([[11., 20., 29.],
[11., 20., 29.]], dtype=float32)
이는 jax.jit()과 비슷하게 함수를 트레이싱하고, 각 입력의 앞쪽에 배치 축을 자동으로 추가해서 동작해요.
배치 차원이 첫 번째가 아니라면, in_axes와 out_axes 인자로 입력·출력에서 배치 차원의 위치를 지정할 수 있어요. 모든 입력·출력의 배치 축이 같다면 정수 하나로 지정하고, 아니면 리스트로 지정해요.
jax.vmap()은 인자 중 하나만 배치되는 경우도 지원해요. 예를 들어 단일 가중치 세트 w와 벡터 배치 x의 컨볼루션을 원한다면, batch가 아닌 인자를 그대로 두고 쓰면 돼요.
변환 조합하기
모든 JAX 변환이 그렇듯 jax.jit()과 jax.vmap()은 서로 조합 가능하게 설계돼 있어요. vmapped 함수를 jit으로 감싸거나 jitted 함수를 vmap으로 감쌀 수 있고, 둘 다 올바르게 동작해요.
jitted_batch_convolve = jax.jit(auto_batch_convolve)
jitted_batch_convolve(xs, ws)
Array([[11., 20., 29.],
[11., 20., 29.]], dtype=float32)
vmap을 자기 자신과 조합하는 것도 유용해요. 예를 들어 두 겹의 vmap으로 함수의 쌍별(pairwise) 평가를 간결하게 표현할 수 있어요:
def pairwise(f, xs):
return jax.vmap(lambda x: jax.vmap(lambda y: f(x, y))(xs))(xs)
정리
vmap은 함수를 자동 벡터화해 배치 입력을 한 번에 처리해요.in_axes/out_axes로 배치 차원 위치를 조정할 수 있어요.- jit·grad 등 다른 변환과 자유롭게 조합돼요.
- 단일 배치 인자, 중첩 vmap 같은 다양한 시나리오를 지원해요.
더 알아보기
- Key concepts — JAX 변환 기본
- Just-in-time compilation — jax.jit 으로 가속
- Automatic differentiation — jax.grad 로 미분