본문으로 건너뛰기

[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_rmsnormNone이 아닌 값을 반환하면, 이는 최적화된 커널이 성공적으로 실행되었음을 의미하므로 해당 결과를 즉시 반환합니다. 그렇지 않으면, 기존의 수동 구현(fallback path)으로 진행됩니다.

5. python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py - 융합 커널 관리

이 파일은 diffusion 파이프라인의 denoising 단계에서 사용되는 융합 커널들을 관리합니다. PR은 mount_lingbot_video_rmsnormunmount_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: 변화 없음

이러한 성능 향상은 다음과 같은 이유로 가능했습니다.

  1. Triton 커널 사용: Triton은 GPU 커널을 효율적으로 작성하기 위한 언어로, 복잡한 연산을 단일 커널로 융합하고 하드웨어에 최적화된 방식으로 실행할 수 있습니다. 기존의 수동 RMSNorm 체인은 pow, mean, rsqrt, multiply 등 여러 개의 CUDA 커널 호출로 이루어져 있었을 가능성이 높습니다. 각 커널 호출에는 오버헤드가 발생하며, GPU의 연산 유닛이 데이터를 처리하는 동안 다른 연산으로 전환되지 못하고 대기하는 경우가 많습니다. Triton 커널은 이러한 연산들을 하나로 묶어(fused) GPU에서 더 적은 오버헤드로, 더 효율적으로 처리할 수 있게 합니다.
  2. 메모리 접근 최적화: Triton 커널은 데이터의 로딩 및 저장 방식을 최적화하여 메모리 대역폭 사용을 효율화할 수 있습니다. 특히 norm_inferrmsnorm_onepass_triton과 같은 커널들은 tiled execution, shared memory 활용 등을 통해 메모리 접근 패턴을 개선합니다.
  3. 조건부 최적화: 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

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글