본문으로 건너뛰기

[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 커널로 구현합니다:

  1. fused_rmsnorm_scale_shift_bitexact: norm(x) * (1 + scale) + shift 연산을 융합합니다.
  2. 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.reshapetl.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 해시값까지 완전히 동일함을 보장합니다.

이는 모델의 추론 품질을 전혀 희생하지 않으면서 성능을 개선할 수 있음을 의미하며, 특히 모델의 미세한 수치적 차이에도 민감한 연구 및 프로덕션 환경에서 매우 중요합니다.

일반적인 교훈

  1. 융합의 힘: 메모리 대역폭이 병목인 연산들은 융합을 통해 상당한 성능 향상을 얻을 수 있습니다. 특히, 연속적으로 호출되는 작은 연산들을 하나로 묶는 것이 효과적입니다.
  2. 비트-정확성의 중요성: 딥러닝 모델, 특히 생성 모델에서는 미세한 수치적 차이가 최종 결과물의 품질에 큰 영향을 미칠 수 있습니다. 성능 최적화를 추구할 때, 기존 연산의 산술 연산 순서와 반올림 방식을 정확히 모방하는 'rounding-faithful' 접근 방식이 중요합니다.
  3. Triton의 유연성: Triton은 CUDA 커널의 특정 동작(예: 특정 명령어 사용, 연산 순서)을 세밀하게 제어하고 모방할 수 있는 강력한 도구입니다. 이를 통해 복잡한 기존 커널의 동작을 재현하고 최적화할 수 있습니다.
  4. 런타임 검증 및 Fallback: 새로운 최적화 기법을 도입할 때는 항상 런타임 시 실제 환경에서의 검증 메커니즘과 실패 시 안전하게 이전 방식으로 돌아가는 Fallback 전략을 마련해야 합니다. 이는 예상치 못한 문제를 방지하고 시스템의 안정성을 높입니다.

References

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글