FlashQLA: GDN용 CP·역전파 친화적 Fused Linear Attention 커널 라이브러리
FlashQLA는 GDN Chunked Prefill을 CP·fused kernel로 최적화해 FLA Triton 대비 최대 3배 빨라졌다.
FlashQLA가 공개됐다. TileLang 기반 linear attention kernel 라이브러리로, GDN Chunked Prefill의 forward와 backward를 함께 최적화해 NVIDIA Hopper에서 FLA Triton 대비 forward는 2~3배, backward는 2배 빨라졌다. 효율 향상은 특히 pretraining과 edge-side agentic inference에서 두드러졌다.
Qwen3-Next 이후 GDN이 Qwen3.5/3.6까지 확산되면서, 397A17B / 122A10B / 35B / 27B 급 모델과 256K+ 컨텍스트에서 GDN block 오버헤드가 눈에 띄기 시작했다.
기존 FLA GDN Chunked Prefill의 병목은 두 가지였다.
- HBM 왕복이 잦은 memory-bound 구조라 K, V, W, U, S 같은 중간 상태를 계속 읽고 써야 했다.
- 재귀 상태 때문에 thread block 수가 제한돼, 소형 모델·소배치·TP 환경에서 SM utilization이 낮았다.
해결책은 전체를 완전 fusion하는 대신, CP(컨텍스트 병렬) 를 고려해 forward를 두 개의 fused kernel로 나누고 중간에 CP preprocessing을 넣는 방식이다. 병렬도는 chunk 수 N 과 rank당 chunk 수 L 을 두고 L = λ√N 로 정했으며, 실전에서는 batch_size * num_heads <= 40 또는 batch_size * num_heads <= 56 && seq_len >= 8192일 때만 CP를 켰다. 그렇지 않은 구간에서는 v_head_dim 분할로 2~4배 병렬도를 올리는 기존 경로를 유지했다.
게이트 감쇠를 이용한 warmup 도 추가됐다. 기사에 따르면 linear attention head의 **60~80%**에서 sliding-window 성질이 관찰됐고, 6~8 chunk만 warmup해도 state 오차가 잡음 수준 아래로 내려갔다. 이 경우는 correction matrix M 을 계산하지 않고도 충분한 정확도를 얻었고, CP preprocessing은 M 과 S 를 모두 계산하는 경로와 S 만 계산하는 경로를 하나의 fused kernel로 처리했다. warmup 길이는 gate statistics를 모으는 별도 kernel이 정하며, 비용은 거의 없다.
구현은 producer 1개와 consumer 3개 warpgroup 이 shared memory와 mbarrier 로 협업하는 warpgroup specialization 패턴이다. forward는 V', S, O 를 겹쳐 계산하고, backward는 앞선 CP preprocessing kernel로 S 를 재계산한 뒤 on-chip resource constraints 때문에 multi-stage pipelining 대신 긴 compute chain으로 메모리 트래픽을 가렸다. 이후 bwd_dv, bwd_dhu, bwd_dqkwg, bwd_wy 를 하나로 묶어 추가적인 중복 접근을 줄였다.
벤치마크는 FLA 0.5.0, Triton 3.5.1, FlashInfer 0.6.9, TileLang 0.1.8 기준으로 측정됐고, H200에서 397B/122B TP8 1x32768 시나리오는 0.310ms로 FlashInfer 1.653ms, FLA 0.913ms를 앞섰다. 397B/122B TP4, 27B TP2, 2B/0.8B TP1, Sym h32에서도 일관된 개선이 나왔고, TP가 커질수록 intra-card AutoCP의 효과가 더 커졌다. 코드와 벤치마크는 github.com/QwenLM/FlashQLA 에 공개됐다.
이 요약은 원문 이해를 돕기 위한 큐레이션입니다. 저작권은 원저작자에게 있으며, 정확한 내용과 맥락은 원문을 확인하세요.
요약 오류, 출처 표기 문제, 삭제 요청은 문의 · 건의로 알려주세요.