[sglang] SGLang, NVIDIA Blackwell GPU를 위한 Wan2.2 모델의 ML P 연산 최적화
PR 링크: sgl-project/sglang#37075 상태: Merged | 변경: +411 / -5
들어가며
최근 SGLang 프로젝트에서는 NVIDIA Blackwell GPU 아키텍처에 최적화된 새로운 커널 융합(Kernel Fusion) 기법을 도입했습니다. 이번 PR은 특히 Wan2.2 모델의 fc_in 경로에서 발생하는 메모리 병목 현상을 해결하는 데 중점을 둡니다. 기존에는 NVFP4 GEMM -> BF16 bias add -> tanh GELU 와 같이 세 단계로 나뉘어 처리되던 연산을, NVFP4 GEMM 이후의 bias add와 tanh GELU 연산을 하나의 CUDA 커널로 융합하여 처리 속도를 향상시키는 것을 목표로 합니다. 이는 특히 고품질(high quality) 생성 시나리오에서 렌더링 시간을 단축하는 데 기여합니다.
코드 분석
이번 PR의 핵심은 새로운 JIT CUDA 커널 bias_gelu_tanh_kernel의 도입과 이를 기존 SGLang 파이프라인에 통합하는 것입니다.
1. 새로운 JIT CUDA 커널: bias_gelu_tanh_kernel (python/sglang/kernels/jit/csrc/elementwise/bias_gelu.cuh)
이 커널은 기존의 bias add와 GELU 연산을 하나로 융합합니다. 주요 특징은 다음과 같습니다:
- FP16/BF16 지원: 다양한 반정밀도 부동소수점 연산을 지원합니다.
- Blackwell 최적화: Blackwell 아키텍처에서 효율적인 32-byte 벡터 연산을 지원하며, 이전 아키텍처를 위한 16-byte 벡터도 지원합니다.
- 점유율 제한(Occupancy-capped) 및 영구 런칭(Persistent Launch): GPU 자원을 효율적으로 사용하기 위한 최적화 기법이 적용되었습니다.
- PDL(Process Data Loading) 지원: 비동기 데이터 로딩을 지원하여 GPU 연산과 데이터 전송 간의 병목을 줄입니다.
- 명시적 반정밀도 반올림 경계:
bias add연산 후, GELU 연산 전에 명시적으로 반정밀도(half-precision)로 반올림하여 기존의 eager 모드에서의 연산 순서와 동일한 수치적 결과를 보장합니다.
// Before (Conceptual - separate operations)
// ... GEMM output ...
vec_t biased = x + b; // BF16 bias add
result[i] = device::cast<T>(gelu_tanh(device::cast<fp32_t>(biased))); // tanh GELU
// After (Fused kernel)
#pragma unroll
for (int i = 0; i < kVecN; ++i) {
const float x_f32 = device::cast<fp32_t>(x[i]);
const float bias_f32 = device::cast<fp32_t>(b[i]);
// Explicit rounding boundary after bias add
const T biased = device::cast<T>(x_f32 + bias_f32);
result[i] = device::cast<T>(gelu_tanh(device::cast<fp32_t>(biased)));
}
위 코드는 bias add 연산 후 T 타입으로 캐스팅하여 반올림 경계를 명시적으로 설정하고, 그 결과를 다시 fp32_t로 캐스팅하여 gelu_tanh 함수에 적용하는 과정을 보여줍니다. 이는 기존의 F.gelu(input + bias, approximate="tanh")와 동일한 수치적 결과를 보장하기 위함입니다.
2. 융합 사이트 통합 (python/sglang/kernels/ops/diffusion/sites/nvfp4_bias_gelu_site.py)
이 파일은 새로운 융합 기능을 기존 SGLang의 품질 게이트(Quality Gate) 시스템에 통합합니다.
QualityGatedFusion: 품질 설정(`quality=
참고 자료
- https://pytorch.org/docs/stable/generated/torch.nn.functional.gelu.html
- https://docs.nvidia.com/deeplearning/cuda/parallel-thread-execution/index.html#ld-st-instructions
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [sglang] FLUX.2 모델 성능 최적화: Token Concatenation과 NVFP4 양자화의 커널 융합
- [flashinfer] NVIDIA Blackwell 아키텍처를 위한 FlashInfer의 Router GEMM 최적화
- [flashinfer] FlashInfer, Blackwell 아키텍처를 위한 Recurrent KDA Prefill 최적화: Small-BH 커널 도입
- [flashinfer] FlashInfer, Blackwell GPU를 위한 Gated MoE 커널 최적화로 성능 대폭 향상
- [flashinfer] Blackwell NVFP4 양자화 최적화: TMA OOB Zero-fill을 이용한 메모리 복사 오버헤드 제거
PR Analysis 의 다른글
- 이전글 [onnxruntime] ONNX Runtime CUDA 커널 최적화: Speculative Decoding을 위한 GEMV 확장
- 현재글 : [sglang] SGLang, NVIDIA Blackwell GPU를 위한 Wan2.2 모델의 ML P 연산 최적화
- 다음글 [vllm] vLLM의 ROCm 환경에서 듀얼 스트림 디코드를 통한 성능 최적화
댓글