데이터 병렬 학습의 표준 패턴은 pmap, grad, 집합 통신의 조합이야. 한 번 익히면 어느 모델이든 같은 형태로 확장할 수 있어.
import jax
import jax.numpy as jnp
from functools import partial
# ============ 1. params 를 모든 device 에 복제 ============
def replicate(params, n_devices):
return jax.tree.map(
lambda x: jnp.broadcast_to(x, (n_devices,) + x.shape),
params,
)
# ============ 2. data 를 device 별로 분할 ============
def shard(data, n_devices):
'''(B, ...) → (n_devices, B//n_devices, ...)'''
return data.reshape(n_devices, -1, *data.shape[1:])
# ============ 3. data-parallel train step ============
@partial(jax.pmap, axis_name="data")
def train_step(params, x, y):
def loss_fn(p):
pred = x @ p["w"] + p["b"]
return jnp.mean((pred - 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 = jax.tree.map(lambda p, g: p - 0.01 * g, params, grads)
return new_params, loss
# ============ 4. 학습 loop ============
n_devices = jax.device_count()
print(f"{n_devices} devices")
# 초기 params (단일 device 에서 만든 후 replicate)
params = {"w": jnp.zeros(10), "b": jnp.zeros(())}
params = replicate(params, n_devices)
# 한 batch
batch_x = jnp.zeros((128, 10)) # B=128
batch_y = jnp.zeros(128)
batch_x_sharded = shard(batch_x, n_devices) # (4, 32, 10)
batch_y_sharded = shard(batch_y, n_devices) # (4, 32)
for step in range(100):
params, loss = train_step(params, batch_x_sharded, batch_y_sharded)
# loss 는 pmap 결과라 (n_devices,) shape — 모든 device 가 같은 값
if step % 10 == 0:
print(f"step {step}: loss = {loss[0]:.4f}")
핵심 포인트:
- 매개변수 replicate: 모든 장치가 동일한 매개변수를 가짐.
(n_devices, ...)모양. - 데이터 shard: 입력의 첫 axis가 장치 axis이고, 각 장치가 배치의 1/n_devices를 담당해.
- local 순전파 + grad: 각 장치는 자기 batch로 그래디언트 계산해.
- pmean 그래디언트: 모든 장치의 그래디언트 평균 → 모든 장치가 동일한 update.
- 매개변수 동기화 유지: 모든 장치의 매개변수가 항상 같음 (수학적으로).
체크포인트 저장 / 복원
pmap 후 매개변수는 (n_devices, ...) 모양이지만 모든 장치가 같으므로 0번째만 저장해:
def get_first_device_params(params):
'''device axis 제거'''
return jax.tree.map(lambda x: x[0], params)
# 저장
single_params = get_first_device_params(params)
# ... pickle 또는 orbax 로 저장 ...
# 복원 후 다시 replicate
loaded_params = ...
params = replicate(loaded_params, n_devices)
📐 데이터-병렬 처리 scaling rule
n_devices를 늘릴 때 effective batch size도 장치 수만큼 늘어나. 같은 학습 동역학을 유지하려면 학습률도 장치 수에 맞춰 늘려야 해(linear scaling rule). 너무 큰 배치라면 warmup + cosine 스케줄로 안정화. Track 11에서 다뤄.
이 패턴이 JAX 데이터 병렬 처리의 정석이야. Track 7-4의 sharding API는 같은 일을 더 깨끗하게, 그리고 모델 병렬 처리까지 자연스럽게 확장할 수 있게 해.