[sglang] Sana 모델의 BCG 성능 향상: 비트-정확 Triton 커널을 활용한 컨볼루션 후처리 최적화
PR 링크: sgl-project/sglang#34928 상태: Merged | 변경: +296 / -8
들어가며
최근 이미지 생성 모델 분야에서 Stable Diffusion과 같은 확산 모델(Diffusion Models)이 뛰어난 성능을 보여주고 있습니다. 특히, sglang/sgl-project 레포지토리에서는 이러한 모델들의 추론 성능을 극대화하기 위한 다양한 최적화 기법을 적용하고 있습니다.
이번 PR(#21)은 Sana 모델의 Backward Compatibility Guarantee (BCG) 단계에서 발생하는 성능 병목 현상을 해결하는 데 중점을 둡니다. 기존에는 Sana 모델의 GLUMB 컨볼루션 후처리 과정에서 여러 개의 연산(bias, SiLU, GLU)이 개별적으로 수행되어 성능 저하의 원인이 되었습니다. 이 PR은 이러한 후처리 과정을 비트-정확(bit-exact) Triton 커널로 융합하여 BCG 캡처 및 재생 성능을 획기적으로 개선하는 것을 목표로 합니다.
특히, B300 GPU 환경에서 측정된 성능 향상은 주목할 만합니다. 최적화 이전 대비 Denoise 단계에서 약 30%의 지연 시간 감소를 달성했으며, 이는 torch.compile의 성능보다도 뛰어난 결과입니다. 이 글에서는 해당 PR의 코드 변경 사항을 상세히 분석하고, 이러한 최적화가 왜 효과적인지, 그리고 어떤 기술적 교훈을 얻을 수 있는지 살펴보겠습니다.
코드 분석
이번 PR은 주로 sglang/kernels/ops/diffusion/triton/sana_conv_post.py 파일에 새로운 Triton 커널을 추가하고, sglang/multimodal_gen/runtime/models/dits/sana.py 파일에서 기존의 컨볼루션 후처리 로직을 이 새로운 커널을 사용하도록 수정하는 내용을 담고 있습니다.
1. sglang/kernels/ops/diffusion/triton/sana_conv_post.py - 새로운 Triton 커널 추가
이 파일은 Sana 모델의 GLUMB 컨볼루션 후처리를 위한 두 가지 새로운 Triton 커널을 정의합니다:
_bias_silu_kernel: 컨볼루션 결과에 bias를 더하고 SiLU 활성화 함수를 적용하는 과정을 융합합니다._bias_glu_kernel: 컨볼루션 결과에 bias를 더하고 GLU (Gated Linear Unit) 활성화 함수를 적용하는 과정을 융합합니다.
이 커널들은 다음과 같은 특징을 가집니다:
- 비트-정확성 보존: PyTorch의 eager 모드에서 BF16 연산의 중간 반올림 경계를 보존하여
quality=lossless출력의 비트-정확성을 유지합니다.round_bf16_to_fp32함수를 사용하여 이를 구현합니다. - Triton 사용: 고성능 병렬 처리를 위해 Triton 언어를 사용하여 GPU 커널을 작성했습니다.
- 채널스-라스트(channels_last) 최적화:
_is_channels_last_bf16함수를 통해 입력 텐서가 channels_last 형식이고 BF16 타입인지 확인하여 최적화된 경로를 사용할 수 있는지 판단합니다. - 융합 가능성 검증:
can_use_fused_bias_silu및can_use_fused_bias_glu함수는 입력 텐서와 bias 텐서의 속성(CUDA, dtype, device, 차원, 연속성 등)을 검증하여 융합 커널 사용 가능 여부를 결정합니다. - 래퍼 함수:
fused_bias_silu및fused_bias_glu함수는 Triton 커널을 호출하고, 입력 검증 및 예외 처리를 담당합니다.
_bias_silu_kernel 코드 예시:
@triton.jit
def _bias_silu_kernel(
out_ptr,
x_ptr,
bias_ptr,
numel,
channels: tl.constexpr,
):
offsets = tl.program_id(0).to(tl.int64) * 1024 + tl.arange(0, 1024)
mask = offsets < numel
channel = offsets % channels
x = tl.load(x_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
bias = tl.load(bias_ptr + channel, mask=mask, other=0.0).to(tl.float32)
# nn.Conv2d applies its bf16 bias before nn.SiLU, so preserve the
# intermediate bf16 rounding boundary rather than contracting the chain.
biased = round_bf16_to_fp32(x + bias)
tl.store(out_ptr + offsets, biased * tl.sigmoid(biased), mask=mask)
_bias_glu_kernel 코드 예시:
@triton.jit
def _bias_glu_kernel(
out_ptr,
x_ptr,
bias_ptr,
out_numel,
channels: tl.constexpr,
):
offsets = tl.program_id(0).to(tl.int64) * 1024 + tl.arange(0, 1024)
mask = offsets < out_numel
channel = offsets % channels
pixel = offsets // channels
in_base = pixel * (2 * channels) + channel
hidden = tl.load(x_ptr + in_base, mask=mask, other=0.0).to(tl.float32)
gate = tl.load(x_ptr + in_base + channels, mask=mask, other=0.0).to(tl.float32)
hidden_bias = tl.load(bias_ptr + channel, mask=mask, other=0.0).to(tl.float32)
gate_bias = tl.load(bias_ptr + channels + channel, mask=mask, other=0.0).to(
tl.float32
)
hidden = round_bf16_to_fp32(hidden + hidden_bias)
gate = round_bf16_to_fp32(gate + gate_bias)
# SiLU materializes a bf16 tensor before the following multiply in eager.
gate = round_bf16_to_fp32(gate * tl.sigmoid(gate))
tl.store(out_ptr + offsets, hidden * gate, mask=mask)
2. sglang/multimodal_gen/runtime/models/dits/sana.py - 기존 로직 수정
이 파일에서는 Sana 모델의 각 레이어에서 컨볼루션 및 잔차 연결(residual connection) 부분을 수정하여 새로운 Triton 커널을 활용하도록 변경했습니다.
-
컨볼루션 후처리 로직 변경: 기존의
nn.Conv2d호출 후F.silu또는torch.chunk와F.silu를 사용하는 부분을_sana_conv_bias_silu및_sana_conv_bias_glu함수 호출로 대체했습니다.Before:
# ... (이전 코드) ... hidden_states = _mps_safe_conv2d(self.conv_inverted, hidden_states) hidden_states = self.nonlinearity(hidden_states) hidden_states = _mps_safe_conv2d(self.conv_depth, hidden_states) hidden_states, gate = torch.chunk(hidden_states, 2, dim=1) hidden_states = hidden_states * self.nonlinearity(gate) # ... (이후 코드) ...After:
# ... (이전 코드) ... hidden_states = _sana_conv_bias_silu(self.conv_inverted, hidden_states) hidden_states = _sana_conv_bias_glu(self.conv_depth, hidden_states) # ... (이후 코드) ..._sana_conv_bias_silu와_sana_conv_bias_glu함수는 내부적으로can_use_fused_bias_silu/can_use_fused_bias_glu를 호출하여 최적화된 Triton 커널을 사용할 수 있는지 확인하고, 사용할 수 없는 경우 기존의 eager 모드 연산을 수행합니다. 또한,BitExactFusionGate를 사용하여 플랫폼별 비트-정확성 검증을 수행하고, 문제가 발생하면 자동으로 eager 모드로 fallback합니다. -
잔차 연결 최적화:
residual_gate_add함수를 사용하여 어텐션 및 MLP 블록의 잔차 연결을 최적화했습니다. 이는hidden_states + gate * update연산을 더 효율적으로 처리합니다.Before:
hidden_states = hidden_states + gate_msa * attn_output # ... hidden_states = hidden_states + gate_mlp * ff_outputAfter:
hidden_states = _sana_residual_gate_add(hidden_states, attn_output, gate_msa) # ... hidden_states = _sana_residual_gate_add(hidden_states, ff_output, gate_mlp)_sana_residual_gate_add함수는torch.compiler.is_compiling()상태가 아니면sglang.kernels.ops.diffusion.residual_gate_add의 최적화된 커널을 사용합니다. -
패치 임베딩 레이아웃 최적화: 패치 임베딩에서
contiguous()호출을 추가하여, 여러 다운스트림 LayerNorm 연산이 동일한 전치된 뷰를 복사하는 비효율성을 제거했습니다.Before:
hidden_states = hidden_states.flatten(2).transpose(1, 2)After:
# One layout conversion here prevents every downstream LayerNorm from # copying the transposed patch view independently. hidden_states = hidden_states.flatten(2).transpose(1, 2).contiguous()
왜 이게 좋은가?
이 PR의 핵심은 연산 융합(Operator Fusion)과 커널 최적화를 통해 GPU 연산의 효율성을 극대화하는 것입니다. 여러 개의 작은 연산을 하나의 큰 커널로 묶으면 다음과 같은 이점을 얻을 수 있습니다:
- 메모리 접근 감소: 개별 연산 간의 중간 결과를 GPU 레지스터나 공유 메모리에 유지하고, 최종 결과만 메인 메모리에 쓰면서 메모리 접근 횟수를 줄일 수 있습니다. 이는 특히 메모리 대역폭이 병목인 경우 큰 성능 향상을 가져옵니다.
- 커널 실행 오버헤드 감소: 각 연산마다 GPU 커널을 실행하는 오버헤드가 발생합니다. 이를 하나의 커널로 융합하면 커널 실행 횟수가 줄어들어 전체적인 지연 시간을 단축할 수 있습니다.
- 데이터 재사용성 증대: 융합된 커널 내에서 데이터가 레지스터나 L1 캐시 등 빠른 메모리에 더 오래 머무를 수 있어 데이터 재사용률이 높아집니다.
B300 성능 향상 분석:
PR 설명에 따르면, B300 GPU에서 Sana 1.5 1.6B 모델을 1024px 워크로드로 실행했을 때 다음과 같은 성능 향상이 관찰되었습니다:
| Mode | Denoise | E2E |
|---|---|---|
| eager | 0.350 s | 0.476 s |
| eager + BCG, before this optimization | ~0.333 s | — |
| eager + BCG, this change | 0.233 s | 0.363 s |
torch.compile |
0.272 s | 0.388 s |
- Denoise 단계: 최적화된 BCG 경로는 0.233초로, 최적화 이전의 0.333초 대비 약 30%의 성능 향상을 보였습니다. 이는
torch.compile의 0.272초보다도 약 14.2% 더 빠릅니다. - E2E (End-to-End) 단계: 최적화된 BCG 경로는 0.363초로,
torch.compile의 0.388초보다 약 6.4% 더 빠릅니다.
이러한 성능 향상은 기존에 분산되어 있던 bias + SiLU, bias + GLU 연산을 비트-정확 Triton 커널로 융합하고, 잔차 연결 로직을 최적화한 결과입니다. 특히, torch.compile보다도 빠른 성능을 달성했다는 점은 수동으로 최적화된 커널이 특정 워크로드에서 얼마나 강력한 성능을 발휘할 수 있는지를 보여줍니다.
일반적인 교훈:
- 병목 구간 식별 및 최적화: 프로파일링을 통해 성능 병목 구간을 정확히 식별하고, 해당 구간에 집중하여 최적화를 수행하는 것이 중요합니다. 이 PR에서는 BCG 후처리 단계의 여러 연산이 병목임을 파악했습니다.
- 연산 융합의 힘: GPU에서 반복적으로 발생하는 작은 연산들을 융합하는 것은 메모리 접근 및 커널 실행 오버헤드를 줄여 상당한 성능 향상을 가져올 수 있습니다.
- Triton의 활용: Triton은 PyTorch와 같은 프레임워크에서 고성능 사용자 정의 커널을 쉽게 작성할 수 있도록 지원하는 강력한 도구입니다. 복잡한 연산이나 특정 하드웨어에 최적화된 로직을 구현할 때 유용합니다.
- 비트-정확성 유지의 중요성: 특히 이미지 생성 모델과 같이 출력 품질이 중요한 경우, 성능 최적화 과정에서 비트-정확성을 유지하는 것이 필수적입니다. 이 PR은 이를 위해
round_bf16_to_fp32함수와BitExactFusionGate를 사용하여 검증 메커니즘을 도입했습니다. - 레이아웃 최적화: 데이터 레이아웃(예: channels_last)을 적절히 관리하고, 불필요한 데이터 복사나 변환을 줄이는 것도 성능에 영향을 미칩니다. 패치 임베딩 부분의
contiguous()호출이 좋은 예시입니다.
리뷰 피드백 분석
제공된 PR 정보에는 리뷰 댓글이 포함되어 있지 않아, 해당 부분에 대한 분석은 생략합니다. 하지만 실제 프로젝트에서는 리뷰어들의 피드백을 통해 코드의 안정성, 정확성, 유지보수성 등을 더욱 향상시킬 수 있습니다.
References
- Triton Language Documentation
- torch.compile
- BitExactFusionGate (해당 PR에서 사용된 검증 로직의 소스 코드)
- residual_gate_add (해당 PR에서 사용된 잔차 연결 최적화 커널의 소스 코드)
- fused_layernorm_modulate_raw (참고: 이 PR에서 직접 수정되지는 않았으나, 관련 최적화 기법)
참고 자료
- https://triton-lang.org/docs/
- https://pytorch.org/docs/stable/generated/torch.compile.html
- https://github.com/sglang/sglang/blob/main/python/sglang/kernels/ops/diffusion/bitexact_gate.py
- https://github.com/sglang/sglang/blob/main/python/sglang/kernels/ops/diffusion/residual_gate_add.py
- https://github.com/sglang/sglang/blob/main/python/sglang/kernels/ops/diffusion/triton/layernorm_mod.py
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
PR Analysis 의 다른글
- 이전글 [sglang] Cosmos3 T2I 가속화: QKNorm + RoPE 커널 퓨전과 BF16 정밀도 최적화
- 현재글 : [sglang] Sana 모델의 BCG 성능 향상: 비트-정확 Triton 커널을 활용한 컨볼루션 후처리 최적화
- 다음글 [sglang] [AMD gfx950] GLM-5.2 MLA 최적화: FP8 양자화와 Zero-Copy 레이아웃 전환
댓글