본문으로 건너뛰기

[sglang] ROCm DSA Indexer Top-K 최적화: 정확성과 성능을 동시에 잡다

PR 링크: sgl-project/sglang#37591 상태: Merged | 변경: +1672 / -2

들어가며

대규모 언어 모델(LLM)의 성능은 단순히 모델의 크기뿐만 아니라, 이를 효율적으로 구동하는 인프라 소프트웨어의 최적화에도 크게 좌우됩니다. 특히, sglang과 같은 프레임워크에서 DSA(DeepSeek-V4 Architecture) 인덱서의 top-k 연산은 다음 토큰 예측 과정에서 핵심적인 역할을 합니다. 이 연산은 수많은 후보 토큰 중에서 가장 확률이 높은 k개의 토큰을 정확하고 빠르게 찾아내는 것이 중요합니다.

이번에 분석할 PR("[ROCm] Make DSA indexer top-k exact with cooperative selection")은 sglang의 ROCm(AMD GPU) 환경에서 DSA 인덱서의 top-k 구현을 대폭 개선하여, 기존의 정확성 문제를 해결하고 동시에 성능을 크게 향상시킨 사례입니다. 이 PR은 기존의 hipify된 코드의 한계를 극복하고, 네이티브 HIP 커널을 통해 fp32 정밀도와 오버플로우 재스캔 메커니즘을 도입하여 top-k 연산의 정확성을 torch.topk 수준으로 끌어올리면서도 탁월한 성능을 달성했습니다.

문제점: 기존 ROCm Top-K 구현의 한계

이전 ROCm topk 구현은 topk.cu 파일을 hipify 도구를 통해 변환한 topk.hip을 사용했습니다. 이 방식은 몇 가지 심각한 문제를 야기했습니다.

  1. 낮은 정밀도 (fp16 coarse key): fp16 (half-precision floating-point)을 사용하여 coarse histogram을 생성했는데, 이는 넓은 범위의 실제 인덱서 행 값을 하나의 버킷으로 뭉개버리는(collapse) 현상을 초래했습니다. 특히 GLM-5.2 로짓과 같이 값이 밀집된 분포에서는 fp16 키가 4-126개의 빈만 채우고, 가장 큰 빈이 행의 6%에서 88%를 차지하는 경우가 발생했습니다. 이로 인해 top-k 경계가 이 지배적인 빈 안에 들어갈 경우, 정확한 top-k 값을 선별하기 어려웠습니다.
  2. 후보 버퍼 오버플로우 및 침묵하는 오류: 후보 버퍼가 오버플로우될 경우, 유효한 top-k 후보들이 아무런 경고 없이 조용히 누락되는 문제가 있었습니다. 이는 GLM-5.2 로짓에서 selected-value recall이 0.53까지 떨어지는 결과를 초래했습니다. 즉, top-k 연산이 정확한 결과를 반환하지 못하고 있었습니다.

이러한 문제점들은 모델의 예측 품질에 직접적인 악영향을 미치며, 특히 정확성이 중요한 LLM 추론 과정에서는 치명적일 수 있습니다.

해결책: 네이티브 HIP 커널과 협력적 Top-K

이 PR은 위에서 언급된 문제들을 해결하기 위해 다음과 같은 핵심적인 개선 사항들을 도입했습니다.

  1. 네이티브 HIP 커널 구현: topk.cuhipify하는 대신, ROCm 아키텍처에 최적화된 네이티브 HIP 커널인 topk.hip을 직접 구현하고 관리합니다. 이는 hipify 도구의 자동 변환이 가질 수 있는 한계를 극복하고, ROCm GPU의 특성을 최대한 활용할 수 있게 합니다.
  2. fp32 정밀도 coarse histogram: 기존의 fp16 대신 fp32 (single-precision floating-point)를 사용하여 coarse histogram을 생성함으로써, 값의 정밀도를 높여 top-k 연산의 정확성을 보장합니다.
  3. 정확한 Radix Tie Refinement: 동점(tie) 상황에서 fp32 키를 사용하여 정확하게 순위를 매기는 radix tie refinement 기법을 도입하여, top-k 경계에 있는 값들의 순서가 정확하게 유지되도록 합니다.
  4. 오버플로우 재스캔 (Overflow Rescan): 후보 버퍼가 오버플로우될 경우, 행 전체를 다시 스캔하여 누락된 유효한 top-k 후보가 없도록 보장합니다. 이로써 torch.topk와 동일한 1.0의 selected-value recall을 달성합니다.
  5. 협력적 선택 (Cooperative Selection): 여러 스레드 블록이 협력하여 top-k 연산을 수행하는 방식을 사용합니다. 이는 특히 긴 행(long rows) 처리 시 GPU 활용도를 높이는 데 기여합니다.
  6. 로우 분할 (Row Splitting): 낮은 배치 크기(batch size)로 인해 GPU 활용도가 떨어질 때, 긴 행을 여러 블록에 걸쳐 분할하여 처리합니다. 이는 LDS(Local Data Share) histogram의 경쟁을 줄이고, 행에 더 많은 메모리 병렬성을 제공하여 GPU를 효율적으로 활용하게 합니다.

