[sglang] FLUX.2 모델 성능 최적화: Token Concatenation과 NVFP4 양자화의 커널 융합
PR 링크: sgl-project/sglang#37141 상태: Merged | 변경: +474 / -3
들어가며
최신 대규모 언어 모델 및 확산 모델(Diffusion Model)의 추론 성능을 극대화하기 위해서는 메모리 대역폭을 절약하고 커널 실행 횟수를 줄이는 것이 핵심입니다. 이번 SGLang의 PR은 FLUX.2 모델의 단일 스트림(single-stream) 블록에서 발생하는 torch.cat([attention_bf16, swiglu_bf16]) 연산과 이후 이어지는 NVFP4 양자화 과정을 하나의 SM103 전용 JIT CUDA 커널로 융합(fuse)하여 성능을 개선했습니다.
기존에는 BF16 데이터를 먼저 결합한 뒤 별도의 커널에서 양자화를 수행했으나, 이 과정에서 불필요한 메모리 읽기/쓰기가 발생했습니다. 이를 단일 커널로 통합함으로써 24,576-wide BF16 데이터의 중간 생성 과정을 제거하고, 양자화된 데이터를 직접 메모리에 기록하도록 최적화했습니다.
코드 분석
1. JIT 커널 구현 (flux2_token_cat_nvfp4.cuh)
핵심은 tensorrt_llm의 cvt_warp_fp16_to_fp4 프리미티브를 활용하여, 데이터를 읽어오는 즉시 양자화하여 저장하는 것입니다.
Before (개념적):
# 기존 방식: 두 번의 커널 실행
cat_data = torch.cat([attention, mlp], dim=-1)
quantized, scales = flashinfer.fp4_quantize(cat_data)
After (핵심 커널 로직):
// 융합된 커널: 읽기 -> 즉시 양자화 -> 쓰기
const uint64_t packed = tensorrt_llm::kernels::cvt_warp_fp16_to_fp4<__nv_bfloat16, kGroupSize, kGroupSize, false>(
quant_vec, global_scale, quant_scales + scale_offset);
static_cast<uint64_t*>(params.quantized)[int64_t(row) * kScaleColumns + group] = packed;
이 코드는 attention과 mlp 브랜치를 각각 읽어와서 별도의 cat 연산 없이 바로 NVFP4 포맷으로 변환합니다.
2. JIT 모듈 로드 및 검증 (flux2_token_cat_nvfp4_jit.py)
이 커널은 torch.compile이나 CUDA 그래프 캡처 등 특정 조건에서만 동작하도록 설계되었으며, 입력 데이터의 정렬(alignment)과 디바이스 capability(SM103)를 엄격히 검사합니다.
# 조건부 실행 로직
if (
torch.compiler.is_compiling()
or not _is_dense_bf16(attention, _ATTENTION_HIDDEN)
or torch.cuda.get_device_capability(attention.device) != (10, 3)
): return None
왜 이게 좋은가
성능 향상
벤치마크 결과, Denoise 단계당 약 2.413%의 속도 향상을 확인했습니다. 특히 커널 실행 횟수가 크게 줄어들었습니다.
| Metric | Parent | PR | Delta |
|---|---|---|---|
| Total kernel launches | 1089 | 1041 | -48 |
| Target cat + quant GPU time | 5464 us | 2261 us | -58.6% |
교훈
- Memory Access Minimization: 중간 결과물을 메모리에 쓰지 않고 레지스터 수준에서 처리하는 'Kernel Fusion'은 대역폭 제한적인(bandwidth-bound) 연산에서 매우 효과적입니다.
- Bit-Exactness: 최적화 과정에서 결과값이 달라지지 않음을 보장하기 위해
SHA256비교 및pixel-exact검증을 수행하여, 성능 향상과 정확성을 동시에 잡았습니다. - Fallback Strategy: JIT 커널이 지원하지 않는 환경(비 SM103 디바이스 등)에서는 기존의 eager 방식을 사용하도록 설계하여 안정성을 확보했습니다.
리뷰어 피드백 반영
리뷰 과정에서 초기 이미지 생성 결과의 미세한 차이가 지적되었으나, 이는 본 PR의 문제가 아닌 다른 의존성(PR #37096) 때문임이 밝혀졌습니다. 작성자는 해당 의존성을 제거하고 독립적인 PR로 재구성하여, 최종적으로 원본과 비트 단위로 동일한(byte-exact) 결과를 보장하도록 개선했습니다.
참고 자료
- https://pytorch.org/docs/stable/generated/torch.compile.html
- https://docs.nvidia.com/cuda/cuda-c-programming-guide/index.html#warp-level-matrix-operations
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
PR Analysis 의 다른글
- 이전글 [flashinfer] FlashInfer: Blackwell 아키텍처를 위한 결정론적 BGMV MoE 최적화
- 현재글 : [sglang] FLUX.2 모델 성능 최적화: Token Concatenation과 NVFP4 양자화의 커널 융합
- 다음글 [uv] macOS에서 uv 캐시 정리가 3.8배 빨라진 비결: getattrlistbulk를 활용한 일괄 메타데이터 조회
댓글