포스트

JAX가 NumPy보다 느리게 나오는 이유: jit, grad, vmap 벤치마크 함정

JAX를 NumPy와 한 번 dot해 보고 느리다고 결론 내리면 안 된다. 데이터가 어느 장치에 있는지, 첫 컴파일 시간을 포함했는지, 비동기 계산을 기다렸는지가 맞아야 비교가 성립한다.

먼저 같은 dtype, shape, 값에서 결과가 일치하는지 확인하고, compile, transfer, steady-state 실행을 분리한다. JAX의 핵심은 단일 연산 승부가 아니라 순수한 수치 함수를 jit, grad, vmap으로 변환해 조합하는 데 있다.

JAX는 어떤 관점으로 읽어야 할까?

JAX는 NumPy와 닮은 API에 자동 미분과 컴파일, 벡터화, 여러 가속기 병렬화를 위한 변환을 더한다. 핵심은 별도의 tensor 문법을 외우는 것보다 함수에 변환을 적용한다는 관점이다.

jnp.dot 하나로 속도를 판단하면 안 되는 이유

원문 quickstart 코드는 다음 import로 시작한다.

1
2
3
4
import numpy as np

import jax.numpy as jnp
from jax import device_put, grad, jit, random, vmap

원문에는 경고를 피하려고 플랫폼을 CPU로 고정하는 줄도 있었다.

1
2
import jax
jax.config.update('jax_platform_name', 'cpu')

이 설정을 켠 채 결과를 “GPU 연산 시간”이라고 부르면 비교 자체가 틀린다. CPU 실험인지 가속기 실험인지 먼저 정하고, 실제 실행 장치와 설정을 기록해야 한다.

JAX 연산은 비동기로 실행될 수 있으므로 시간을 잴 때 결과가 끝날 때까지 기다린다.

1
2
3
4
5
key = random.PRNGKey(0)
size = 3000

x = random.normal(key, (size, size), dtype=jnp.float32)
%timeit jnp.dot(x, x.T).block_until_ready()

NumPy 배열을 매번 JAX 연산에 넘기면 장치 이동 비용이 측정에 섞일 수 있다. 원문은 device_put으로 먼저 올려둔 경우도 비교했다.

1
2
3
x = np.random.normal(size=(size, size)).astype(np.float32)
x = device_put(x)
%timeit jnp.dot(x, x.T).block_until_ready()

원문 결과에서는 NumPy dot이 더 빠르게 나오기도 했다. 이 현상과 관련한 논의는 Stack Overflow 질문에 있다. 한 번의 숫자보다 dtype, 장치, 전송, 동기화 조건을 같게 맞췄는지가 먼저다.

jit은 함수를 컴파일 가능한 단위로 만든다

jit은 Python 함수를 XLA로 컴파일할 수 있게 감싼다. XLA 개념은 Google Developers 설명에서 볼 수 있다.

1
2
3
4
5
6
7
8
9
10
11
12
def selu(x, alpha=1.67, lmbda=1.05):
    return lmbda * jnp.where(
        x > 0,
        x,
        alpha * jnp.exp(x) - alpha,
    )

x = random.normal(key, (1_000_000,))
selu_jit = jit(selu)

%timeit selu(x).block_until_ready()
%timeit selu_jit(x).block_until_ready()

원문 측정에서는 일반 호출 약 5.66ms, JIT 호출 약 1.15ms가 나왔다. 이 숫자는 당시 환경의 결과이며 내 환경의 보장은 아니다. 특히 첫 호출에는 컴파일이 포함될 수 있으므로 warm-up과 반복 호출을 구분해야 한다.

jit이 모든 코드에 자동 이득을 주는 것도 아니다. 함수가 너무 작거나 한 번만 호출된다면 컴파일 비용을 회수하지 못할 수 있다. 반복되는 수치 계산처럼 컴파일된 함수를 여러 번 쓸 구간을 고르는 것이 중요하다.

grad는 미분식을 직접 쓰지 않게 한다

grad는 scalar를 반환하는 함수의 미분 함수를 만든다.

1
2
3
4
5
6
def sum_logistic(x):
    return jnp.sum(1.0 / (1.0 + jnp.exp(-x)))

x = jnp.arange(3.0)
derivative_fn = grad(sum_logistic)
print(derivative_fn(x))

원문 출력은 다음과 같았다.

1
[0.25       0.19661197 0.10499357]

수치 미분으로 대략적인 값을 비교할 수 있다.

1
2
3
4
5
6
7
8
def first_finite_differences(function, x):
    eps = 1e-3
    return jnp.array([
        (function(x + eps * v) - function(x - eps * v)) / (2 * eps)
        for v in jnp.eye(len(x))
    ])

print(first_finite_differences(sum_logistic, x))

두 결과가 가깝다고 해서 모든 미분 구현이 검증되는 것은 아니지만, 작은 함수에서 방향과 크기를 확인하는 sanity check로 쓸 수 있다. dtype과 eps가 달라지면 수치 미분 오차도 달라질 수 있다는 점을 함께 본다.

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를 평가하는 가장 작은 실험 순서는 다음과 같다.

  1. dtype과 실행 장치를 출력해 기록한다.
  2. 입력을 장치에 미리 두고 전송 시간과 계산 시간을 분리한다.
  3. block_until_ready()로 비동기 계산을 기다린다.
  4. 첫 JIT 호출과 이후 반복 호출을 나눠 잰다.
  5. grad 결과를 작은 입력의 수치 미분과 비교한다.
  6. Python loop와 vmap 결과 shape, 값이 같은지 먼저 확인한 뒤 시간을 잰다.

