FlashQLA: GDN을 위한 CP-/Bwd-친화적 융합 선형 어텐션 커널

Qwen은 Gated Delta Network (GDN) 레이어를 최적화하도록 설계된 고성능 선형 어텐션 커널 라이브러리인 FlashQLA를 출시했습니다. TileLang을 기반으로 하며, FlashQLA는 NVIDIA Hopper GPU에서 FLA Triton 커널에 비해 전방 2-3배, 후방 2배의 속도 향상을 달성하여 사전 학습 및 엣지 측 에이전시 추론에 특히 도움이 됩니다.

GDN 청크 프리필 최적화

FlashQLA는 원래 FLA 구현에서 발견된 Gated Delta Network (GDN) Chunked Prefill 프로세스의 두 가지 주요 효율성 병목 현상을 해결합니다:

  1. Memory-Bound Kernels: 표준 흐름은 중간 변수 ($W, U, S$)를 고대역폭 메모리(HBM)로 반복적으로 읽고 쓰며, 상당한 오버헤드를 발생시킵니다.
  2. Low GPU Utilization: State Space Model (SSM) 상태의 재귀적 특성으로 인해 동시에 실행 가능한 스레드 블록 수가 batch_size * num_heads 로 제한됩니다. 모델이 작거나 배치가 작거나 Tensor Parallelism (TP) 환경에서는 GPU 스트리밍 멀티프로세서(SM)가 유휴 상태가 됩니다.

이러한 상충되는 문제를 해결하기 위해 FlashQLA는 완전 융합 커널(소규모 배치 상황에서 실패함)을 피하고, 대신 전방 계산을 두 개의 융합 커널로 나누고 그 사이에 Context Parallelism (CP) 전처리 단계를 삽입합니다.

핵심 기술 혁신

Gate-Driven Automatic Intra-Card Context Parallelism (AutoCP)

FlashQLA는 TP, 긴 시퀀스, 그리고 헤드 수가 적은 환경에서 SM 활용도를 높이기 위해 자동 인트라 카드 CP 메커니즘을 구현합니다. 이는 수학적 모델을 사용해 최적의 병렬도 ($L = \lambda \sqrt{N}$)를 결정하는데, 여기서 $N$은 청크 수이며 $L$은 CP 랭크당 청크 수를 의미합니다.

오버헤드를 추가로 줄이기 위해 FlashQLA는 GDN 게이트의 지수 감쇠 특성을 활용합니다. 게이트 $\alpha_i \in (0,1)$ 인 헤드의 경우 이전 상태의 영향이 지수적으로 감소합니다. FlashQLA는 일반적으로 6–8 청크의 "warmup" 과정을 사용해 상태 오류를 노이즈 플로어 이하로 낮추어, 비용이 많이 드는 보정 항 $M$ 행렬 계산을 생략하고 정확한 부분 시퀀스 $S_0$를 직접 얻을 수 있게 합니다.

TileLang 워프 특화 커널

FlashQLA는 TileLang를 활용해 워프그룹 특화 커널을 구현합니다. 이 아키텍처는 동일한 SM 내에서 하나의 프로듀서 워프그룹과 세 개의 컨슈머 워프그룹을 사용하며, 공유 메모리를 통해 데이터를 교환하고 mbarriers로 동기화합니다.

  • Forward Pass: 세 개의 컨슈머 워프그룹이 각각 $V'$, $S$, $O$ 를 계산하며, ping-pong 구조를 사용해 계산과 메모리 트래픽을 겹칩니다.
  • CP Preprocessing: 단일 융합 커널이 원래 $M$ 및 $S$ 계산과 가벼운 슬라이딩 윈도우 warmup 방식을 모두 처리합니다.
  • Backward Pass: FlashQLA는 bwd_dv, bwd_dhu, bwd_dqkwg, bwd_wy 를 하나의 커널로 융합합니다. 온칩 자원 제한으로 인해 다단계 파이프라인 대신 긴 연산 체인을 이용해 메모리 트래픽을 숨깁니다.

성능 벤치마크

NVIDIA H200 GPU에서 FLA Triton 및 FlashInfer 기준과 비교한 벤치마크 결과, 특히 Tensor Parallelism (TP) 정도가 증가할수록 큰 성능 향상이 나타났습니다.

모델 / TP 시퀀스 길이 $h_{qk}$ $h_v$ FlashQLA FlashInfer FLA 대비 FlashQLA vs FI
397B/122B TP8 1x32768 2 8 0.310ms 1.653ms 2.95×
397B/122B TP4 1x32768 4 16 0.486ms 1.654ms 2.57×
27B TP2 1x32768 8 24 0.659ms 1.616ms 2.37×
2B/0.8B TP1 1x32768 16 16 0.493ms 1.640ms 2.60×

구현 및 요구 사항

FlashQLA는 FLA 서명과 호환되는 고수준 API와 전방 및 후방 패스를 위한 저수준 진입점을 제공합니다.

시스템 요구 사항:

  • 하드웨어: NVIDIA SM90 (Hopper)
  • 소프트웨어: CUDA 12.8+, PyTorch 2.8+

코드와 벤치마크는 github.com/QwenLM/FlashQLA 에서 확인할 수 있습니다.

Sources