pmap 안에서 여러 장치가 각자의 결과를 합치거나 공유해야 할 때는 집합 통신 연산을 사용해.
가장 흔한 집합 통신: jax.lax.psum, 모든 장치의 값을 합쳐.
def f(x):
'''각 device 의 x 합을 모든 device 에 broadcast'''
return jax.lax.psum(x, axis_name="i")
# pmap 으로 device axis 에 이름 부여
parallel = jax.pmap(f, axis_name="i")
# 4 device, 각각 다른 값
x = jnp.array([1.0, 2.0, 3.0, 4.0]) # device 4개
result = parallel(x)
print(result) # [10. 10. 10. 10.] — 모든 device 가 같은 합
주요 집합 통신 연산:
psum(x, axis_name), 모든 장치의 합pmean(x, axis_name), 모든 장치의 평균pmax(x, axis_name), 모든 장치의 최댓값pmin(x, axis_name), 모든 장치의 최솟값all_gather(x, axis_name), 모든 장치의 x를 이어 붙이기all_to_all(x, ...), 모든 장치가 모든 장치에게 데이터 전송
데이터-병렬 처리 그래디언트 평균 내기, 표준 사용 사례:
def train_step(params, batch_x, batch_y):
def loss_fn(p):
return jnp.mean((batch_x @ p - batch_y) ** 2)
loss, grads = jax.value_and_grad(loss_fn)(params)
# 모든 device 의 gradient 평균
grads = jax.lax.pmean(grads, axis_name="data")
loss = jax.lax.pmean(loss, axis_name="data")
new_params = params - 0.01 * grads
return new_params, loss
parallel_step = jax.pmap(
train_step,
in_axes=(None, 0, 0),
axis_name="data",
)
예를 들어 4개 GPU에서 batch_size=128로 학습할 때 각 GPU가 32개 예제를 처리한 뒤 그래디언트의 평균을 내. 결과적으로 batch_size=128 학습과 동등해.
여러 axis
# 2D mesh — data parallel + model parallel
parallel = jax.pmap(
jax.pmap(f, axis_name="model"),
axis_name="data",
)
# 두 axis 다 통합한 reduction
def f(x):
s_data = jax.lax.psum(x, axis_name="data")
s_model = jax.lax.psum(x, axis_name="model")
s_all = jax.lax.psum(x, axis_name=("data", "model"))
return s_all
큰 모델에서는 데이터 축으로 배치를 나누고 모델 axis로 매개변수를 나누는 2차원 패턴이 흔해.
🔄 집합 통신의 비용
집합 통신 연산에는 네트워크 비용이 들어. 4개 GPU가 psum을 수행하면 각 GPU의 데이터를 다른 GPU로 전송해. 큰 모델에서 그래디언트 (수 GB)를 매 스텝 psum을 수행하는 게, 학습 속도의 병목이 되는 일이 많아. NVIDIA NCCL과 JAX 집합 통신의 효율 차이가 학습 시간에 큰 영향을 줘.
통신 비용을 줄이려면 그래디언트 누적으로 통신 빈도를 낮추고, 혼합 정밀도로 데이터 크기를 줄이며, all_reduce 대신 reduce_scatter와 all_gather를 조합하는 ZeRO 패턴을 사용할 수 있어.