[sglang] Diffusion DiT 모델의 FFN 성능 최적화: cublasLt GELU Epilogue 융합
PR 링크: sgl-project/sglang#33536 상태: Merged | 변경: +317 / -0
들어가며
Diffusion DiT(Diffusion Transformer) 모델의 FeedForward Network(FFN)는 일반적으로 gelu(up_proj(x), approximate="tanh") 연산을 수행합니다. 기존 방식은 up_proj GEMM을 먼저 수행하고, 그 결과를 HBM(High Bandwidth Memory)에 쓴 뒤, 별도의 GELU 커널을 실행하는 구조였습니다. 이 과정은 불필요한 커널 호출과 HBM 왕복(round-trip)을 유발하여 성능 병목이 됩니다. 본 PR은 torch._addmm_activation을 활용하여 GEMM과 GELU를 하나의 cublasLt 에필로그(epilogue)로 융합함으로써 이 문제를 해결합니다.
코드 분석
1. 커널 헬퍼 구현 (python/sglang/kernels/ops/diffusion/fused_linear_gelu.py)
새로운 헬퍼 모듈은 cublasLt의 GELU 에필로그를 활용하는 fused_linear_gelu_tanh 함수를 정의합니다.
@register_custom_op(
op_name="diffusion_fused_linear_gelu_tanh",
mutates_args=[],
fake_impl=_fused_linear_gelu_tanh_fake,
)
def fused_linear_gelu_tanh(
x: torch.Tensor, weight: torch.Tensor, bias: torch.Tensor
) -> torch.Tensor:
x2d = x.reshape(-1, x.shape[-1])
out = torch._addmm_activation(bias, x2d, weight.t(), use_gelu=True)
return out.view(*x.shape[:-1], weight.shape[0])
이 코드는 torch._addmm_activation을 호출하여 GEMM 연산 직후에 GELU를 적용합니다. register_custom_op를 통해 torch.compile 환경에서도 그래프 브레이크 없이 단일 연산으로 처리되도록 설계되었습니다.
2. 모델 적용 (python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py)
기존의 분리된 연산을 융합된 커널로 교체하되, quality="high" 옵션이 켜져 있을 때만 동작하도록 게이팅 로직을 추가했습니다.
Before:
hidden_states, _ = self.proj(hidden_states)
return F.gelu(hidden_states, approximate="tanh")
After:
if self._sgl_fused_gelu_enabled and can_fuse_linear_gelu(self.proj, hidden_states):
return fused_linear_gelu_tanh(hidden_states, self.proj.weight, self.proj.bias)
hidden_states, _ = self.proj(hidden_states)
return F.gelu(hidden_states, approximate="tanh")
왜 이게 좋은가
이 최적화는 단순히 연산을 합치는 것을 넘어, 메모리 대역폭 제한(bandwidth-bound) 문제를 해결합니다. cublasLt 에필로그를 사용하면 중간 결과물을 HBM에 쓰지 않고 레지스터/L2 캐시 수준에서 GELU를 적용할 수 있습니다.
- 성능 향상: Qwen-Image 1024^2 모델에서 Denoising 단계가 12.36초에서 12.05초로 약 2.5% 단축되었습니다.
- 품질 보존:
quality="high"모드에서만 활성화하여, 기본값인lossless모드에서는 기존의 비트 단위 정확도(bit-exact)를 유지합니다. 이는cublasLt의 GELU 근사치가 표준F.gelu와 매우 유사하다는 점을 활용한 전략적 선택입니다. - 일반적 교훈: 커널 융합(Kernel Fusion)은 특히 Transformer의 MLP 블록과 같이 연산 강도가 낮은 곳에서 큰 효과를 발휘합니다. 특히
torch._addmm_activation과 같은 저수준 API를 활용하면 커널 런타임 오버헤드를 획기적으로 줄일 수 있습니다.
참고 자료
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
PR Analysis 의 다른글
- 이전글 [triton] Triton FPSAN의 MMA 에뮬레이션 오버헤드 최적화 분석
- 현재글 : [sglang] Diffusion DiT 모델의 FFN 성능 최적화: cublasLt GELU Epilogue 융합
- 다음글 [flashinfer] Blackwell NVFP4 양자화 최적화: TMA OOB Zero-fill을 이용한 메모리 복사 오버헤드 제거
댓글