본문으로 건너뛰기

[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 addtanh 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 addGELU 연산을 하나로 융합합니다. 주요 특징은 다음과 같습니다:

  • 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=

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글