[sglang] LingBot Video 성능 개선: 수동 RMSNorm 체인을 Triton 커널로 최적화하기
PR 링크: sgl-project/sglang#35969 상태: Merged | 변경: +174 / -0
들어가며
최근 sglang-project/sglang 레포지토리에서는 LingBot Video 모델의 성능 개선을 위한 흥미로운 Pull Request(PR)가 올라왔습니다. 이 PR의 핵심은 quality=high 설정에서 LingBot Video 모델이 사용하던 수동으로 구현된 RMSNorm(Root Mean Square Normalization) 연산 체인을, 이미 잘 최적화된 diffusion Triton RMSNorm 커널로 대체하는 것입니다. 기존의 RMSNorm 구현은 여러 개의 개별 연산(cast, square, mean, rsqrt, multiply)으로 구성되어 있어 GPU 활용 측면에서 비효율적이었습니다. 이 PR은 이러한 비효율성을 제거하고, 특히 H200 GPU 환경에서 상당한 성능 향상을 가져왔습니다.
본 블로그 글에서는 이 PR의 코드 변경 사항을 상세히 분석하고, 왜 이러한 변경이 성능 향상으로 이어졌는지, 그리고 이 최적화가 주는 일반적인 교훈은 무엇인지 살펴보겠습니다.
코드 분석
이 PR은 LingBot Video 모델의 RMSNorm 연산을 최적화하기 위해 여러 파일에서 변경을 가했습니다. 주요 변경 사항은 다음과 같습니다.
1. docs/docs/sglang-diffusion/fused_kernels.mdx 업데이트
이 파일은 sglang의 융합 커널(fused kernels)에 대한 문서를 관리합니다. PR은 새로운 융합 기능으로 'LingBot Video fused RMSNorm'을 추가하고, 이 기능이 기존의 수동 구현을 대체한다는 설명을 덧붙였습니다.
변경 전 (일부):
... (기존 내용) ...
변경 후 (일부):
+| LingBot Video fused RMSNorm | Replaces the handwritten cast, square, mean, rsqrt, and multiply chain with existing Triton RMSNorm kernels |
... (기존 내용) ...
이는 새로운 최적화 기능이 추가되었음을 문서화하여 사용자들에게 알리는 중요한 단계입니다.
2. python/sglang/kernels/ops/diffusion/__init__.py - 커널 등록
이 파일은 diffusion 연산에 사용되는 다양한 커널들을 등록하고 관리합니다. PR은 LingBot Video RMSNorm과 관련된 새로운 함수들을 등록하여 시스템이 이를 인식하고 사용할 수 있도록 합니다.
변경 전 (일부):
... (기존 내용) ...
변경 후 (일부):
+ "lingbot_video_rmsnorm_active": "sites.lingbot_video_rmsnorm_site",
+ "mark_lingbot_video_rmsnorm_site": "sites.lingbot_video_rmsnorm_site",
+ "mount_lingbot_video_rmsnorm": "sites.lingbot_video_rmsnorm_site",
+ "try_lingbot_video_rmsnorm": "sites.lingbot_video_rmsnorm_site",
+ "unmount_lingbot_video_rmsnorm": "sites.lingbot_video_rmsnorm_site",
...
mark_, mount_, try_, unmount_, active와 같은 함수들은 해당 RMSNorm 연산을 활성화하거나 비활성화하고, 실제 연산을 수행하는 데 필요한 인터페이스를 제공합니다. 이는 quality=high 설정 시 최적화된 커널을 사용하도록 하는 메커니즘의 일부입니다.
3. python/sglang/kernels/ops/diffusion/sites/lingbot_video_rmsnorm_site.py - 새로운 RMSNorm 사이트 구현
이 파일은 LingBot Video 모델의 RMSNorm 연산을 위한 새로운 최적화 사이트(site)를 정의합니다. QualityGatedFusion 클래스를 사용하여 quality=high 설정에서만 이 최적화가 활성화되도록 합니다.
새로 추가된 파일 내용 (핵심 부분):
# ... (import 및 QualityGatedFusion 초기화) ...
def mark_lingbot_video_rmsnorm_site(module: nn.Module) -> None:
"""Mark a LingBot RMSNorm module; it starts on the reference path."""
_FUSION.mark(module)
def lingbot_video_rmsnorm_active(module: nn.Module) -> bool:
return _FUSION.is_enabled(module)
def _site_reject_reason(site: nn.Module) -> str | None:
# ... (Triton 사용 가능 여부, weight dtype, shape 등 검사) ...
return None
def mount_lingbot_video_rmsnorm(root: nn.Module) -> bool:
return _FUSION.mount(root, reject_reason=_site_reject_reason, logger=logger)
def unmount_lingbot_video_rmsnorm(root: nn.Module) -> None:
_FUSION.unmount(root)
def try_lingbot_video_rmsnorm(
site: nn.Module,
hidden_states: torch.Tensor,
weight: torch.Tensor,
eps: float,
) -> torch.Tensor | None:
"""Return the quality-gated RMSNorm result, or ``None`` to fall back."""
if not (
_FUSION.is_enabled(site)
and hidden_states.is_cuda
and hidden_states.dtype in (torch.float16, torch.bfloat16)
and hidden_states.is_contiguous()
and hidden_states.shape[-1] == weight.numel()
and weight.is_cuda
and weight.device == hidden_states.device
and weight.dtype in (hidden_states.dtype, torch.float32)
):
return None
hidden_size = hidden_states.shape[-1]
if weight.dtype == torch.float32 and hidden_size > 128:
from sglang.kernels.ops.diffusion.norm.norm_triton import norm_infer
shape = hidden_states.shape
return norm_infer(
hidden_states.view(-1, hidden_size),
weight,
bias=None,
eps=eps,
is_rms_norm=True,
).view(shape)
from sglang.kernels.ops.diffusion.norm.rmsnorm_onepass_triton import (
triton_one_pass_rms_norm,
)
return triton_one_pass_rms_norm(hidden_states, weight, eps)
try_lingbot_video_rmsnorm 함수는 최적화 커널을 사용할 수 있는 조건(CUDA 사용, 특정 데이터 타입, contiguous 텐서 등)을 검사합니다. 조건이 만족되면, weight의 dtype과 hidden_size에 따라 두 가지 Triton 커널 중 하나를 선택하여 실행합니다:
weight.dtype == torch.float32이고hidden_size > 128인 경우:norm_infer커널 사용 (Row Kernel).- 그 외의 경우:
triton_one_pass_rms_norm커널 사용 (Tiled One-Pass Kernel).
이러한 조건부 로직은 다양한 하드웨어 및 모델 설정에 대해 최적의 성능을 보장하기 위함입니다.
4. python/sglang/multimodal_gen/runtime/models/dits/lingbot_video_moe.py - RMSNorm 모듈 수정
LingBot Video 모델의 RMSNorm 모듈(RMSNorm 클래스)에서 forward 메서드를 수정하여 새로 구현된 try_lingbot_video_rmsnorm 함수를 호출하도록 변경했습니다.
변경 전 (일부):
--- a/python/sglang/multimodal_gen/runtime/models/dits/lingbot_video_moe.py
+++ b/python/sglang/multimodal_gen/runtime/models/dits/lingbot_video_moe.py
@@ -75,8 +79,15 @@ def __init__(self, dim: int, eps: float = 1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.variance_epsilon = eps
+ mark_lingbot_video_rmsnorm_site(self)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
+ fused = try_lingbot_video_rmsnorm(
+ self, hidden_states, self.weight, self.variance_epsilon
+ )
+ if fused is not None:
+ return fused
+
input_dtype = hidden_states.dtype
hidden_states = hidden_states.to(torch.float32)
variance = hidden_states.pow(2).mean(-1, keepdim=True)
__init__ 메서드에서 mark_lingbot_video_rmsnorm_site(self)를 호출하여 해당 모듈이 최적화 대상임을 표시하고, forward 메서드 시작 부분에서 try_lingbot_video_rmsnorm을 호출합니다. 만약 try_lingbot_video_rmsnorm이 None이 아닌 값을 반환하면, 이는 최적화된 커널이 성공적으로 실행되었음을 의미하므로 해당 결과를 즉시 반환합니다. 그렇지 않으면, 기존의 수동 구현(fallback path)으로 진행됩니다.
5. python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py - 융합 커널 관리
이 파일은 diffusion 파이프라인의 denoising 단계에서 사용되는 융합 커널들을 관리합니다. PR은 mount_lingbot_video_rmsnorm과 unmount_lingbot_video_rmsnorm 함수를 융합 커널 목록에 추가하여, quality=high 설정 시 이 최적화가 활성화되도록 합니다.
변경 전 (일부):
... (기존 내용) ...
변경 후 (일부):
+ (
+ "LingBot Video fused RMSNorm",
+ mount_lingbot_video_rmsnorm,
+ unmount_lingbot_video_rmsnorm,
+ ),
...
이는 시스템이 이 새로운 융합 커널을 인지하고, 적절한 시점에 마운트(활성화) 및 언마운트(비활성화)할 수 있도록 합니다.
6. test/registered/kernels/ops/diffusion/test_sites.py - 테스트 케이스 추가
새로운 최적화 기능의 정확성과 안정성을 검증하기 위해 test_lingbot_video_rmsnorm_quality_path_and_guards라는 테스트 케이스가 추가되었습니다. 이 테스트는 다양한 조건(FP32, BF16 가중치, CUDA 사용 등)에서 try_lingbot_video_rmsnorm 함수가 올바르게 작동하는지, 그리고 가드(guard) 조건들이 제대로 작동하는지를 검증합니다.
테스트 코드 (일부):
@requires_cuda
@torch.no_grad()
def test_lingbot_video_rmsnorm_quality_path_and_guards():
# ... (테스트 설정) ...
assert (
lingbot_video_rmsnorm.try_lingbot_video_rmsnorm(
site, hidden_states, site.weight, 1e-6
)
is None
)
assert lingbot_video_rmsnorm.mount_lingbot_video_rmsnorm(site)
output = lingbot_video_rmsnorm.try_lingbot_video_rmsnorm(
site, hidden_states, site.weight, 1e-6
)
# ... (참조 값과 비교) ...
torch.testing.assert_close(output, reference, atol=2e-2, rtol=2e-2)
# ... (BF16 가중치 테스트) ...
# ... (가드 조건 테스트) ...
lingbot_video_rmsnorm.unmount_lingbot_video_rmsnorm(site)
assert not lingbot_video_rmsnorm.lingbot_video_rmsnorm_active(site)
이 테스트는 최적화가 의도한 대로 작동하며, 예상치 못한 입력이나 조건에서는 fallback path로 올바르게 전환됨을 보장합니다.
왜 이게 좋은가?
이 PR의 가장 큰 장점은 성능 향상입니다. H200 GPU에서의 벤치마크 결과는 이를 명확히 보여줍니다.
H200 결과 요약:
- Denoise 시간: 3.9293s (main) -> 2.9145s (PR) (약 25.65% 향상)
- Saved-request e2e 시간: 4.5685s (main) -> 3.5533s (PR) (약 22.05% 향상)
- Peak reserved memory: 변화 없음
이러한 성능 향상은 다음과 같은 이유로 가능했습니다.
- Triton 커널 사용: Triton은 GPU 커널을 효율적으로 작성하기 위한 언어로, 복잡한 연산을 단일 커널로 융합하고 하드웨어에 최적화된 방식으로 실행할 수 있습니다. 기존의 수동 RMSNorm 체인은
pow,mean,rsqrt,multiply등 여러 개의 CUDA 커널 호출로 이루어져 있었을 가능성이 높습니다. 각 커널 호출에는 오버헤드가 발생하며, GPU의 연산 유닛이 데이터를 처리하는 동안 다른 연산으로 전환되지 못하고 대기하는 경우가 많습니다. Triton 커널은 이러한 연산들을 하나로 묶어(fused) GPU에서 더 적은 오버헤드로, 더 효율적으로 처리할 수 있게 합니다. - 메모리 접근 최적화: Triton 커널은 데이터의 로딩 및 저장 방식을 최적화하여 메모리 대역폭 사용을 효율화할 수 있습니다. 특히
norm_infer와rmsnorm_onepass_triton과 같은 커널들은 tiled execution, shared memory 활용 등을 통해 메모리 접근 패턴을 개선합니다. - 조건부 최적화:
quality=high설정에서만 이 최적화가 적용되고,quality=lossless에서는 기존 구현을 유지하는 것은 매우 현명한 접근입니다. 이는 성능과 결과의 정확성(bit-exactness) 사이의 균형을 맞추기 위함입니다. 또한,try_lingbot_video_rmsnorm함수 내에서 다양한 조건(dtype, shape, CUDA 사용 여부 등)을 검사하여 최적화 커널을 적용할 수 있을 때만 적용하고, 그렇지 않을 경우 fallback path로 전환하는 것은 안정성을 높입니다.
일반적인 교훈:
- 표준 라이브러리/프레임워크 활용: PyTorch나 Triton과 같이 잘 최적화된 라이브러리나 커널을 활용하는 것은 직접 구현하는 것보다 훨씬 효율적이고 안정적입니다. 특히 고성능 컴퓨팅에서는 이러한 최적화된 빌딩 블록을 재사용하는 것이 중요합니다.
- 융합(Fusion)의 힘: 여러 개의 작은 연산을 하나의 큰 연산으로 묶는 융합은 GPU 성능 향상의 핵심 기법 중 하나입니다. 이는 커널 호출 오버헤드를 줄이고, 데이터 재사용성을 높이며, 메모리 접근 패턴을 개선합니다.
- 조건부 최적화 및 Fallback: 모든 상황에서 단일 최적화 기법이 최선은 아닙니다. 특정 조건(예:
quality설정, 데이터 타입, 하드웨어)에 따라 최적화 기법을 선택하고, 최적화가 불가능할 경우 안정적으로 동작하는 fallback path를 제공하는 것이 중요합니다. - 정확한 벤치마킹 및 테스트: 성능 개선 효과를 정량적으로 입증하고, 기능의 정확성을 보장하기 위해 철저한 벤치마킹과 테스트 케이스 작성이 필수적입니다.
결론
이 PR은 LingBot Video 모델에서 비효율적인 수동 RMSNorm 연산을 고성능 Triton 커널로 대체함으로써, 특히 quality=high 설정에서 상당한 성능 향상을 달성했습니다. 이는 최적화된 라이브러리 활용, 융합 커널의 이점, 그리고 조건부 최적화 전략의 중요성을 잘 보여주는 사례입니다. 이러한 개선은 sglang 라이브러리의 전반적인 성능을 향상시키고, 사용자들에게 더 빠르고 효율적인 AI 모델 추론 경험을 제공할 것입니다.
References
- Triton Documentation
- PyTorch RMSNorm Documentation (참고용, 실제 사용된 커널은 sglang 내부 구현)
- sglang-project/sglang GitHub Repository
참고 자료
- https://triton-lang.org/
- https://pytorch.org/docs/stable/generated/torch.nn.RMSNorm.html
- https://github.com/sgl-project/sglang
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
PR Analysis 의 다른글
- 이전글 [triton] Triton 컴파일러의 Reduce 연산 최적화: OptimizeThreadLocality 개선 분석
- 현재글 : [sglang] LingBot Video 성능 개선: 수동 RMSNorm 체인을 Triton 커널로 최적화하기
- 다음글 [open-webui] Open WebUI 스트리밍 성능 190배 개선: O(N^2)에서 O(N)으로의 최적화
댓글