8.8 KiB
์์ ๋ถํ ๋ฐ์ดํฐ ๋ณ๋ ฌ ์ฒ๋ฆฌ(FSDP) fully-sharded-data-parallel
Fully Sharded Data Parallel (FSDP)์ ๋ชจ๋ธ์ ๋งค๊ฐ๋ณ์, ๊ทธ๋ ์ด๋์ธํธ ๋ฐ ์ตํฐ๋ง์ด์ ์ํ๋ฅผ ์ฌ์ฉ ๊ฐ๋ฅํ GPU(์์
์ ๋๋ ๋ญํฌ๋ผ๊ณ ๋ ํจ) ์์ ๋ฐ๋ผ ๋ถํ ํ๋ ๋ฐ์ดํฐ ๋ณ๋ ฌ ์ฒ๋ฆฌ ๋ฐฉ์์
๋๋ค. DistributedDataParallel (DDP)์ ๋ฌ๋ฆฌ, FSDP๋ ๊ฐ GPU์ ๋ชจ๋ธ์ ๋ณต์ ํ๊ธฐ ๋๋ฌธ์ ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋์ ์ค์
๋๋ค. ์ด๋ GPU ๋ฉ๋ชจ๋ฆฌ ํจ์จ์ฑ์ ํฅ์์ํค๋ฉฐ ์ ์ ์์ GPU๋ก ํจ์ฌ ๋ ํฐ ๋ชจ๋ธ์ ํ๋ จํ ์ ์๊ฒ ํฉ๋๋ค. FSDP๋ ๋ถ์ฐ ํ๊ฒฝ์์์ ํ๋ จ์ ์ฝ๊ฒ ๊ด๋ฆฌํ ์ ์๋ ๋ผ์ด๋ธ๋ฌ๋ฆฌ์ธ Accelerate์ ํตํฉ๋์ด ์์ผ๋ฉฐ, ๋ฐ๋ผ์ [Trainer] ํด๋์ค์์ ์ฌ์ฉํ ์ ์์ต๋๋ค.
์์ํ๊ธฐ ์ ์ Accelerate๊ฐ ์ค์น๋์ด ์๊ณ ์ต์ PyTorch 2.1.0 ์ด์์ ๋ฒ์ ์ด ์ค์น๋์ด ์๋์ง ํ์ธํ์ธ์.
pip install accelerate
FSDP ๊ตฌ์ฑ fsdp-configuration
์์ํ๋ ค๋ฉด accelerate config ๋ช
๋ น์ ์คํํ์ฌ ํ๋ จ ํ๊ฒฝ์ ๋ํ ๊ตฌ์ฑ ํ์ผ์ ์์ฑํ์ธ์. Accelerate๋ ์ด ๊ตฌ์ฑ ํ์ผ์ ์ฌ์ฉํ์ฌ accelerate config์์ ์ ํํ ํ๋ จ ์ต์
์ ๋ฐ๋ผ ์๋์ผ๋ก ์ฌ๋ฐ๋ฅธ ํ๋ จ ํ๊ฒฝ์ ์ค์ ํฉ๋๋ค.
accelerate config
accelerate config๋ฅผ ์คํํ๋ฉด ํ๋ จ ํ๊ฒฝ์ ๊ตฌ์ฑํ๊ธฐ ์ํ ์ผ๋ จ์ ์ต์
๋ค์ด ๋ํ๋ฉ๋๋ค. ์ด ์น์
์์๋ ๊ฐ์ฅ ์ค์ํ FSDP ์ต์
์ค ์ผ๋ถ๋ฅผ ๋ค๋ฃน๋๋ค. ๋ค๋ฅธ ์ฌ์ฉ ๊ฐ๋ฅํ FSDP ์ต์
์ ๋ํด ๋ ์์๋ณด๊ณ ์ถ๋ค๋ฉด fsdp_config ๋งค๊ฐ๋ณ์๋ฅผ ์ฐธ์กฐํ์ธ์.
๋ถํ ์ ๋ต sharding-strategy
FSDP๋ ์ฌ๋ฌ ๊ฐ์ง ๋ถํ ์ ๋ต์ ์ ๊ณตํฉ๋๋ค:
FULL_SHARD- ๋ชจ๋ธ ๋งค๊ฐ๋ณ์, ๊ทธ๋ ์ด๋์ธํธ ๋ฐ ์ตํฐ๋ง์ด์ ์ํ๋ฅผ ์์ ์ ๊ฐ์ ๋ถํ ; ์ด ์ต์ ์ ์ ํํ๋ ค๋ฉด1์ ์ ํํ์ธ์SHARD_GRAD_OP- ๊ทธ๋ ์ด๋์ธํธ ๋ฐ ์ตํฐ๋ง์ด์ ์ํ๋ฅผ ์์ ์ ๊ฐ์ ๋ถํ ; ์ด ์ต์ ์ ์ ํํ๋ ค๋ฉด2๋ฅผ ์ ํํ์ธ์NO_SHARD- ์๋ฌด ๊ฒ๋ ๋ถํ ํ์ง ์์ (DDP์ ๋์ผ); ์ด ์ต์ ์ ์ ํํ๋ ค๋ฉด3์ ์ ํํ์ธ์HYBRID_SHARD- ๊ฐ ์์ ์๊ฐ ์ ์ฒด ๋ณต์ฌ๋ณธ์ ๊ฐ์ง๊ณ ์๋ ์ํ์์ ๋ชจ๋ธ ๋งค๊ฐ๋ณ์, ๊ทธ๋ ์ด๋์ธํธ ๋ฐ ์ตํฐ๋ง์ด์ ์ํ๋ฅผ ์์ ์ ๋ด์์ ๋ถํ ; ์ด ์ต์ ์ ์ ํํ๋ ค๋ฉด4๋ฅผ ์ ํํ์ธ์HYBRID_SHARD_ZERO2- ๊ฐ ์์ ์๊ฐ ์ ์ฒด ๋ณต์ฌ๋ณธ์ ๊ฐ์ง๊ณ ์๋ ์ํ์์ ๊ทธ๋ ์ด๋์ธํธ ๋ฐ ์ตํฐ๋ง์ด์ ์ํ๋ฅผ ์์ ์ ๋ด์์ ๋ถํ ; ์ด ์ต์ ์ ์ ํํ๋ ค๋ฉด5๋ฅผ ์ ํํ์ธ์
์ด๊ฒ์ fsdp_sharding_strategy ํ๋๊ทธ๋ก ํ์ฑํ๋ฉ๋๋ค.
CPU ์คํ๋ก๋ cpu-offload
์ฌ์ฉํ์ง ์๋ ๋งค๊ฐ๋ณ์์ ๊ทธ๋ ์ด๋์ธํธ๋ฅผ CPU๋ก ์คํ๋ก๋ํ์ฌ ๋ ๋ง์ GPU ๋ฉ๋ชจ๋ฆฌ๋ฅผ ์ ์ฝํ๊ณ FSDP๋ก๋ ์ถฉ๋ถํ์ง ์์ ํฐ ๋ชจ๋ธ์ GPU์ ์ ์ฌํ ์ ์๋๋ก ํ ์ ์์ต๋๋ค. ์ด๋ accelerate config๋ฅผ ์คํํ ๋ fsdp_offload_params: true๋ก ์ค์ ํ์ฌ ํ์ฑํ๋ฉ๋๋ค.
๋ํ ์ ์ฑ wrapping-policy
FSDP๋ ๋คํธ์ํฌ์ ๊ฐ ๋ ์ด์ด๋ฅผ ๋ํํ์ฌ ์ ์ฉ๋ฉ๋๋ค. ๋ํ์ ์ผ๋ฐ์ ์ผ๋ก ์ค์ฒฉ ๋ฐฉ์์ผ๋ก ์ ์ฉ๋๋ฉฐ ๊ฐ๊ฐ ์๋ฐฉํฅ์ผ๋ก ์ง๋๊ฐ ํ ์ ์ฒด ๊ฐ์ค์น๋ฅผ ์ญ์ ํ์ฌ ๋ค์ ๋ ์ด์ด์์ ์ฌ์ฉํ ๋ฉ๋ชจ๋ฆฌ๋ฅผ ์ ์ฝํฉ๋๋ค. ์๋ ๋ํ ์ ์ฑ
์ ์ด๋ฅผ ๊ตฌํํ๋ ๊ฐ์ฅ ๊ฐ๋จํ ๋ฐฉ๋ฒ์ด๋ฉฐ ์ฝ๋๋ฅผ ๋ณ๊ฒฝํ ํ์๊ฐ ์์ต๋๋ค. Transformer ๋ ์ด์ด๋ฅผ ๋ํํ๋ ค๋ฉด fsdp_auto_wrap_policy: TRANSFORMER_BASED_WRAP๋ฅผ ์ ํํ๊ณ ๋ํํ ๋ ์ด์ด๋ฅผ ์ง์ ํ๋ ค๋ฉด fsdp_transformer_layer_cls_to_wrap๋ฅผ ์ ํํ์ธ์ (์: BertLayer).
๋๋ ํน์ ๋งค๊ฐ๋ณ์ ์๋ฅผ ์ด๊ณผํ ๊ฒฝ์ฐ FSDP๊ฐ ๋ ์ด์ด์ ์ ์ฉ๋๋ ํฌ๊ธฐ ๊ธฐ๋ฐ ๋ํ ์ ์ฑ
์ ์ ํํ ์ ์์ต๋๋ค. ์ด๋ fsdp_wrap_policy: SIZE_BASED_WRAP ๋ฐ min_num_param์ ์ํ๋ ํฌ๊ธฐ์ ์๊ณ๊ฐ์ผ๋ก ์ค์ ํ์ฌ ํ์ฑํ๋ฉ๋๋ค.
์ฒดํฌํฌ์ธํธ checkpointing
์ค๊ฐ ์ฒดํฌํฌ์ธํธ๋ fsdp_state_dict_type: SHARDED_STATE_DICT๋ก ์ ์ฅํด์ผ ํฉ๋๋ค. CPU ์คํ๋ก๋๊ฐ ํ์ฑํ๋ ๋ญํฌ 0์์ ์ ์ฒด ์ํ ๋์
๋๋ฆฌ๋ฅผ ์ ์ฅํ๋ ๋ฐ ์๊ฐ์ด ๋ง์ด ๊ฑธ๋ฆฌ๊ณ , ๋ธ๋ก๋์บ์คํ
์ค ๋ฌด๊ธฐํ ๋๊ธฐํ์ฌ NCCL Timeout ์ค๋ฅ๊ฐ ๋ฐ์ํ ์ ์๊ธฐ ๋๋ฌธ์
๋๋ค. [~accelerate.Accelerator.load_state] ๋ฉ์๋๋ฅผ ์ฌ์ฉํ์ฌ ๋ถํ ๋ ์ํ ๋์
๋๋ฆฌ๋ก ํ๋ จ์ ์ฌ๊ฐํ ์ ์์ต๋๋ค.
# ๊ฒฝ๋ก๊ฐ ๋ด์ฌ๋ ์ฒดํฌํฌ์ธํธ
accelerator.load_state("ckpt")
๊ทธ๋ฌ๋ ํ๋ จ์ด ๋๋๋ฉด ์ ์ฒด ์ํ ๋์ ๋๋ฆฌ๋ฅผ ์ ์ฅํด์ผ ํฉ๋๋ค. ๋ถํ ๋ ์ํ ๋์ ๋๋ฆฌ๋ FSDP์๋ง ํธํ๋๊ธฐ ๋๋ฌธ์ ๋๋ค.
if trainer.is_fsdp_enabled:
trainer.accelerator.state.fsdp_plugin.set_state_dict_type("FULL_STATE_DICT")
trainer.save_model(script_args.output_dir)
TPU tpu
PyTorch XLA๋ TPU์ ๋ํ FSDP ํ๋ จ์ ์ง์ํ๋ฉฐ accelerate config๋ก ์์ฑ๋ FSDP ๊ตฌ์ฑ ํ์ผ์ ์์ ํ์ฌ ํ์ฑํํ ์ ์์ต๋๋ค. ์์์ ์ง์ ํ ๋ถํ ์ ๋ต ๋ฐ ๋ํ ์ต์
์ธ์๋ ์๋์ ํ์๋ ๋งค๊ฐ๋ณ์๋ฅผ ํ์ผ์ ์ถ๊ฐํ ์ ์์ต๋๋ค.
xla: True # PyTorch/XLA๋ฅผ ํ์ฑํํ๋ ค๋ฉด True๋ก ์ค์ ํด์ผ ํฉ๋๋ค
xla_fsdp_settings: # XLA ํน์ FSDP ๋งค๊ฐ๋ณ์
xla_fsdp_grad_ckpt: True # gradient checkpointing์ ์ฌ์ฉํฉ๋๋ค
xla_fsdp_settings๋ FSDP์ ๋ํ ์ถ๊ฐ์ ์ธ XLA ํน์ ๋งค๊ฐ๋ณ์๋ฅผ ๊ตฌ์ฑํ ์ ์๊ฒ ํฉ๋๋ค.
ํ๋ จ ์์ launch-training
์์ FSDP ๊ตฌ์ฑ ํ์ผ์ ๋ค์๊ณผ ๊ฐ์ ์ ์์ต๋๋ค:
compute_environment: LOCAL_MACHINE
debug: false
distributed_type: FSDP
downcast_bf16: 'no'
fsdp_config:
fsdp_auto_wrap_policy: TRANSFORMER_BASED_WRAP
fsdp_backward_prefetch_policy: BACKWARD_PRE
fsdp_cpu_ram_efficient_loading: true
fsdp_forward_prefetch: false
fsdp_offload_params: true
fsdp_sharding_strategy: 1
fsdp_state_dict_type: SHARDED_STATE_DICT
fsdp_sync_module_states: true
fsdp_transformer_layer_cls_to_wrap: BertLayer
fsdp_use_orig_params: true
machine_rank: 0
main_training_function: main
mixed_precision: bf16
num_machines: 1
num_processes: 2
rdzv_backend: static
same_network: true
tpu_env: []
tpu_use_cluster: false
tpu_use_sudo: false
use_cpu: false
ํ๋ จ์ ์์ํ๋ ค๋ฉด accelerate launch ๋ช
๋ น์ ์คํํ์ธ์. ์ด ๋ ์ ์ accelerate config๋ก ์์ฑํ ๊ตฌ์ฑ ํ์ผ์ ์๋์ผ๋ก ์ฌ์ฉํฉ๋๋ค.
accelerate launch my-trainer-script.py
accelerate launch --fsdp="full shard" --fsdp_config="path/to/fsdp_config/ my-trainer-script.py
๋ค์ ๋จ๊ณ next-steps
FSDP๋ ๋งค์ฐ ํฐ ๋ชจ๋ธ์ ํ๋ จํ ๋ ๊ฐ๋ ฅํ ๋๊ตฌ๊ฐ ๋ ์ ์์ผ๋ฉฐ, ์ฌ๋ฌ ๊ฐ์ GPU๋ TPU๋ฅผ ์ฌ์ฉํ ์ ์์ต๋๋ค. ๋ชจ๋ธ ๋งค๊ฐ๋ณ์, ์ตํฐ๋ง์ด์ ๋ฐ ๊ทธ๋ ์ด๋์ธํธ ์ํ๋ฅผ ๋ถํ ํ๊ณ ๋นํ์ฑ ์ํ์ผ ๋, CPU๋ก ์คํ๋ก๋ํ๋ฉด FSDP๋ ๋๊ท๋ชจ ํ๋ จ์ ๋์ ์ฐ์ฐ ๋น์ฉ์ ์ค์ผ ์ ์์ต๋๋ค. ๋ ์์๋ณด๊ณ ์ถ๋ค๋ฉด ๋ค์ ์๋ฃ๊ฐ ๋์์ด ๋ ์ ์์ต๋๋ค:
- FSDP์ ๋ํ ๋ ๊น์ด ์๋ Accelerate ๊ฐ์ด๋๋ฅผ ๋ฐ๋ผ๊ฐ ๋ณด์ธ์.
- PyTorch์ ์์ ๋ถํ ๋ฐ์ดํฐ ๋ณ๋ ฌ ์ฒ๋ฆฌ (FSDP) API๋ฅผ ์๊ฐํฉ๋๋ค ๋ธ๋ก๊ทธ ๊ธ์ ์ฝ์ด๋ณด์ธ์.
- FSDP๋ฅผ ์ฌ์ฉํ์ฌ ํด๋ผ์ฐ๋ TPU์์ PyTorch ๋ชจ๋ธ ํฌ๊ธฐ ์กฐ์ ํ๊ธฐ ๋ธ๋ก๊ทธ ๊ธ์ ์ฝ์ด๋ณด์ธ์.