[flashinfer] FlashInfer의 FP4 GEMM 최적화: 휴리스틱 개선과 Autotuning 효율화
PR 링크: flashinfer-ai/flashinfer#3948 상태: Merged | 변경: +161 / -127
들어가며
FlashInfer의 mm_fp4(backend='cute-dsl') 연산은 고성능 추론을 위해 다양한 커널 전술(tactic)을 지원합니다. 하지만 기존에는 모든 가능한 전술(~120개)을 일일이 프로파일링하여 최적의 설정을 찾는 방식이었고, 커널 하나를 컴파일하는 데 약 2.5초가 소요되어 전체 오토튜닝(autotuning) 시간이 매우 길어지는 문제가 있었습니다. 본 PR은 휴리스틱 기반의 랭킹 시스템을 도입하고, 상위 N개의 후보군만 오토튜닝함으로써 성능 저하를 최소화하면서 튜닝 시간을 획기적으로 개선했습니다.
코드 분석
1. flashinfer/gemm/gemm_base.py: 랭킹 기반 필터링 도입
기존에는 모든 valid_tactics를 오토튜닝 대상에 포함했으나, 이제는 휴리스틱 점수를 기준으로 상위 후보군만 선별합니다.
Before:
# 모든 전술을 오토튜닝 대상으로 반환
return valid_tactics
After:
# 휴리스틱 점수로 정렬 후 상위 N개만 선택
ranked_configs = sorted(
config_tactics.values(),
key=lambda ts: _score_sm100_mm_fp4_tactic(
m, n, real_k, sm_count, ts[0][0], ts[0][1], ts[0][2]
),
reverse=True,
)
return [t for ts in ranked_configs[: _MM_FP4_CUTE_DSL_MAX_TUNING_CONFIGS // 2] for t in ts]
2. flashinfer/gemm/kernels/utils.py: 통합된 휴리스틱 스코어링
전술의 효율성을 평가하는 _score_sm100_mm_fp4_tactic 함수를 새로 정의하여, 타일 크기, 웨이브 양자화, 클러스터 형태 등을 종합적으로 고려하도록 개선했습니다.
Before:
# 단순화된 로직으로 전술 선택
score = prob_m * total_ctas * ns / (tile_m * num_waves * sm_count)
After:
# 타일 효율성, 웨이브 양자화, 처리량(throughput)을 고려한 다차원 스코어링
score = m_eff * n_eff * throughput / (num_waves * tile_n)
# 2-CTA MMA 및 클러스터 멀티캐스트 페널티 적용
if cta_group == 2: score *= 1.05
score *= 0.95 ** (cga_n.bit_length() - 1)
왜 이게 좋은가
이번 최적화의 핵심은 '전수 조사'에서 '휴리스틱 기반 선별 조사'로의 패러다임 전환입니다.
- 컴파일 시간 단축: 120여 개의 전술을 모두 컴파일하던 기존 방식에서 32개의 후보군으로 제한함으로써 오토튜닝에 소요되는 시간을 대폭 줄였습니다.
- 성능 유지: 벤치마크 결과, 휴리스틱 기반의 Top-1 선택만으로도 기존 대비 1.01x~1.05x의 성능 향상을 보였으며, 오토튜닝을 병행할 경우 성능 손실은 거의 0%에 수렴합니다.
- 일반적 교훈: 대규모 탐색 공간(Search Space)을 가진 커널 튜닝에서는, 모든 경우의 수를 확인하기보다 도메인 지식(Tile/Wave quantization 등)을 활용한 휴리스틱으로 후보군을 좁히는 것이 실무적인 성능과 개발 생산성 사이의 최적의 균형점임을 보여줍니다.
리뷰어 피드백 반영
리뷰 과정에서 2048×36864×4608 형상에서 튜닝 시 성능이 다소 낮게 측정되는 이슈가 제기되었습니다. 이에 대해 측정 분산(variance)에 의한 일시적 현상임을 확인하였고, 전체적인 기하평균(geomean) 성능 손실이 0.5% 미만임을 입증하여 튜닝 시간 단축이라는 이득이 훨씬 크다는 점을 명확히 했습니다.
참고 자료
⚠️ 알림: 이 분석은 AI가 실제 코드 diff를 기반으로 작성했습니다.
관련 포스트
PR Analysis 의 다른글
- 이전글 [ultralytics] [Ultralytics] NDJSON 변환 최적화: 보안과 성능을 동시에 잡는 설계 전략
- 현재글 : [flashinfer] FlashInfer의 FP4 GEMM 최적화: 휴리스틱 개선과 Autotuning 효율화
- 다음글 [sglang] SGLang SM100 CuteDSL Prefill 최적화: State I/O의 커널 퓨전
댓글