jax.pmap은 병렬 map이야. 여러 장치(GPU, TPU)에서 같은 프로그램을 다른 데이터로 동시에 돌려.
약자: SPMD = Single Program, Multiple 데이터. 모든 장치가 같은 코드를 실행하지만, 입력 데이터는 장치마다 다른 조각을 받아.
import jax
import jax.numpy as jnp
# 사용 가능한 device 확인
print(jax.devices())
print(f"device 개수: {jax.device_count()}")
# 단일 device 환경에서도 시뮬레이션 가능
import os
os.environ["XLA_FLAGS"] = "--xla_force_host_platform_device_count=4"
# (다음 import 부터 4 device 인 척)
기본 사용:
def f(x):
return x ** 2 + jnp.sin(x)
# pmap — 첫 번째 axis 가 device axis (vmap 과 비슷)
parallel_f = jax.pmap(f)
# 입력 첫 axis 의 길이가 device 개수와 일치해야 함
x = jnp.arange(8).reshape(4, 2) # 4 devices, 각 device 가 (2,) 받음
result = parallel_f(x) # (4, 2)
각 장치가 입력의 slice를 받아 같은 함수 실행해. 결과의 첫 axis도 장치 axis.
학습 스텝의 SPMD화:
def train_step(params, batch_x, batch_y):
'''단일 device 의 train step'''
def loss_fn(p):
return jnp.mean((batch_x @ p - batch_y) ** 2)
loss, grads = jax.value_and_grad(loss_fn)(params)
new_params = params - 0.01 * grads
return new_params, loss
# pmap — 모든 device 가 같은 step 실행
parallel_step = jax.pmap(train_step, in_axes=(None, 0, 0))
# params 는 모든 device 에 broadcast, batch 는 device 별로 나눔
# batch_x: (n_devices, B/n_devices, D)
# batch_y: (n_devices, B/n_devices)
new_params, losses = parallel_step(params, batch_x_sharded, batch_y_sharded)
중요한 점은 이대로면 각 장치의 매개변수가 따로 업데이트되어 발산한다는 거야. 집합 통신 연산으로 그래디언트의 평균을 모든 장치에 걸쳐 내야 해. 다음 레슨에서 다룰게.
vmap vs pmap 차이:
- vmap: 단일 장치, 배치 축 자동화. 메모리 / op flow는 한 장치 안에서.
- pmap: 여러 장치, 각 장치가 독립적인 메모리. communication은 명시적 집합 통신.
vmap: 하나의 array (B, D) → 하나의 device 가 배치 처리
pmap: 여러 array, 각각 (B', D) → N 개 device 가 각자 처리, 필요시 collective
🌐 pmap = MPI의 ML 버전
SPMD는 HPC에서 오래전부터 사용한 패턴이야. MPI 프로그래밍에서는 모든 노드가 같은 바이너리를 실행하고 입력 데이터만 다르며, 명시적 메시지 전달로 동기화해. pmap은 그 모델을 머신러닝에 가져온 거야. 데이터-병렬 처리 학습이 SPMD의 정석 사례야.
최근에는 JAX 팀이 jax.sharding + Mesh로 pmap을 점진적으로 대체하는 중. pmap은 단순 데이터-병렬 처리에 강하지만, 모델 병렬 처리 같은 복잡 패턴에선 샤딩 API가 더 깔끔해. Track 7-4에서 다뤄.