launch 함수는 이 로우 분할 로직에 따라 단일 블록 커널(coop_topk_kernel)을 실행하거나, 여러 단계의 멀티 블록 커널(coop_mb_hist0, coop_mb_hist1, coop_mb_scatter, coop_mb_refine)을 순차적으로 실행합니다. 이 다단계 접근 방식은 대규모 데이터셋에서 효율적인 top-k 연산을 가능하게 합니다.

코드 분석

3rdparty/amd/wheel/sgl-kernel/rocm_hipify.py

이 파일의 변경은 topk 구현 방식의 근본적인 변화를 보여줍니다. 기존에는 topk.cuhipify하여 ROCm 환경에서 사용했지만, 이제는 이 목록에서 topk.cu를 제외하고 네이티브 topk.hip을 직접 관리합니다.

Before:

-    "csrc/elementwise/topk.cu",

After:

+    # topk.hip is maintained as native HIP instead of being generated from topk.cu.
+    "csrc/elementwise/topk.cu",

이 변경은 sglang 팀이 topk 연산의 정확성과 성능을 위해 hipify된 코드의 한계를 인정하고, ROCm 플랫폼에 특화된 최적화를 직접 수행하기로 결정했음을 명확히 보여줍니다. 주석(topk.hip is maintained as native HIP instead of being generated from topk.cu.)은 이러한 전략적 변화를 설명합니다.

python/sglang/kernels/aot/csrc/elementwise/topk.hip

이 파일은 새로 추가된 네이티브 HIP 커널 구현입니다. 핵심적인 부분들을 살펴보겠습니다.

상수 정의

constexpr uint32_t kTopK = 2048;
constexpr uint32_t kHistBits = 12;
constexpr uint32_t kTieCap = 4096;
constexpr uint32_t kBlock = 1024;

kTopKDeepSeek-V3.2GLM-5.2 인덱서에서 사용하는 top-k 값인 2048로 고정되어 있습니다. kHistBits는 히스토그램의 빈 개수를 결정하고, kTieCap은 동점 후보를 저장할 버퍼의 크기를, kBlock은 스레드 블록당 스레드 수를 정의합니다. 이 상수들은 커널의 메모리 사용량과 병렬 처리 전략에 직접적인 영향을 미칩니다.

row_split 함수

이 함수는 주어진 batch 크기와 row_len_hint (행 길이 힌트)를 기반으로 행을 여러 블록으로 분할할지 여부와 분할할 블록 수를 결정합니다. 이는 GPU의 multiProcessorCount를 고려하여 GPU 활용도를 최적화합니다.

int row_split(int batch, int64_t row_len_hint) {
  // ... 환경 변수 SGL_DSA_TOPK_ROW_SPLIT 처리 ...

  // ... hipGetDeviceProperties를 통해 GPU의 multiProcessorCount 캐싱 ...

  constexpr int64_t kMinRowLen = 65536;
  constexpr int kMinSplit = 4;
  constexpr int kMaxSplit = 32;

  if (row_len_hint < kMinRowLen) {
    return 0; // 짧은 행은 분할하지 않음
  }
  const int g = std::min(cached_cu / std::max(batch, 1), kMaxSplit);
  return g < kMinSplit ? 0 : g; // 충분히 많은 CU가 없거나 배치 크기가 크면 분할하지 않음
}

kMinRowLen (65536)보다 짧은 행은 분할하지 않고, batch 크기와 GPU의 multiProcessorCount를 고려하여 kMinSplit (4)에서 kMaxSplit (32) 사이의 값으로 g (블록 수)를 결정합니다. 이 로직은 낮은 배치 크기에서 GPU가 유휴 상태가 되는 것을 방지하고, 긴 행에서 LDS histogram 경쟁을 줄이며, 메모리 병렬성을 높이는 데 기여합니다. PR 설명에 따르면, GLM-5.2 디코드 형태에서 로우 분할은 기존 커널 대비 2.1-2.8배의 추가 속도 향상을 제공합니다.

launch 함수

launch 함수는 row_split의 결과에 따라 단일 블록 커널 또는 멀티 블록 커널을 실행합니다.

