학습을 위한 텐서 병렬화

학습을 위한 텐서 병렬화 (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,
)

다음 단계