[sglang] ROCm 환경에서 SGLang HiCache JIT 전송 커널 최적화 및 유연성 개선
PR 링크: sgl-project/sglang#37152 상태: Merged | 변경: +167 / -37
들어가며
SGLang의 HiCache는 KV 캐시 전송을 가속화하는 핵심 컴포넌트입니다. 기존 구현은 JIT 전송 커널에서 128바이트 단위의 고정된 타일링 전략을 사용했습니다. 하지만 최신 모델 아키텍처(예: MLA)의 fp8 데이터 행 크기인 576바이트와 같이 128로 나누어떨어지지 않는 경우, JIT 경로를 타지 못하고 성능이 저하되는 문제가 있었습니다. 본 PR은 ROCm 환경에서 이러한 제약을 해결하고, 더 넓은 범위의 데이터 크기를 효율적으로 처리할 수 있도록 JIT 커널을 개선했습니다.
코드 분석
1. hicache.cuh: 복사 단위의 유연성 확보
기존에는 128바이트 단위로만 처리가 가능했으나, pick_group_bytes 함수를 도입하여 하드웨어 지원 패키지 크기(4, 8, 16바이트)에 맞춰 128/64/32/16바이트 단위로 복사 라운드를 동적으로 선택하도록 변경되었습니다.
// Before: 128바이트 고정
static_assert(kBytes % 128 == 0, "kBytes must be multiple of 128 bytes");
// After: 동적 선택
inline constexpr uint32_t pick_group_bytes(int64_t bytes, uint32_t lanes_per_worker) {
#ifdef USE_ROCM
return group_fits(bytes, lanes_per_worker, 128) ? 128u
: group_fits(bytes, lanes_per_worker, 64) ? 64u
: group_fits(bytes, lanes_per_worker, 32) ? 32u
: group_fits(bytes, lanes_per_worker, 16) ? 16u
: 0u;
#else
return group_fits(bytes, lanes_per_worker, 128) ? 128u : 0u;
#endif
}
2. MHATokenToKOnlyPoolHost JIT 경로 활성화
기존에는 can_use_jit이 CUDA 전용으로 제한되어 있었으나, 커널 로직 자체가 특정 하드웨어에 종속적이지 않음을 확인하여 ROCm에서도 JIT 경로를 사용할 수 있도록 게이트를 열었습니다. 이를 통해 호스트 메모리 풀 활용 효율이 크게 개선되었습니다.
왜 이게 좋은가
- 범용성 확장: MLA fp8 행 크기인 576바이트와 같이 128의 배수가 아닌 데이터도 이제 JIT 최적화 경로를 통과할 수 있습니다. 이는 특정 모델 아키텍처에 대한 성능 병목을 제거합니다.
- ROCm 최적화: ROCm 환경에서의 블록 쿼터(Block Quota)를 16으로 조정하여 대역폭 활용도를 높였습니다. 이는 CUDA의 2와 비교하여 ROCm 아키텍처의 특성에 더 적합한 튜닝입니다.
- 안전성 강화:
static_assert를 통해 컴파일 타임에 하드웨어 제약 조건을 검증함으로써, 런타임 오류를 방지하고 유지보수성을 높였습니다.
일반적 교훈
- 하드웨어 추상화: 특정 하드웨어(CUDA)에 종속적인 가정을 코드에서 제거하고, 하드웨어의 패키지 크기(4/8/16B)와 같은 근본적인 제약 조건 기반으로 로직을 재설계하는 것이 이식성에 유리합니다.
- JIT 커널의 유연성: 고정된 타일링 크기보다는 데이터 크기에 따른 동적 타일링 선택이 다양한 모델 아키텍처를 지원하는 데 필수적입니다.
리뷰어 피드백 반영
리뷰어들은 ROCm 환경에서의 wave64와 같은 하드웨어 특성과 소프트웨어 스레드 그룹 간의 차이를 명확히 할 것을 요구했습니다. 이에 따라 kCopyGroupThreads를 명시적으로 정의하고, 하드웨어 워프 크기와 소프트웨어 복사 그룹 간의 관계를 문서화하여 코드의 의도를 명확히 했습니다.
참고 자료
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [sglang] [AMD ROCm] GLM-5.x Prefill 성능을 66% 끌어올린 Top-K 커널 최적화 분석
- [sglang] AMD MI355X에서 GLM-5.2 성능 극대화하기: 왜 다시 HIP Top-K인가?
- [sglang] SGLang aiter 백엔드의 Sliding Window Attention(SWA) 최적화 및 안정성 개선
- [sglang] ROCm DSA Indexer Top-K 최적화: 정확성과 성능을 동시에 잡다
- [sglang] ROCm 환경에서 BF16 All-Reduce의 수치 안정성 확보하기: QuickReduce의 FP16 Saturation 이슈 해결
PR Analysis 의 다른글
- 이전글 [vllm] NVIDIA RTX PRO 6000 GPU에서 VLLM의 행렬 곱셈 성능 최적화: sm120 아키텍처 지원 추가
- 현재글 : [sglang] ROCm 환경에서 SGLang HiCache JIT 전송 커널 최적화 및 유연성 개선
- 다음글 [sglang] H200 GPU에서 GLM-5.2 MoE를 위한 W4A8 GEMM 커널 최적화 분석
댓글