본문으로 건너뛰기

[sglang] SGLang LongCat-Flash Router GEMM 최적화: HPC-Ops bf16xfp32 커널로 H200에서 최대 4.31배 성능 향상

PR 링크: sgl-project/sglang#30247 상태: Merged | 변경: +336 / -7

들어가며

SGLang은 LLM(Large Language Model) 추론을 위한 고성능 프레임워크로, 다양한 최적화 기법을 통해 모델의 처리량을 극대화합니다. 이번 글에서는 SGLang의 LongCat-Flash 모델에서 중요한 구성 요소인 라우터(router)의 GEMM(General Matrix Multiply) 연산 최적화에 대해 다룹니다. LongCat-Flash 라우터는 bf16 타입의 활성화(activations)와 fp32 타입의 가중치(router weight)를 사용하여 연산을 수행합니다. 기존 방식은 fp32 GEMM을 사용하여 성능 병목이 발생할 수 있었으며, 이는 특히 NVIDIA Hopper 아키텍처의 Tensor Core가 bf16 연산에 최적화되어 있다는 점을 고려할 때 개선의 여지가 있었습니다.

이 PR의 핵심 목표는 fp32 가중치의 정확도를 유지하면서도 bf16 Tensor Core의 이점을 최대한 활용하여 라우터 GEMM의 성능을 획기적으로 향상시키는 것입니다. 이를 위해 Tencent의 HPC-Ops 라이브러리에서 제공하는 gemm_bf16xfp32 커널을 도입하여, fp32 가중치를 두 개의 bf16 절반으로 분해하고 이를 퓨전(fused)된 bf16 GEMM으로 처리하는 방식을 채택했습니다.

코드 분석

이번 최적화는 주로 sglang/jit_kernel/dsv4/gemm.py 파일에서 새로운 GEMM 커널을 통합하고, sglang/srt/models/longcat_flash.py 파일에서 이 커널을 LongCat-Flash 라우터에 적용하는 방식으로 이루어졌습니다.

sglang/jit_kernel/dsv4/gemm.py 변경사항

이 파일은 SGLang의 JIT(Just-In-Time) 컴파일된 커널들을 관리하는 곳입니다. HPC-Ops 커널을 활용하기 위한 여러 헬퍼 함수와 메인 linear_bf16_fp32 함수의 로직이 추가되었습니다.

  1. HPC-Ops 커널 가용성 및 조건 확인: _hpc_gemm_bf16xfp32_available 함수는 HPC-Ops 라이브러리 설치 여부와 현재 GPU가 Hopper 아키텍처(sm90a)인지 확인합니다. _can_use_hpc_gemm_bf16xfp32 함수는 입력 텐서의 차원, 데이터 타입, 연속성, min_m (토큰 수) 등의 추가적인 조건을 검사하여 HPC-Ops 커널 사용 가능 여부를 판단합니다.

  2. FP32 가중치 분해 및 캐싱: _get_bf16xfp32_weight_split 함수는 fp32 가중치를 HPC-Ops 커널이 요구하는 두 개의 bf16 절반(w_high, w_low)으로 분해합니다. 이 분해된 가중치와 split-K 플래그 워크스페이스는 가중치 텐서에 캐싱되어 재사용 효율을 높입니다. w_high = w.bf16 이고 w_low = ((w - w_high.float()) / _HPC_GEMM_WEIGHT_SCALE).bf16 형태로 분해됩니다.

  3. HPC-Ops 커널 호출: _linear_bf16_fp32_hpc 함수는 위에서 정의된 조건들을 만족할 경우 hpc.gemm_bf16xfp32 커널을 호출합니다. 이 커널은 분해된 bf16 가중치들을 사용하여 연산을 수행하고 fp32 출력을 반환합니다.

  4. linear_bf16_fp32 함수 로직 변경: 기존 linear_bf16_fp32 함수는 이제 HPC-Ops 커널 경로를 우선적으로 시도하고, 조건이 맞지 않거나 HPC-Ops가 설치되지 않은 경우 기존 cublas 기반의 fp32 GEMM 경로(_linear_bf16_fp32_cublas)로 폴백(fallback)하도록 변경되었습니다.

Before:

