본문으로 건너뛰기

[sglang] FLUX.2 모델의 추론 속도 향상: Residual-Gate 커널 최적화

PR 링크: sgl-project/sglang#33823 상태: Merged | 변경: +51 / -8

들어가며

FLUX.2와 같은 최신 Diffusion Transformer(DiT) 모델은 매 스텝마다 수많은 residual stream 업데이트를 수행합니다. 기존에는 gate * updateresidual + ...를 별도의 두 커널로 처리하는 eager 방식을 사용했으나, 이는 GPU 커널 호출 오버헤드를 유발합니다. 본 PR은 이미 검증된 residual_gate_add_cuda 커널을 FLUX.2의 모든 게이트 사이트에 적용하여, 비트 단위로 동일한(bit-exact) 결과를 유지하면서도 추론 성능을 최적화했습니다.

코드 분석

flux_2.py에서의 커널 통합

핵심 변경 사항은 _flux2_residual_gate_add 헬퍼 함수를 정의하고 이를 Flux2SingleTransformerBlockFlux2TransformerBlock 내의 모든 게이트 연산 지점에 적용한 것입니다.

Before:

hidden_states = hidden_states + mod_gate * attn_output

After:

hidden_states = _flux2_residual_gate_add(hidden_states, attn_output, mod_gate)

이 함수는 half dtype(float16, bfloat16)에서만 커널을 호출하며, can_use_residual_gate_add_cuda를 통해 커널 사용 가능 여부를 사전에 체크합니다. 또한 torch.compiler.is_compiling()을 사용하여 컴파일된 그래프 내에서 예기치 않은 폴백(fallback)이 발생하지 않도록 설계되었습니다.

테스트 코드 확장

test_residual_gate_add.py를 수정하여 FLUX.2에서 사용하는 실제 텐서 형상(D=3072, D=6144)을 커널 테스트 케이스에 추가했습니다.

# FLUX.2-dev (D=6144) joint sequence.
((1, 4608, 6144), (1, 1, 6144)),

왜 이게 좋은가

성능 향상

H200 GPU 환경에서 FLUX.2-klein-4B 모델을 사용하여 50단계 디노이징을 수행한 결과, 전체 추론 시간에서 약 1.2%의 성능 향상을 확인했습니다. 특히 커널 마이크로벤치마크에서는 특정 형상에서 최대 2.5배 이상의 속도 향상을 보였습니다.

최적화의 교훈

  1. Bit-Exactness 유지: 성능을 위해 정확도를 희생하지 않았습니다. half dtype에서 eager 방식과 동일한 반올림 결과를 보장하여 모델 품질 저하 없이 최적화를 달성했습니다.
  2. Fallback 전략: 커널 호출 실패 시 자동으로 eager 방식으로 전환되는 안전장치를 마련하여 런타임 안정성을 확보했습니다.
  3. 커널 퓨전(Kernel Fusion): 여러 연산을 하나의 커널로 통합하는 것은 메모리 대역폭 제한적인(memory-bound) 연산에서 매우 효과적입니다. 특히 gate * updateresidual 더하기를 하나의 커널로 묶어 GPU 커널 실행 횟수를 획기적으로 줄였습니다.

리뷰어 피드백

리뷰 과정에서 torch.compiler.is_compiling()을 통한 컴파일러 안전성 확보와 half dtype 한정 적용에 대한 논의가 이루어졌으며, 이는 프로덕션 환경에서의 안정적인 배포를 가능하게 했습니다.

참고 자료

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

댓글

관련 포스트

PR Analysis 의 다른글