vmap은 Python loop를 batch 연산으로 바꾼다
벡터 하나를 행렬에 곱하는 함수를 여러 입력에 적용해 보자.
1
2
3
4
5
| matrix = random.normal(key, (150, 100))
batched_x = random.normal(key, (10, 100))
def apply_matrix(vector):
return jnp.dot(matrix, vector)
|
Python loop로 쌓으면 다음과 같다.
1
2
| def naively_batched_apply_matrix(vectors):
return jnp.stack([apply_matrix(v) for v in vectors])
|
이 문제는 직접 matrix-matrix product로 바꿀 수 있다.
1
2
3
| @jit
def manually_batched_apply_matrix(vectors):
return jnp.dot(vectors, matrix.T)
|
연산이 복잡해 직접 batch 식으로 다시 쓰기 어렵다면 vmap으로 입력 축에 함수를 적용한다.
1
2
3
| @jit
def vmap_batched_apply_matrix(vectors):
return vmap(apply_matrix)(vectors)
|
원문 측정은 단순 loop 약 7.68ms, 직접 batch 약 68.8µs, vmap 약 105µs였다. 직접 행렬곱이 가능한 예에서는 손으로 쓴 batch가 빨랐지만, vmap의 가치는 복잡한 함수를 자동으로 벡터화할 때 드러난다.
JAX를 평가하는 가장 작은 실험 순서는 다음과 같다.
- dtype과 실행 장치를 출력해 기록한다.
- 입력을 장치에 미리 두고 전송 시간과 계산 시간을 분리한다.
block_until_ready()로 비동기 계산을 기다린다.- 첫 JIT 호출과 이후 반복 호출을 나눠 잰다.
grad 결과를 작은 입력의 수치 미분과 비교한다.- Python loop와
vmap 결과 shape, 값이 같은지 먼저 확인한 뒤 시간을 잰다.
결론적으로 JAX의 장점은 “NumPy 한 줄보다 항상 빠르다”가 아니라, jit, grad, vmap 같은 변환을 조합해 반복되는 수치 함수를 컴파일하고 미분하고 벡터화할 수 있다는 데 있다.