본문으로 건너뛰기

[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_llmcvt_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;

이 코드는 attentionmlp 브랜치를 각각 읽어와서 별도의 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%

교훈

  1. Memory Access Minimization: 중간 결과물을 메모리에 쓰지 않고 레지스터 수준에서 처리하는 'Kernel Fusion'은 대역폭 제한적인(bandwidth-bound) 연산에서 매우 효과적입니다.
  2. Bit-Exactness: 최적화 과정에서 결과값이 달라지지 않음을 보장하기 위해 SHA256 비교 및 pixel-exact 검증을 수행하여, 성능 향상과 정확성을 동시에 잡았습니다.
  3. Fallback Strategy: JIT 커널이 지원하지 않는 환경(비 SM103 디바이스 등)에서는 기존의 eager 방식을 사용하도록 설계하여 안정성을 확보했습니다.

리뷰어 피드백 반영

리뷰 과정에서 초기 이미지 생성 결과의 미세한 차이가 지적되었으나, 이는 본 PR의 문제가 아닌 다른 의존성(PR #37096) 때문임이 밝혀졌습니다. 작성자는 해당 의존성을 제거하고 독립적인 PR로 재구성하여, 최종적으로 원본과 비트 단위로 동일한(byte-exact) 결과를 보장하도록 개선했습니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글