본문 바로가기
C.W.K.
Stream
Lesson 02 of 05 · published

Collective Operations: 장치 간 통신

~8 min · pmap, jax, tutorial

Level 0호기심
0 XP0/73 lessons0/17 achievements
0/100 XP to next level100 XP to go0% complete

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 패턴을 사용할 수 있어.

Code

import jax
import jax.numpy as jnp

# psum: sum values across all devices (all-reduce sum)
@jax.pmap
def sum_across_devices(x):
    # x is this device's local value
    total = jax.lax.psum(x, axis_name='devices')
    return total

# Note: pmap needs to know the axis name for collectives
sum_across_devices = jax.pmap(
    lambda x: jax.lax.psum(x, axis_name='i'),
    axis_name='i'
)

# pmean: average across devices
mean_fn = jax.pmap(
    lambda x: jax.lax.pmean(x, axis_name='i'),
    axis_name='i'
)

# pmax: maximum across devices
max_fn = jax.pmap(
    lambda x: jax.lax.pmax(x, axis_name='i'),
    axis_name='i'
)
import jax
import jax.numpy as jnp

def loss_fn(params, x, y):
    pred = jnp.dot(x, params)
    return jnp.mean((pred - y) ** 2)

def train_step(params, x, y, lr):
    """Training step that runs on each device."""
    loss, grads = jax.value_and_grad(loss_fn)(params, x, y)

    # Average gradients across all devices
    grads = jax.lax.pmean(grads, axis_name='devices')
    loss = jax.lax.pmean(loss, axis_name='devices')

    # Update (all devices now have the same gradients → same params)
    new_params = params - lr * grads
    return new_params, loss

# Parallelize with pmap
parallel_train_step = jax.pmap(train_step, axis_name='devices')

External links

Exercise

pmap으로 실행되는 함수 안에서 장치별 손실을 계산하고 jax.lax.psum으로 전역 합계를 구해. 모든 장치가 같은 결과를 받는지 확인한 다음 pmean으로 바꿔 평균도 검증해.

Progress

Progress is local-only — sign in to sync across devices.
이 페이지에서 버그를 발견하셨거나 피드백이 있으세요?문제 신고

댓글 0

🔔 답글 알림 (로그인 필요)
로그인댓글을 남기려면 로그인해 주세요.

아직 댓글이 없어요. 첫 댓글을 남겨보세요.