def linear_bf16_fp32(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
    if _use_aiter:
        return tgemm.mm(x, y, otype=x.dtype).float()
    elif _linear_bf16_fp32_algo == "deep_gemm":
        z = torch.empty(x.size(0), y.size(0), dtype=torch.float32, device=x.device)
        deep_gemm_wrapper.gemm_nt_bf16bf16f32(x, y, z)
        return z
    else:
        return torch.mm(x, y.t(), out_dtype=torch.float32)

After:

def linear_bf16_fp32(
    x: torch.Tensor,
    y: torch.Tensor,
    *,
    hpc_kernel_min_m: Optional[int] = None,
) -> torch.Tensor:
    if _use_aiter and y.dtype == torch.bfloat16:
        return tgemm.mm(x, y, otype=x.dtype).float()
    elif hpc_kernel_min_m is not None:
        output = _linear_bf16_fp32_hpc(x, y, min_m=hpc_kernel_min_m)
        if output is not None:
            return output
        return _linear_bf16_fp32_cublas(x, y)
    elif _linear_bf16_fp32_algo == "hpc":
        output = _linear_bf16_fp32_hpc(x, y)
        if output is not None:
            return output
        return _linear_bf16_fp32_cublas(x, y)
    elif _linear_bf16_fp32_algo == "deep_gemm" and y.dtype == torch.bfloat16:
        from sglang.srt.layers import deep_gemm_wrapper

        z = torch.empty(x.size(0), y.size(0), dtype=torch.float32, device=x.device)
        deep_gemm_wrapper.gemm_nt_bf16bf16f32(x, y, z)
        return z
    else:
        return _linear_bf16_fp32_cublas(x, y)

sglang/srt/models/longcat_flash.py 변경사항

이 파일은 LongCat-Flash 모델의 라우터 구현을 포함합니다. 라우터의 forward 메서드에서 새로운 linear_bf16_fp32 함수를 호출하도록 변경되었습니다.

  1. _LONGCAT_FLASH_ROUTER_HPC_GEMM_MIN_M 정의: 벤치마크를 통해 LongCat-Flash Chat/Lite 모델의 라우터 형태(hidden_size, n_routed_experts)별로 HPC-Ops GEMM이 cublas보다 성능이 우위를 보이는 최소 m (토큰 수) 값을 정의했습니다. 이는 불필요한 오버헤드를 피하고 최적의 성능을 보장하기 위한 가드(guard) 역할을 합니다.

  2. LongcatFlashRouter.forward 수정: forward 메서드 내에서 hpc_kernel_min_m이 설정되어 있고, 라우터 파라미터가 torch.float32이며 바이어스(bias)가 없는 경우에만 최적화된 linear_bf16_fp32 경로를 사용하도록 조건부 로직이 추가되었습니다. 이 외의 경우에는 기존의 classifier (일반적인 torch.nn.Linear 또는 유사한 구현)를 통해 연산을 수행합니다.

Before:

    def forward(self, hidden_states):
        logits, _ = self.classifier(hidden_states.to(self.rounter_params_dtype))
        return logits

After:

    def forward(self, hidden_states):
        if (
            self.hpc_kernel_min_m is not None
            and self.rounter_params_dtype == torch.float32
            and self.classifier.bias is None
        ):
            return linear_bf16_fp32(
                hidden_states,
                self.classifier.weight,
                hpc_kernel_min_m=self.hpc_kernel_min_m,
            )
        logits, _ = self.classifier(hidden_states.to(self.rounter_params_dtype))
        return logits

테스트 파일 추가

새로운 최적화 경로의 정확성과 동작을 검증하기 위해 두 개의 테스트 파일이 추가되었습니다:

  • test/registered/gemm/test_linear_bf16_fp32_hpc.py: HPC-Ops 경로가 fp32 레퍼런스와 수치적으로 일치하는지, min_m 디스패치 로직이 올바르게 작동하는지, 그리고 가중치 분할 캐시가 재사용되는지 등을 검증합니다.
  • test/registered/unit/models/test_longcat_flash_router_hpc_gemm.py: LongCat-Flash 라우터가 HPC-Ops bf16xfp32 커널로 올바르게 디스패치되는지 단위 테스트합니다.

왜 이게 좋은가

이번 최적화는 SGLang의 LongCat-Flash 모델 성능에 여러 가지 긍정적인 영향을 미칩니다.

성능 향상

NVIDIA H200 GPU에서의 벤치마크 결과는 이 최적화의 효과를 명확히 보여줍니다. 라우터 GEMM 단일 연산에서 HPC-Ops 커널은 기존 fp32 GEMM 경로 대비 최대 4.31배의 속도 향상을 달성했습니다.

shape m HPC (us) fp32 mm (us) speedup
k=6144, n=768 64 15.5 35.6 2.31x
k=6144, n=768 512 38.4 120.4 3.14x
k=6144, n=768 8192 430.0 1661.1 3.86x
k=3072, n=384 64 14.8 23.1 1.56x
k=3072, n=384 512 20.7 44.1 2.13x
k=3072, n=384 8192 113.1 487.8 4.31x

이러한 라우터 GEMM의 최적화는 전체 모델의 엔드-투-엔드(end-to-end) 성능에도 기여합니다. 특히, 토큰 수가 많은 Prefill 단계에서 +2.8%에서 +5.4%의 입력 처리량(input tok/s) 향상을 가져왔습니다. 디코딩(decode) 단계에서는 min_m 가드에 의해 기존 cublas 경로를 사용하므로 성능 변화는 미미합니다. 이는 최적화가 특정 병목 지점에 효과적으로 작용했음을 보여줍니다.

정확도 유지

가장 중요한 점은 이러한 성능 향상이 fp32 가중치의 수치적 정확도를 유지하면서 달성되었다는 것입니다. HPC-Ops 커널은 fp32 가중치를 bf16으로 분해하여 Tensor Core를 활용하되, 연산 과정에서 fp32 레벨의 정확도를 보존하도록 설계되었습니다. 모델 수준의 검증(GSM8K 정확도)에서도 기존 main 브랜치와 비교하여 정확도 변화가 거의 없거나 미미한 향상을 보였습니다.

기술적 교훈 및 리뷰 반영

  1. 하드웨어 특화 커널의 중요성: NVIDIA Hopper 아키텍처와 같은 최신 GPU는 bf16 Tensor Core를 통해 높은 성능을 제공합니다. HPC-Ops와 같은 하드웨어 특화 커널을 활용하는 것은 특정 연산의 성능을 극대화하는 데 필수적입니다.

  2. 정확도와 성능의 균형: bf16 Tensor Core를 활용하면서도 fp32 가중치의 정확도를 유지하는 기법은 딥러닝 모델의 실용적인 최적화에서 중요한 과제입니다. 이 PR은 가중치를 분해하여 이 두 가지 목표를 동시에 달성하는 좋은 예시입니다.

  3. 병목 지점 집중 최적화: 라우터 GEMM과 같이 전체 워크로드에서 작은 부분일지라도, 그 부분이 병목이라면 집중적인 최적화는 전체 시스템 성능에 유의미한 영향을 미칠 수 있습니다.

  4. 리뷰 피드백을 통한 견고성 확보: PR 리뷰 과정에서 VAthree님은 온라인 가중치 업데이트 시 캐시된 가중치가 무효화되지 않아 발생할 수 있는 'silent correctness failure' 문제를 지적했습니다. 이는 param.data.copy_()와 같은 인플레이스(in-place) 업데이트가 _version을 변경하지 않아 캐시 키가 동일하게 유지되는 문제였습니다.

    초기에는 이 문제를 해결하기 위해 캐시를 가중치 업데이트 라이프사이클에 통합하는 방안이 논의되었으나, 최종적으로는 이 HPC-Ops 최적화 경로가 고정 가중치 서빙(fixed-weight serving)에 초점을 맞추도록 결정되었습니다. 이에 따라, 온라인 가중치 업데이트 API가 활성화된 split cache를 감지하면 명확한 오류를 발생시켜 실패(fail fast)하도록 변경되었습니다. 또한, CUDA Graph 관련 문제를 해결하기 위해 캐시 키에서 _version을 제거하여 split buffer가 한 번 할당되면 주소가 변경되지 않도록 하여 CUDA Graph의 안정성을 확보했습니다. 이러한 결정은 시스템의 예측 가능성과 안정성을 높이는 중요한 설계 선택입니다.

결론

이번 SGLang LongCat-Flash 라우터 GEMM 최적화는 HPC-Ops의 bf16xfp32 커널을 성공적으로 통합하여 H200 GPU에서 라우터 연산의 성능을 최대 4배 이상 향상시켰습니다. 이는 모델의 Prefill 처리량을 최대 5.4%까지 끌어올리며, fp32 가중치의 정확도를 유지하는 동시에 bf16 Tensor Core의 이점을 극대화했습니다. 또한, 리뷰 과정을 통해 온라인 가중치 업데이트 시 발생할 수 있는 잠재적인 문제를 명확히 인지하고, 고정 가중치 서빙이라는 이 최적화의 핵심 목표에 맞춰 시스템의 견고성을 확보한 점은 기술 블로그 작성자로서 매우 인상 깊은 부분입니다. 앞으로 SGLang의 GEMM 인터페이스 통합과 같은 더 큰 규모의 리팩토링도 기대됩니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글