학습을 위한 텐서 병렬화
학습을 위한 텐서 병렬화 (Tensor parallelism for training)
텐서 병렬화(TP, Tensor Parallelism)는 가중치 행렬을 열 방향 또는 행 방향으로 GPU들에 분할해요. 각 GPU는 샤드(shard)를 보유하고 부분 결과를 계산한 뒤, all-reduce로 동기화해서 전체 출력을 만들어 내요.
출처: 문서
본문
TP는 잦은 GPU 간 통신에 의존해요. NVLink처럼 빠른 인트라-노드(intra-node) 링크를 갖춘 하드웨어에서 가장 잘 작동해요.
┌─────────────────────────────┐
│ X (replicated) │
└────┬──────────┬─────────┬───┘
│ │ │
┌────▼───┐ ┌────▼───┐ ┌───▼────┐
│ ▓▓▓ W₀ │ │ ░░░ W₁ │ │ ███ W₂ │
│ X@W₀ │ │ X@W₁ │ │ X@W₂ │
└────┬───┘ └────┬───┘ └───┬────┘
└──────────┼─────────┘
Y₀+Y₁+Y₂
┌────────────────────────────┐
│ Y (full) │
└────────────────────────────┘
Transformers는 config가 base_model_tp_plan을 정의하는 아키텍처에 대해 TP를 지원해요. 먼저 그 필드를 확인해서 모델이 네이티브 TP를 지원하는지 살펴보세요.
from transformers import AutoConfig
config = AutoConfig.from_pretrained("Qwen/Qwen3-0.6B")
print(config.base_model_tp_plan is not None)
print(config.base_model_tp_plan)
모델이 TP를 지원한다면 tp_size에 장치 개수를 넣은 DistributedConfig를 만들어 from_pretrained()에 전달해요. Transformers가 모델의 사전 정의된 플랜을 사용해 디바이스 메시(device mesh)를 초기화하고, 지원되는 계층을 알아서 샤딩해요.
DistributedConfig에서 tp_plan="auto"로 설정할 수도 있어요. tp_size를 생략하면 WORLD_SIZE에서 추론돼요. tp_plan을 from_pretrained()에 직접 전달하는 방식은 더 이상 권장되지 않으며 v5.18에서 제거될 예정이에요.
[!WARNING]
distributed_config와 함께device_map을 사용하지 마세요. 둘은 가중치 로딩 수준에서 충돌해요.device_map은 전체 모듈을 특정 GPU에 배치하지만, 텐서 병렬화는 그 동일한 파라미터를 모든 GPU에 걸쳐 샤딩하기 때문이에요.
import torch
from transformers import AutoModelForCausalLM, DistributedConfig
distributed_config = DistributedConfig(tp_size=4)
model = AutoModelForCausalLM.from_pretrained(
"Qwen/Qwen3-0.6B",
dtype=torch.bfloat16,
distributed_config=distributed_config,
)
Trainer가 텐서 병렬 플랜을 감지하고, 모델에서 tp_size를 읽어 ParallelismConfig를 자동으로 만들어요.
4개의 GPU가 있는 한 노드에서 학습을 시작해요.
torchrun --nproc-per-node 4 train_tp.py
ParallelismConfig
TP를 FSDP 같은 다른 병렬화 기법과 결합할 때는 ParallelismConfig를 명시적으로 전달해요.
import torch
from accelerate import ParallelismConfig
from transformers import AutoModelForCausalLM, DistributedConfig, TrainingArguments
distributed_config = DistributedConfig(tp_size=4)
model = AutoModelForCausalLM.from_pretrained(
"Qwen/Qwen3-0.6B",
dtype=torch.bfloat16,
distributed_config=distributed_config,
)
parallelism_config = ParallelismConfig(tp_size=4)
args = TrainingArguments(
...,
parallelism_config=parallelism_config,
)
다음 단계
- 작동 방식에 대한 더 자세한 내용은 The Ultra-Scale Playbook의 Tensor Parallelism 챕터를 읽어보세요.
- 분할 전략, 수동 TP 플랜, 구현 세부 사항에 대해 더 배우려면 tensor parallelism inference guide를 읽어보세요.