본문으로 건너뛰기

[sglang] [SGLang] Diffusion 모델의 TP/FSDP 체크포인트 로딩 속도를 20% 이상 개선하는 방법

PR 링크: sgl-project/sglang#33960 상태: Merged | 변경: +757 / -107

들어가며

대규모 Diffusion 모델(DiT 등)을 서빙하거나 학습할 때, Tensor Parallelism(TP)이나 Fully Sharded Data Parallelism(FSDP)은 필수적입니다. 하지만 기존 SGLang의 구현에서는 모델을 로드할 때 치명적인 비효율이 있었습니다. 바로 모든 GPU 랭크가 전체 체크포인트 파일을 처음부터 끝까지 스트리밍하며 읽은 뒤, 자신에게 필요한 부분만 골라내는 방식이었습니다.

예를 들어, 22.9 GiB 크기의 Z-Image 모델을 2개의 GPU(TP2)에서 로드할 때, 각 GPU는 실제로 절반의 가중치만 보유하면 됨에도 불구하고 두 GPU 모두 22.9 GiB 전체를 읽어 들였습니다. 이는 불필요한 I/O 오버헤드를 발생시키고, 특히 네트워크 스토리지 환경에서 심각한 병목 현상을 초래합니다.

최근 SGLang에 병합된 이 PR은 safetensors의 랜덤 액세스 기능을 활용하여, 각 랭크가 자신에게 필요한 가중치 슬라이스만 선택적으로 로드(Rank-local loading)하도록 개선했습니다. 이를 통해 로딩 속도를 최대 23%까지 끌어올렸습니다.

코드 분석: 전체 로딩에서 슬라이스 로딩으로

1. 로더 로직의 변화 (fsdp_load.py)

기존에는 safetensors_weights_iterator를 통해 모든 가중치를 순회하며 체크포인트를 읽었습니다. 변경된 코드에서는 rank_local_checkpoint 모듈을 도입하여, 로컬 랭크에 필요한 state_dict만 미리 구성합니다.

Before:

# 기존: 모든 가중치를 순회하는 iterator 생성
weight_iterator = safetensors_weights_iterator(weight_dir_list)
preprocess_loaded_state_dict = getattr(model, "preprocess_loaded_state_dict", None)
if preprocess_loaded_state_dict is not None:
    weight_iterator = preprocess_loaded_state_dict(weight_iterator)

load_model_from_full_model_state_dict(
    model,
    weight_iterator,
    # ... 생략
)

After:

# 개선: 랭크별로 필요한 부분만 골라 담은 preconverted_state_dict 생성
preconverted_state_dict = None
if (use_fsdp and weight_dir_list and preprocess_loaded_state_dict is None and not is_bnb_quantized):
    preconverted_state_dict = rank_local_checkpoint.try_load_rank_local_fsdp_state_dict(
        model, weight_dir_list, param_names_mapping_fn,
    )
elif (not use_fsdp and weight_dir_list and preprocess_loaded_state_dict is None and not is_bnb_quantized):
    preconverted_state_dict = rank_local_checkpoint.try_load_rank_local_tp_state_dict(
        model, weight_dir_list, param_names_mapping_fn,
    )

# preconverted_state_dict가 있으면 iterator는 비워둠
if preconverted_state_dict is None:
    weight_iterator = safetensors_weights_iterator(weight_dir_list)
else:
    weight_iterator = iter(())

load_model_from_full_model_state_dict(
    model,
    weight_iterator,
    preconverted_state_dict=preconverted_state_dict,
    # ... 생략
)

2. DTensor 및 Shard 메타데이터 활용

새로운 로더는 PyTorch의 DTensorDeviceMesh 정보를 활용하여 각 파라미터가 어떤 오프셋으로 샤딩되어야 하는지 계산합니다. safetensors 라이브러리의 slice 기능을 사용하면 파일 전체를 메모리에 올리지 않고도 특정 바이트 범위만 읽어올 수 있습니다.

After (load_model_from_full_model_state_dict 내부):

if is_rank_local_fsdp_shard:
    # 이미 랭크에 맞게 잘려진(sliced) 텐서를 DTensor로 변환
    local_tensor = full_tensor.to(device=checkpoint_load_device, dtype=target_dtype)
    sharded_tensor = dist_tensor.DTensor.from_local(
        local_tensor,
        meta_sharded_param.device_mesh,
        meta_sharded_param.placements,
        run_check=False,
        shape=meta_sharded_param.shape,
        stride=meta_sharded_param.stride(),
    )

이 과정에서 QKV(Query, Key, Value)나 W13(Gate, Up Proj)처럼 여러 파라미터가 하나로 합쳐진(Merged) 레이어의 경우, 소스 텐서에서 각 부분을 먼저 슬라이싱한 뒤 로컬에서 결합(Concatenation)하는 정교한 처리가 추가되었습니다.

왜 이게 좋은가?

1. I/O 대역폭 절감 및 속도 향상

PR 본문의 벤치마크 결과에 따르면, Z-Image TP2 설정에서 로딩 시간이 43.2초에서 34.4초로 약 20.3% 단축되었습니다. 특히 FSDP2 설정에서는 22.8%의 성능 향상을 보였습니다. 랭크당 읽기 데이터 양이 22.9 GiB에서 약 11.5 GiB로 줄어든 것이 결정적이었습니다.

2. 메모리 효율성

전체 체크포인트를 스트리밍할 때는 사용하지 않는 텐서도 일시적으로 메모리에 로드될 수 있습니다. 랭크 로컬 로딩은 필요한 데이터만 메모리에 올리므로, 로딩 과정에서의 Peak Memory 사용량을 억제할 수 있습니다.

3. 견고한 폴백(Fallback) 설계

모든 케이스에 대해 이 최적화를 강제하지 않았습니다. 양자화(Quantization)된 모델이나 커스텀 레이아웃을 사용하는 모델(예: FLUX.2-klein-4B)은 기존의 전체 로딩 방식을 유지하도록 설계되었습니다. 이는 성능 최적화를 시도하되 시스템의 안정성을 해치지 않는 훌륭한 엔지니어링 사례입니다.

일반적인 교훈

  1. 데이터 포맷의 특성을 이해하라: safetensors가 단순한 가중치 저장소가 아니라, 효율적인 랜덤 액세스를 지원하는 포맷이라는 점을 활용해 병목을 해결했습니다.
  2. 분산 환경에서의 중복 제거: 분산 시스템에서 모든 노드가 동일한 작업을 수행하고 있다면, 그것은 대개 최적화의 기회입니다. "각자 자기 것만 하기"는 분산 컴퓨팅의 기본 원칙입니다.
  3. 메타데이터 기반의 선제적 계산: 런타임에 데이터를 자르는 대신, 로드 시점에 DTensor 메타데이터를 기반으로 필요한 범위를 미리 계산하여 I/O 자체를 줄였습니다.

결론

이번 업데이트는 SGLang이 대규모 Diffusion 모델을 더 빠르고 효율적으로 서빙할 수 있는 기반을 마련했습니다. 특히 클라우드 환경처럼 스토리지 I/O 비용이 비싼 곳에서 그 효과는 더욱 극대화될 것입니다.

참고 자료

⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.

댓글

관련 포스트

PR Analysis 의 다른글