FSDP2

FSDP2

Fully Sharded Data Parallel (FSDP2)은 모델, 그라디언트, 옵티마이저 상태를 GPU 간에 샤딩해요. 계산 전에 각 GPU가 모든 샤드에서 파라미터의 완전한 세트를 모은 다음, 이후에 해제해요. 샤딩을 통해 단일 GPU 메모리보다 큰 모델을 학습할 수 있으며, 대가로 DDP보다 더 많은 통신이 발생해요. 모델이나 옵티마이저 상태가 단일 GPU에 들어가지 않을 때 FSDP를 사용해요.

                      ┌─────────────────┐
                      │  training data  │
                      └────────┬────────┘
            ┌──────────────────┼──────────────────┐
            │ shard 0          │ shard 1          │ shard 2
            ▼                  ▼                  ▼
     ┌─────────────┐    ┌─────────────┐    ┌─────────────┐
     │  param      │    │  param      │    │  param      │
     │  shard 0    │    │  shard 1    │    │  shard 2    │
     │  GPU 0      │    │  GPU 1      │    │  GPU 2      │
     └──────┬──────┘    └──────┬──────┘    └──────┬──────┘
            │                  │                  │
            └──────── all-gather (params) ────────┘
                               │
                    full params on each GPU
                               │
            ┌──────────────────┼──────────────────┐
            ▼                  ▼                  ▼
         forward             forward             forward
            │                  │                  │
            └───── reduce-scatter (grads) ────────┘
                               │
            ┌──────────────────┼──────────────────┐
            ▼                  ▼                  ▼
     grad shard 0       grad shard 1       grad shard 2
     optim shard 0      optim shard 1      optim shard 2
        step               step               step

출처: 문서

본문

샤딩 전략

FSDP2는 ~TrainingArguments.fsdp_config로 샤딩을 제어해요. FSDP를 활성화하려면 fsdp=True를 설정하고, FSDP config에서 reshard_after_forward를 설정해서 메모리와 처리량 간의 트레이드오프를 선택해요.

reshard_after_forward 동작
true 더 많은 메모리를 절약하기 위해 forward pass 후에 파라미터를 reshard해요.
false 재-all-gather를 피하기 위해 forward와 backward 사이에 파라미터를 gathered 상태로 유지하되, 더 높은 최고(peak) 메모리를 감수해요.

auto_wrap_policy는 모듈이 FSDP 유닛으로 어떻게 래핑되는지 제어해요. 기본값은 "TRANSFORMER_BASED_WRAP"로, 모델의 transformer 레이어를 래핑해요. 래핑하지 않으면("NO_WRAP") 전체 모델이 하나의 FSDP 유닛이 되어 샤딩의 메모리 이점을 잃어요.

FSDP 구성하기

이 필드들은 FSDP2가 모델을 어떻게 래핑하고, 샤딩하고, 로드하는지 제어해요. reshard_after_forward와 auto_wrap_policy는 Sharding strategies에서 다뤄요.

  • cpu_offload는 사용하지 않을 때 파라미터와 그라디언트를 CPU로 오프로드해서 GPU 메모리를 절약해요.

  • transformer_layer_cls_to_wrap는 auto_wrap_policy가 "TRANSFORMER_BASED_WRAP"일 때 FSDP 유닛으로 래핑할 transformer 레이어를 정의해요. 각 유닛은 자신의 gather와 scatter 연산을 관리해요. forward pass 동안 현재 유닛의 파라미터만 gathered돼요. 이전 유닛들의 파라미터는 메모리를 절약하기 위해 해제돼요.

    최상위 모델만 래핑하면 GPU 메모리 절약 효과가 없어요. 개별 Linear 레이어마다 래핑하면 유닛 간 통신이 매우 비싸져요. 이 필드를 비워두면 FSDP가 모델 정의에서 값을 읽어요.

  • min_num_params는 크기 기반 래핑을 위한 모듈당 최소 파라미터 수를 설정해요. auto_wrap_policy가 "SIZE_BASED_WRAP"일 때만 사용돼요.

  • state_dict_type는 체크포인트 형식을 제어해요. 단일 Transformers 호환 체크포인트를 위해 기본값은 "FULL_STATE_DICT"예요. rank당 체크포인트 파일 하나를 원하면 "SHARDED_STATE_DICT"를 사용해요. 대형 모델에서 더 빠르죠. 샤딩된 체크포인트는 FSDP에만 다시 로드할 수 있으므로, 공유하거나 FSDP 밖에서 로드할 최종 체크포인트는 "FULL_STATE_DICT"로 저장해요.

  • cpu_ram_efficient_loading은 rank 0에서만 디스크에서 체크포인트를 로드해요. 다른 GPU는 빈 모델을 초기화하고 broadcast로 가중치를 받아서, 여러 프로세스가 큰 모델을 CPU RAM에 로드하는 것을 피해요.

  • activation_checkpointing은 활성화를 저장하는 대신 backward pass에서 다시 계산해요. TrainingArguments의 gradient checkpointing 대신 이걸 사용해요. 둘 다 설정하면 에러가 발생해요.

FSDP 학습은 Accelerate config file 또는 fsdp_config에 전달되는 FSDP config 파일로 구성할 수 있어요.

accelerate config 명령을 실행하고 하드웨어와 학습 설정에 대한 질문에 답해요. 이렇게 하면 캐시에 default_config.yaml 파일이 생성돼요.

Trainer 기반 스크립트로 accelerate launch를 실행해요. Accelerate config 파일이 동일한 설정을 다루므로 fsdp_config는 불필요해요.

accelerate launch train.py
{
  "version": 2,
  "reshard_after_forward": true,
  "cpu_offload": false,
  "auto_wrap_policy": "TRANSFORMER_BASED_WRAP",
  "transformer_layer_cls_to_wrap": ["LlamaDecoderLayer"],
  "state_dict_type": "FULL_STATE_DICT",
  "cpu_ram_efficient_loading": true,
  "activation_checkpointing": true
}

fsdp=True로 설정하고 FSDP config 파일을 fsdp_config에 전달해요.

from transformers import TrainingArguments

TrainingArguments(
    ...,
    fsdp=True,
    fsdp_config="path/to/fsdp.json",
)

더 알아보기 (Learn more)

  • 모델이 GPU 하나에 들어갈 때 데이터 병렬 학습은 DDP를 참고해요.
  • ZeRO 최적화와 NVMe offloading은 DeepSpeed를 참고해요.
  • PyTorch/XLA로 TPU에서 FSDP를 사용하려면 ~TrainingArguments.fsdp_config에서 xla, xla_fsdp_settings, xla_fsdp_grad_ckpt를 설정해요.
  • FSDP가 어떻게 동작하는지 더 자세히 알고 싶다면 The Ultra-Scale Playbook의 FSDP chapter를 읽어 보세요.