void launch(
    const float* input,
    int32_t* out_idx,
    const int32_t* row_starts,
    const int32_t* lengths,
    const OutMap& map,
    int batch,
    int64_t stride,
    at::Device device,
    hipStream_t stream) {
  const int g = row_split(batch, stride);

  if (g == 0) {
    // 단일 블록 커널 실행
    coop_topk_kernel<kTopK, kHistBits, kTieCap, kBlock><<<batch, kBlock, 0, stream>>>(p);
    return;
  }

  // 멀티 블록 (로우 분할) 커널 실행
  const size_t ws_bytes = Ws::bytes(static_cast<size_t>(batch));
  at::Tensor ws = at::empty({static_cast<int64_t>(ws_bytes)}, at::TensorOptions().dtype(at::kByte).device(device));

  // 워크스페이스 초기화
  C10_HIP_CHECK(hipMemsetAsync(p.ws, 0, Ws::zero_bytes(static_cast<size_t>(batch)), stream));

  const dim3 grid(static_cast<unsigned>(batch), static_cast<unsigned>(g), 1);
  coop_mb_hist0<kTopK, kHistBits, kTieCap, kBlock><<<grid, kBlock, 0, stream>>>(p, static_cast<uint32_t>(g));
  coop_mb_hist1<kTopK, kHistBits, kTieCap, kBlock><<<grid, kBlock, 0, stream>>>(p, static_cast<uint32_t>(g));
  coop_mb_scatter<kTopK, kHistBits, kTieCap, kBlock><<<grid, kBlock, 0, stream>>>(p, static_cast<uint32_t>(g));
  coop_mb_refine<kTopK, kHistBits, kTieCap, kBlock><<<batch, kBlock, 0, stream>>>(p);
}

g == 0일 때는 각 행이 하나의 블록으로 처리되는 coop_topk_kernel이 실행됩니다. g > 0일 때는 워크스페이스를 할당하고 초기화한 후, coop_mb_hist0, coop_mb_hist1, coop_mb_scatter, coop_mb_refine의 네 가지 멀티 블록 커널이 순차적으로 실행됩니다. 이 다단계 접근 방식은 히스토그램 생성, 후보 수집, 최종 정제 과정을 분리하여 복잡한 top-k 로직을 효율적으로 처리합니다.

성능 및 정확성 개선

이 PR은 MI355X (gfx950) GPU에서 광범위한 성능 테스트를 거쳤습니다. 결과는 new native 커널의 압도적인 우위를 보여줍니다.

old generated vs new native

distribution batch old generated (μs) new native (μs) speedup
tiny 1 95.2 31.6 3.02x
tiny 256 121.9 62.7 1.94x
diffuse 1 54.3 31.6 1.72x
diffuse 256 103.4 63.2 1.64x
clustered 1 160.9 102.2 1.57x
clustered 256 186.9 174.2 1.07x

new native 커널은 old generated 커널 대비 최소 1.07배에서 최대 3.18배까지의 속도 향상을 보여줍니다. 특히 tiny 분포에서 큰 폭의 개선이 두드러집니다. 이는 fp32 정밀도와 최적화된 커널 로직 덕분입니다.

native topk.hip vs topk_v2

sglang에는 topk_v2라는 또 다른 top-k 커널이 존재하며, SGLANG_OPT_USE_TOPK_V2=1 환경 변수를 통해 선택적으로 사용될 수 있습니다. 이 PR의 native topk.hip 커널은 topk_v2와도 비교되었습니다.

distribution batch native topk.hip (μs) topk_v2 (μs)
tiny 1 31.5 73.1
tiny 256 63.2 76.1
diffuse 1 31.9 34.2
diffuse 256 62.9 46.5
clustered 1 102.2 109.2
clustered 256 174.5 111.0

성능 면에서는 tiny 분포에서 native topk.hip이 1.20-2.43배 더 빠릅니다. diffuse 분포에서는 낮은 배치에서 native가 빠르지만, 높은 배치에서는 topk_v2가 1.04-1.49배 더 빠릅니다. clustered 분포에서도 유사한 경향을 보입니다.

그러나, 가장 중요한 차이점은 정확성입니다. topk_v2fp16 coarse histogram을 사용하고 오버플로우 재스캔이 없기 때문에, tinyclustered 분포에서 심각한 정확성 문제를 보입니다.

distribution native topk.hip topk_v2
tiny 1.0000 0.0420
diffuse 1.0000 1.0000
clustered 1.0000 0.0220

topk_v2tiny 분포에서 0.0420, clustered 분포에서 0.0220이라는 매우 낮은 selected-value recall을 기록했습니다. 이는 topk_v2가 이러한 분포에서 정확한 top-k 결과를 반환하지 못한다는 것을 의미합니다. 반면, native topk.hip은 모든 분포에서 1.0000의 완벽한 recall을 유지합니다.