결론적으로 JAX의 장점은 “NumPy 한 줄보다 항상 빠르다”가 아니라, jit, grad, vmap 같은 변환을 조합해 반복되는 수치 함수를 컴파일하고 미분하고 벡터화할 수 있다는 데 있다.

공정한 JAX 벤치마크를 만드는 순서

첫째, NumPy와 JAX 입력의 dtype과 shape를 같게 만든다. 기본 dtype 차이나 암묵적 변환이 있으면 연산량과 정확도가 달라질 수 있다. 결과 배열의 값과 허용 오차를 먼저 비교해 빠르지만 다른 계산을 측정하지 않는다.

둘째, host에서 array를 만드는 시간과 device로 옮기는 시간을 분리한다. 반복 loop 안에서 매번 변환하면 JAX 연산이 아니라 데이터 이동을 재게 된다. 실제 애플리케이션이 매번 데이터를 받는다면 그 비용을 포함한 end-to-end 시간도 별도로 쓴다.

셋째, JIT 첫 호출과 같은 signature의 반복 호출을 나눈다. 입력 shape나 dtype가 달라지면 새로운 compilation이 일어날 수 있으므로 benchmark loop에서 signature가 고정됐는지 확인한다. Compile 시간이 중요한 일회성 작업과 반복 계산의 선택은 다를 수 있다.

넷째, 결과 준비를 기다린 뒤 timer를 멈춘다. 비동기 실행을 기다리지 않으면 queue에 작업을 넣는 시간만 측정할 수 있다. NumPy와 JAX 양쪽의 계산 범위, 반복 횟수, warm-up을 명시한다.

jit에 넣을 함수는 어떻게 정리하나

입력에서 출력이 결정되는 수치 계산을 작은 함수로 분리한다. 함수 안에서 Python side effect나 입력 값에 따른 동적 구조를 많이 사용하면 tracing과 실제 호출의 차이를 이해하기 어려워진다. 먼저 JIT 없이 결과를 검증한 뒤 장식한다.

같은 함수에 shape가 다른 입력을 계속 넣을지 운영 패턴을 본다. 고정 batch에서 반복하는 경우와 가변 길이 요청은 compile 비용의 비중이 다르다. 함수 경계를 너무 작게 잡아 매 연산을 따로 compile하는 것과 전체 pipeline을 한 번에 묶는 것 사이도 측정한다.

grad 결과를 어떻게 검증하나

출력 하나를 내는 작은 함수에서 시작해 입력과 gradient shape가 같은지 본다. 손으로 미분 가능한 값이나 작은 수치 변화와 방향을 비교한다. 함수가 vector를 반환한다면 어떤 scalar 목적을 미분하는지 먼저 정한다.

Gradient가 0이나 NaN일 때 optimizer부터 바꾸지 않는다. 입력 범위, 미분할 인자, 함수 안의 연산과 dtype을 확인한다. JIT를 함께 쓴다면 JIT 전후 gradient 값이 대응하는지도 본다.

vmap으로 loop를 바꿀 때 확인할 것

먼저 단일 sample 함수의 입력, 출력 shape를 적고 Python loop로 여러 sample 결과를 쌓는다. vmap이 어느 입력 축을 batch로 보고 출력 축을 어디에 두는지 지정한 뒤 두 배열을 비교한다. Axis를 잘못 잡으면 실행은 되지만 다른 차원을 반복할 수 있다.

Nested data나 여러 인자가 있다면 각각 batch 축이 있는지 구분한다. 모든 인자를 같은 방식으로 vectorize한다고 가정하지 않는다. 값이 같은 것을 확인한 뒤 loop overhead와 compiled 실행을 포함한 시간을 비교한다.

변환을 조합할 때의 실패 조건

jit(grad(f))grad(jit(f))처럼 조합을 볼 때는 목표 함수와 입력 signature를 고정한다. 결과, gradient, compile 횟수를 함께 보고 속도만 비교하지 않는다. Vmap과 grad를 섞을 때는 sample별 gradient인지 batch 합의 gradient인지 질문을 먼저 정한다.

CPU 강제 benchmark 결과를 GPU 성능으로 확대하지 않고, 사용한 device와 backend를 출력한다. 작은 문제에서는 compile, transfer가 계산보다 클 수 있고 큰 반복 문제에서는 반대가 될 수 있다. 입력 규모별 전환점을 보는 편이 한 번의 승패보다 유용하다.

함께 읽으면 이해가 이어지는 글

자주 묻는 질문

JAX의 첫 jit 호출이 느린 이유는 무엇인가요?

처음에는 함수의 입력 shape, dtype에 맞는 compilation 비용이 포함될 수 있습니다. 첫 호출과 같은 signature의 반복 호출을 나눠 측정해야 합니다.

JAX 시간을 잴 때 결과를 기다려야 하는 이유는 무엇인가요?

장치 연산이 비동기로 진행되면 Python 호출만 끝난 시각을 재어 실제 계산보다 짧게 보일 수 있습니다. 결과가 준비될 때까지 기다린 뒤 측정해야 합니다.

vmap은 모든 Python loop를 자동으로 빠르게 바꾸나요?

아닙니다. 함수가 어떤 축의 단일 sample을 받는지와 batch 출력 shape를 정확히 정해야 합니다. Loop와 결과 값이 같은지 검증한 뒤 성능을 비교해야 합니다.

THE END / OPSOAI

여기까지 읽었습니다

핵심 장면을 한 번 더 떠올려 보세요. 이해가 남았다면 이 책은 제 역할을 다했습니다.

다른 책 고르기
표지 1

키와 좌우 스와이프를 지원합니다. 읽던 페이지는 이 기기에 저장됩니다.

CONTENTS

이 책의 목차

    12개 장 16 분읽는 시간