[sglang] ERNIE-Image 모델의 성능 향상: 비트-정확성(Bit-Exact)을 유지한 RMSNorm+Scale/Shift 융합 커널 도입
PR 링크: sgl-project/sglang#33854 상태: Merged | 변경: +524 / -4
들어가며
최근 LLM 및 Diffusion 모델 분야에서는 모델의 추론 속도를 높이기 위한 다양한 최적화 기법이 활발히 연구되고 있습니다. 특히, 모델의 정확도를 최대한 유지하면서 연산량을 줄이는 것은 매우 중요한 과제입니다. 이번 PR(#33854)은 sglang 레포지토리에서 ERNIE-Image 모델의 핵심 연산 중 하나인 RMSNorm과 Scale/Shift 연산을 융합하여 성능을 크게 향상시키는 동시에, 기존 연산과 완벽하게 동일한 결과(bit-exact)를 보장하는 새로운 Triton 커널을 도입했습니다.
기존에는 ERNIE-Image 모델의 adaLN 블록에서 RMSNorm, Scale, Shift 연산이 여러 개의 개별 CUDA 커널로 분리되어 실행되었습니다. 이는 GPU의 연산 능력을 충분히 활용하지 못하고 메모리 대역폭에 병목 현상을 일으키는 주요 원인이었습니다. 더 큰 문제는, 이러한 연산들을 하나의 커널로 융합하려는 시도들이 종종 수치적 부정확성(numerical inaccuracy)을 야기하여 모델의 최종 출력 품질을 저하시킨다는 점이었습니다. 이전 PR(#30170)에서 도입되었던 융합 커널들은 FP32 연산을 통해 norm(x)*(1+scale)+shift를 한 번에 계산하려 했으나, 이 과정에서 발생하는 반올림 오차로 인해 50-step 디노이징 과정에서 PSNR 18.83 dB라는 낮은 품질 저하를 보였습니다. 이는 모델의 원래 품질 기준인 PSNR 25 dB에 훨씬 못 미치는 수준이었습니다.
본 PR은 이러한 문제를 해결하기 위해, 기존 연산의 산술 연산 순서와 반올림 방식을 그대로 모방하는 'rounding-faithful' Triton 커널을 새롭게 구현했습니다. 이를 통해 여러 개의 개별 연산을 단일 Triton 커널로 통합하면서도, 기존 연산과 비트 단위까지 동일한 결과를 보장합니다. 이는 별도의 quality=high와 같은 품질 게이트 없이도 기본 경로에서 안전하게 사용할 수 있음을 의미합니다.
ERNIE-Image 모델에서 adaLN 블록은 이미지 생성 과정에서 1024x1024 해상도 기준으로 약 7,200번 (36 블록 x 2 사이트 x 2 CFG x 50 스텝) 호출됩니다. 각 호출 시 여러 개의 CUDA 커널(rmsnorm, 1+scale, mul, add 등)이 실행되며, 이는 주로 메모리 대역폭에 의해 제한됩니다. 본 PR에서 도입된 융합 커널은 이러한 병목 현상을 효과적으로 해소할 것으로 기대됩니다.
코드 분석
python/sglang/kernels/ops/diffusion/triton/rmsnorm_scale_shift_bitexact.py (신규 파일)
이 파일은 ERNIE-Image 모델의 adaLN 블록에 필요한 두 가지 핵심 융합 연산을 Triton 커널로 구현합니다:
fused_rmsnorm_scale_shift_bitexact:norm(x) * (1 + scale) + shift연산을 융합합니다.fused_scale_residual_rmsnorm_scale_shift_bitexact:res = residual + gate * update연산 후,norm(res) * (1 + scale) + shift를 융합합니다. 이 커널은 기존의residual_gate_add_cuda커널과 RMSNorm, Scale/Shift 연산을 하나로 통합합니다.
핵심은 'rounding-faithful'이라는 개념입니다. 이는 단순히 연산 결과를 근사하는 것이 아니라, 기존 CUDA 커널들의 산술 연산 순서, FMA(Fused Multiply-Add) 연산 사용 여부, 그리고 각 단계별 반올림(rounding) 방식까지 정확히 모방하는 것을 의미합니다. 이를 위해 다음과 같은 기법들이 사용되었습니다:
- 정확한 산술 연산 순서 재현: Triton 언어의
tl.inline_asm_elementwise를 사용하여 특정 PTX 명령어(mul.rn.f32,rsqrt.approx.f32)를 직접 호출하고, FMA 연산이 일어나지 않도록 각 곱셈 연산(x*x)을 별도로 처리합니다. 이는 MLIR의vector.reduction시맨틱스를 모방합니다. - 워프(Warp) 수준 환원(Reduction):
shfl.bfly명령어를 활용하여 워프 내의 합계를 계산하는 방식을 정확히 재현합니다. Triton의tl.reshape및tl.split을 사용하여 연산 순서를 유지합니다. - 반올림(Rounding) 처리: FP32 연산 결과를 BF16으로 변환할 때, 기존 CUDA 커널과 동일한 반올림 방식(round-to-nearest-even)을 사용하기 위해
_round_bf16_to_fp32헬퍼 함수를 사용합니다. 이는 각 연산(1+scale,y * that,prod + shift) 후의 BF16 경계를 정확히 모방합니다. - 잔차(Residual) 연결:
fused_scale_residual_rmsnorm_scale_shift_bitexact커널은 기존residual_gate_add_cuda의 동작 방식을 정확히 모방하여,round(gate * update)후round(residual + that)연산을 수행하고 그 결과를 RMSNorm에 입력합니다.
이 커널들은 register_custom_op 데코레이터를 통해 PyTorch의 torch.compile과 호환되도록 등록됩니다. 또한, 런타임 시 첫 번째 실제 입력에 대해 기존 연산과 torch.equal을 사용하여 비트-정확성을 검증하고, 만약 불일치가 발생하면 자동으로 기존의 Eager 모드로 전환하는 안전 장치를 포함합니다.
# 예시: RMSNorm 연산의 산술 연산 순서 재현
# 기존 CUDA 커널의 mul.rn.f32 동작을 모방하여 FMA 방지
@triton.jit
def _mul_rn_f32(x, y):
return tl.inline_asm_elementwise(
asm="mul.rn.f32 $0, $1, $2;",
constraints="=f,f,f",
args=[x, y],
dtype=tl.float32,
is_pure=True,
pack=1,
)
# 예시: BF16 변환 시 반올림 방식 모방
@triton.jit
def _round_bf16_to_fp32(value):
bits = value.to(tl.int32, bitcast=True)
rounding_bias = 0x7FFF + ((bits >> 16) & 1)
rounded_bits = (bits + rounding_bias) & -65536
return rounded_bits.to(tl.float32, bitcast=True)
ernie_image.py 수정
이 파일에서는 새로 구현된 Triton 커널을 ERNIE-Image 모델의 adaLN 레이어에 적용하는 래퍼 함수(_ernie_norm_scale_shift, _ernie_gated_norm_scale_shift)를 수정했습니다. 이 래퍼 함수들은 다음과 같은 역할을 수행합니다:
- 커널 등록 및 호출:
can_use_fused_rmsnorm_scale_shift또는can_use_fused_scale_residual_rmsnorm_scale_shift함수를 통해 융합 커널 사용 가능 여부를 확인하고, 가능하면 Triton 커널을 호출합니다. - 런타임 검증: 첫 번째 실행 시, 융합 커널의 결과와 기존 Eager 연산의 결과를
torch.equal로 비교하여 비트-정확성을 검증합니다. 만약 검증에 실패하거나 예외가 발생하면, 해당 fast path는 비활성화되고 이후에는 Eager 연산으로 대체됩니다. 이는 #33734 PR에서_ernie_residual_gate_add함수에 적용된 방식과 동일합니다. - 기존 커널 대체: 두 번째
adaLN사이트에서는 기존의residual_gate_add_cuda커널과 RMSNorm, Scale/Shift 연산을 하나의 융합 커널로 대체하여 처리합니다.
# 예시: 융합 커널 사용 가능 여부 확인 및 호출 로직 (개념적)
def _ernie_norm_scale_shift(x, weight, scale, shift, eps):
if can_use_fused_rmsnorm_scale_shift(x, weight, scale, shift):
# 런타임 검증 로직 포함
try:
result = fused_rmsnorm_scale_shift_bitexact(x, weight, scale, shift, eps)
# torch.equal 검증...
return result
except Exception:
pass # Eager 모드로 fallback
# Eager 모드 연산
return eager_rmsnorm_scale_shift(x, weight, scale, shift, eps)
test/registered/kernels/ops/diffusion/test_ernie_norm_scale_shift.py (신규 파일)
이 파일은 새로 구현된 융합 커널의 정확성과 성능을 검증하기 위한 테스트 코드를 포함합니다. 주요 테스트 항목은 다음과 같습니다:
- 비트-정확성 검증: ERNIE 모델의 실제 사용 형태(1024x1024 이미지, CFG 배치 등)와 다양한
tpr값(threads per row)에 대해 융합 커널의 결과가 기존 Eager 연산과torch.equal로 동일한지 확인합니다. - 실제 커널 실행 확인: 테스트 실행 시 융합 커널이 실제로 사용되었는지 (즉, Eager 모드로 fallback되지 않았는지) 검증합니다.
왜 이게 좋은가?
성능 향상
H200 GPU 환경에서 1024x1024 해상도의 이미지를 50 스텝 동안 생성하는 실험에서, 본 PR을 적용했을 때 다음과 같은 성능 향상을 보였습니다:
- End-to-End (E2E) 월(wall) 시간: 기존 15.63초에서 14.99초로 약 4.0% 단축되었습니다.
- Denoising Stage 시간: 14.82초에서 14.35초로 약 3.3% 단축되었습니다.
이는 모델의 핵심 연산인 RMSNorm과 Scale/Shift를 융합함으로써 메모리 대역폭 병목을 크게 줄이고 GPU 연산 효율을 높인 결과입니다. 특히, 각 adaLN 블록에서 개별적으로 실행되던 연산들이 하나로 합쳐지면서 다음과 같은 커널 수준의 속도 향상이 관찰되었습니다:
norm*(1+scale)+shift(RMSNorm + 3 aten): 3.47배 빨라짐 (115.8 us -> 33.3 us)res=residual+gate*update; norm(res)*(1+scale)+shift(residual_gate_add_cuda + RMSNorm + 3 aten): 2.23배 빨라짐 (143.8 us -> 64.5 us)
이러한 커널 수준의 속도 향상이 전체 이미지 생성 시간 단축으로 이어진 것입니다.
비트-정확성 보장
이 PR의 가장 큰 기술적 성과는 성능 향상과 더불어 비트-정확성(bit-exactness)을 완벽하게 보장한다는 점입니다. 이전의 융합 시도들이 수치적 부정확성으로 인해 품질 저하를 일으켰던 것과 달리, 본 PR에서는 기존 CUDA 커널의 산술 연산 순서와 반올림 방식을 Triton 커널에서 그대로 모방했습니다. 그 결과, 다음과 같은 정확성 지표를 달성했습니다:
torch.equal검증: 다양한 ERNIE 모델의 실제 사용 형태와 크기에서 기존 Eager 연산과 비트 단위로 동일함을 확인했습니다.- MD5 동일성: 1024x1024 전체 이미지 생성 결과가 기존 모델과 MD5 해시값까지 완전히 동일함을 보장합니다.
이는 모델의 추론 품질을 전혀 희생하지 않으면서 성능을 개선할 수 있음을 의미하며, 특히 모델의 미세한 수치적 차이에도 민감한 연구 및 프로덕션 환경에서 매우 중요합니다.
일반적인 교훈
- 융합의 힘: 메모리 대역폭이 병목인 연산들은 융합을 통해 상당한 성능 향상을 얻을 수 있습니다. 특히, 연속적으로 호출되는 작은 연산들을 하나로 묶는 것이 효과적입니다.
- 비트-정확성의 중요성: 딥러닝 모델, 특히 생성 모델에서는 미세한 수치적 차이가 최종 결과물의 품질에 큰 영향을 미칠 수 있습니다. 성능 최적화를 추구할 때, 기존 연산의 산술 연산 순서와 반올림 방식을 정확히 모방하는 'rounding-faithful' 접근 방식이 중요합니다.
- Triton의 유연성: Triton은 CUDA 커널의 특정 동작(예: 특정 명령어 사용, 연산 순서)을 세밀하게 제어하고 모방할 수 있는 강력한 도구입니다. 이를 통해 복잡한 기존 커널의 동작을 재현하고 최적화할 수 있습니다.
- 런타임 검증 및 Fallback: 새로운 최적화 기법을 도입할 때는 항상 런타임 시 실제 환경에서의 검증 메커니즘과 실패 시 안전하게 이전 방식으로 돌아가는 Fallback 전략을 마련해야 합니다. 이는 예상치 못한 문제를 방지하고 시스템의 안정성을 높입니다.
References
- Triton Language Documentation
- FlashInfer RMSNormKernel (구현 참고)
- torch.compile
- ERNIE-Image-Turbo (모델 정보)
참고 자료
- https://triton-lang.org/main/getting-started/tutorials/01-vector-add.html
- https://github.com/Dao-AILab/flash-inference/blob/main/csrc/flashinfer/kernels/cutedsl/rmsnorm.cpp
- https://pytorch.org/docs/stable/torch.compile.html
- https://huggingface.co/baidu/ERNIE-Image-Turbo
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
- [vllm] vLLM, DeepSeek-V3.2/GLM-5.2 MTP 경로 최적화: All-Reduce 융합 및 로컬 Argmax 도입
- [sglang] SGLang, 레이어별 오프로딩 기본값 설정을 통한 인코더/VAE 성능 최적화
- [sglang] NPU 성능 향상을 위한 causal_conv1d_update_v2 도입
- [sglang] ERNIE-Image의 RoPE와 GELU-mul 융합 및 RoPE cos/sin 호이스팅을 통한 성능 최적화
- [vllm] vLLM Triton 커널 최적화: tl.constexpr 제거를 통한 JIT 컴파일 오버헤드 해결
PR Analysis 의 다른글
- 이전글 [onnxruntime] [CUDA] NVFP4 QMoE GEMV 최적화: ALU 바운드 커널의 한계를 넘어서는 방법
- 현재글 : [sglang] ERNIE-Image 모델의 성능 향상: 비트-정확성(Bit-Exact)을 유지한 RMSNorm+Scale/Shift 융합 커널 도입
- 다음글 [sglang] FLUX.2 모델의 추론 속도 향상: Residual-Gate 커널 최적화
댓글