리뷰어 jiejingzhangamd의 추가 비교에 따르면, 실제 GLM-5.2 로짓(20행, 134868 길이)에서 topk_v2는 160행 중 1행에서 0.957 recall을 기록했는데, 이는 임계값 빈에 3296개의 후보가 있어 2048개 버퍼를 초과했기 때문입니다. 이 오버플로우 상황에서 native topk.hip은 49.4 μs로 정확하게 실행된 반면, topk_v2는 64.0 μs로 느리면서도 부정확했습니다. 이는 native topk.hip이 오버플로우 상황에서 정확성을 유지하면서도 성능 우위를 가짐을 명확히 보여줍니다.

왜 이 최적화가 좋은가?

이 PR의 최적화는 여러 면에서 sglang의 ROCm 지원을 한 단계 끌어올리는 중요한 개선입니다.

  1. 정확성 보장: fp32 정밀도와 오버플로우 재스캔 메커니즘을 통해 top-k 연산의 정확성을 torch.topk와 동일한 수준으로 끌어올렸습니다. 이는 LLM의 예측 품질에 직접적인 영향을 미치며, 특히 민감한 애플리케이션에서 신뢰할 수 있는 결과를 제공하는 데 필수적입니다. 이전의 fp16 기반 구현이 야기했던 selected-value recall 저하 문제를 근본적으로 해결했습니다.
  2. 성능 향상: 다양한 데이터 분포와 배치 크기에서 기존 hipify된 커널 대비 최대 3배 이상의 속도 향상을 달성했습니다. 특히 row_split 로직은 낮은 배치 크기에서도 GPU의 병렬 처리 능력을 최대한 활용하여 GLM-5.2 디코드 형태에서 2.1-2.8배의 추가 속도 향상을 가져왔습니다. 이는 전체 LLM 추론 파이프라인의 처리량을 높이는 데 크게 기여합니다.
  3. GPU 활용 최적화: row_split 함수는 GPU의 multiProcessorCount를 동적으로 고려하여, 행의 길이에 따라 최적의 블록 분할 전략을 적용합니다. 이는 GPU 자원의 낭비를 줄이고, 특히 긴 시퀀스 처리 시 효율성을 극대화합니다.
  4. 네이티브 최적화의 중요성 입증: hipify와 같은 자동 변환 도구는 개발 편의성을 제공하지만, 특정 아키텍처의 미세한 성능 특성을 완전히 반영하기 어렵습니다. 이 PR은 ROCm에 특화된 네이티브 HIP 커널을 직접 구현함으로써, 자동 변환된 코드의 한계를 극복하고 ROCm GPU의 잠재력을 최대한 끌어낼 수 있음을 보여줍니다.
  5. 견고한 엣지 케이스 처리: clustered 로짓과 같이 값이 밀집되어 top-k 경계가 하나의 빈에 몰리거나, 후보 버퍼가 오버플로우되는 엣지 케이스에 대해 정확성을 유지하면서도 효율적으로 처리하는 견고한 로직을 제공합니다.

결론 및 일반적 교훈

이 PR은 sglang의 ROCm top-k 연산을 단순한 성능 개선을 넘어, 정확성이라는 핵심적인 요구사항을 충족시키면서도 탁월한 성능을 달성한 모범적인 사례입니다. 여기서 얻을 수 있는 일반적인 교훈은 다음과 같습니다.

  • 정확성은 타협할 수 없는 가치: 특히 AI/ML 모델의 핵심 연산에서는 성능 최적화만큼이나 결과의 정확성이 중요합니다. 때로는 성능을 약간 희생하더라도 정확성을 보장하는 것이 장기적으로 더 큰 가치를 제공합니다.
  • 자동화 도구의 한계 인식: hipify와 같은 자동 변환 도구는 유용하지만, 최적의 성능과 정확성을 위해서는 플랫폼별 네이티브 최적화가 필수적일 수 있습니다. 특정 하드웨어 아키텍처의 특성을 깊이 이해하고 직접 코드를 작성하는 노력이 필요합니다.
  • 엣지 케이스에 대한 견고한 설계: 데이터 분포의 특이성이나 자원 제약(예: 버퍼 오버플로우)과 같은 엣지 케이스를 예측하고 이에 대한 견고한 처리 로직을 설계하는 것이 안정적인 시스템을 구축하는 데 중요합니다.
  • 협력적 컴퓨팅의 힘: 여러 스레드 블록이 협력하여 작업을 분담하는 cooperative selection과 같은 기법은 대규모 병렬 연산에서 GPU 활용도를 극대화하고 성능을 향상시키는 효과적인 방법입니다.

이러한 최적화는 sglang이 ROCm 플랫폼에서 더욱 강력하고 신뢰할 수 있는 LLM 추론 프레임워크로 자리매김하는 데 중요한 토대가 될 것입니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글