[sglang] SGLang: Wan VAE의 RMSNorm 및 SiLU 연산 융합을 통한 추론 가속
PR 링크: sgl-project/sglang#33546 상태: Merged | 변경: +435 / -53
들어가며
최근 SGLang 프로젝트에서 Wan VAE 모델의 추론 성능을 극대화하기 위한 최적화가 진행되었습니다. 특히 FastWan2.2-TI2V-5B와 같은 모델은 전체 추론 시간의 약 60%를 VAE 디코딩 단계가 차지하고 있습니다. 기존에는 RMSNorm과 SiLU가 별도의 aten 커널로 실행되어 메모리 대역폭 낭비와 커널 실행 오버헤드가 발생했습니다. 본 PR은 이 두 연산을 하나의 Triton 커널로 융합하여 성능을 개선하고, quality=high 옵션을 통해 선택적으로 적용함으로써 정확도와 속도 사이의 균형을 맞췄습니다.
코드 분석
1. Triton 커널 구현 (python/sglang/kernels/ops/diffusion/triton/wan_rmsnorm_silu.py)
핵심은 channels_last_3d 레이아웃을 활용하여 메모리 접근을 최적화한 것입니다. 이전에는 RMSNorm과 SiLU가 개별적으로 호출되었으나, 이를 하나의 커널로 합쳐 중간 결과물의 메모리 쓰기를 줄였습니다.
# Before (Conceptual)
y = F.normalize(x, dim=1) * scale * gamma + bias
y = F.silu(y)
# After (Triton Kernel)
# Fused RMSNorm + SiLU in a single pass
# ...
y = y * tl.sigmoid(y)
tl.store(out_ptr + out_base + offsets * out_stride_c, y, mask=mask)
또한, autocast 환경에서 bf16 입력과 fp32 가중치를 사용하는 경우를 고려하여 aten의 동작과 동일하게 fp32로의 승격(promotion)을 커널 내에서 재현했습니다.
2. Fast Path 게이트 및 래퍼 (python/sglang/multimodal_gen/runtime/models/vaes/wan_vae_cuda_opt.py)
FusedWanRMSNormSiLU 모듈을 도입하여 quality == "high"일 때만 최적화된 커널을 사용하도록 설계했습니다. lossless 모드에서는 기존 연산 경로를 그대로 사용하여 비트 단위의 정확도(bit-exactness)를 보장합니다.
# 래퍼를 통한 분기 처리
if gate.enabled and can_use_wan_rmsnorm_silu(x, self.gamma, self.bias):
return wan_rmsnorm_silu(x, self.gamma, self.bias, self.rms_scale)
else:
return F.silu(self.norm(x))
왜 이게 좋은가
이번 최적화의 핵심은 메모리 대역폭 효율화입니다. 29개의 잔차 블록(residual block)과 출력 헤드에서 발생하는 연쇄적인 RMSNorm -> SiLU 호출을 단일 커널로 융합함으로써, H200 GPU 기준 DecodingStage에서 약 11.5%의 속도 향상을 달성했습니다.
- 성능 수치:
FastWan2.2-TI2V-5B모델에서DecodingStage시간이 5.77초에서 5.11초로 단축되었습니다. - 교훈:
- 선택적 최적화: 모든 경우에 최적화를 강제하지 않고,
quality옵션을 통해 정확도 보존이 필요한 경우와 속도가 중요한 경우를 분리했습니다. - 비트 단위 정확도 검증:
lossless모드에서는sha256해시 비교를 통해 기존 구현과 동일한 출력을 보장하여 CI 안정성을 확보했습니다. - 컴파일러와의 협업:
torch.compile이 활성화된 경우, Inductor가 이미 최적화를 수행하므로 불필요한 커스텀 커널 호출을 피하도록is_compiling()체크를 추가했습니다.
- 선택적 최적화: 모든 경우에 최적화를 강제하지 않고,
리뷰 피드백 반영
리뷰 과정에서 state_dict 로드 시 파라미터 이름이 변경되는 문제를 발견했습니다. 이를 해결하기 위해 래퍼 모듈에서 gamma와 bias를 직접 등록하여 FQN(Fully Qualified Name)이 유지되도록 수정했습니다. 이는 모델 가중치 전송 시 발생할 수 있는 호환성 문제를 방지하는 중요한 조치였습니다.
참고 자료
- https://triton-lang.org/main/index.html
- https://pytorch.org/docs/stable/generated/torch.nn.RMSNorm.html
- https://pytorch.org/docs/stable/generated/torch.nn.SiLU.html
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
PR Analysis 의 다른글
- 이전글 [triton] Triton AMD GPU 최적화: Warp-based Split-K 도입을 통한 MQA 성능 향상
- 현재글 : [sglang] SGLang: Wan VAE의 RMSNorm 및 SiLU 연산 융합을 통한 추론 가속
- 다음글 [flashinfer] FlashInfer의 Blackwell 아키텍처 최적화: CAKE 기반 TinyGEMM2 커널 도입
댓글