[sglang] SGLang: LongCat-Image DiT의 FFN 연산 최적화 - Tanh-GELU 퓨전 적용
PR 링크: sgl-project/sglang#36322 상태: Merged | 변경: +28 / -4
들어가며
최신 Diffusion Transformer(DiT) 모델들은 추론 과정에서 수많은 연산을 수행하며, 특히 FFN(Feed-Forward Network) 블록 내의 GEMM(General Matrix Multiply) 연산과 그 뒤를 잇는 활성화 함수(Activation Function)는 전체 지연 시간(Latency)의 상당 부분을 차지합니다.
이번 SGLang PR에서는 LongCat-Image 모델의 FFN 구조에서 up-projection GEMM과 tanh-GELU 활성화 함수를 별도로 실행하던 기존 방식을, cublasLt의 epilogue 기능을 활용하여 하나의 커널로 퓨전(Fusion)하는 최적화를 진행했습니다. 이를 통해 불필요한 메모리 읽기/쓰기(Memory Bound)를 줄이고 전체적인 추론 속도를 향상시켰습니다.
코드 분석
이번 변경은 python/sglang/multimodal_gen/runtime/models/dits/longcat_image.py 파일 내의 _LongCatFFN과 _SingleTransformerBlock 클래스에 집중되어 있습니다.
1. 퓨전 사이트 등록 (mark_fused_gelu_site)
기존에는 활성화 함수가 별도의 연산으로 존재했으나, 이를 퓨전 가능한 지점으로 명시적으로 등록했습니다.
Before:
self.act = nn.GELU(approximate="tanh")
After:
mark_fused_gelu_site(self.net[0], "proj")
mark_fused_gelu_site를 통해 해당 모듈이 fused_linear_gelu_tanh를 사용할 수 있는 후보임을 시스템에 알립니다.
2. Forward 패스에서의 조건부 퓨전 적용
실제 추론 시점에 fused_gelu_active()와 can_use_linear_gelu()를 체크하여 퓨전된 커널을 사용할지 결정합니다.
Before:
hidden_states, _ = self.net[0]["proj"](hidden_states)
hidden_states = self.act(hidden_states)
After:
if fused_gelu_active(self.net[0]) and can_use_linear_gelu(proj, hidden_states):
hidden_states = fused_linear_gelu_tanh(hidden_states, proj.weight, proj.bias)
else:
hidden_states, _ = proj(hidden_states)
hidden_states = self.act(hidden_states)
이 방식은 기존 로직을 유지하면서도, 최적화가 가능한 환경(예: 특정 하드웨어 및 설정)에서만 효율적인 퓨전 커널을 선택적으로 실행하게 합니다.
왜 이게 좋은가
이 최적화는 메모리 대역폭을 절약하는 전형적인 '커널 퓨전' 사례입니다. GEMM 연산 직후에 활성화 함수를 적용할 때, 중간 결과를 GPU 메모리에 썼다가 다시 읽어오는 오버헤드가 발생합니다. 퓨전을 적용하면 GEMM 연산 중에 레지스터 수준에서 활성화 함수를 즉시 적용하므로 이 오버헤드가 제거됩니다.
성능 지표:
- GELU 영역: 50단계 총합 기준 9638ms에서 9044ms로 약 6.2% 성능 향상.
- End-to-End Denoise: 전체 추론 시간 기준 약 1.5% 향상.
교훈:
- Epilogue 활용:
cublasLt와 같은 라이브러리는 연산 뒤에 붙는 간단한 활성화 함수를 퓨전할 수 있는 기능을 제공합니다. 이를 적극 활용하면 별도의 커널 작성 없이도 큰 성능 이득을 볼 수 있습니다. - 조건부 실행: 모든 상황에서 퓨전을 강제하기보다,
can_use_linear_gelu와 같은 유틸리티를 통해 안전한 경우에만 적용하는 설계가 유지보수와 안정성 측면에서 유리합니다.
결론
이번 PR은 복잡한 커널 수정 없이도 기존의 인프라를 활용하여 모델의 추론 성능을 끌어올린 좋은 사례입니다. 특히 4000번 이상 반복되는 연산 지점을 타겟팅하여 실질적인 사용자 경험 개선을 이끌어냈습니다.
참고 자료
- https://docs.nvidia.com/cuda/cublas/index.html#cublaslt-matmul
- https://pytorch.org/docs/stable/generated/torch.nn.GELU.html
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [sglang] DeepSeek NextN을 위한 Fused EH Norm 최적화: 커널 융합으로 성능 극대화하기
- [sglang] SGLang LTX-2 VAE 디코딩 성능 최적화: channels_last_3d 도입으로 4.5배 속도 향상
- [sglang] SGLang Triton 커널 최적화: libdevice.tanh 도입과 2D Strided Tensor 지원
- [sglang] SGLang: LFM2-MoE 모델을 위한 SM90 커널 퓨전 최적화 분석
- [sglang] FLUX.2 모델 성능 최적화: Token Concatenation과 NVFP4 양자화의 커널 융합
PR Analysis 의 다른글
- 이전글 [vllm] vLLM, Pixtral 모델의 멀티모달 인코더 어텐션 최적화: Packed Sequence Metadata 도입
- 현재글 : [sglang] SGLang: LongCat-Image DiT의 FFN 연산 최적화 - Tanh-GELU 퓨전 적용
- 다음글 [flashinfer] FlashInfer, CUDA 커널 최적화를 통한 LLM 추론 속도 향상: Recurrence-Piece Persistent M128 도입
댓글