Skip to main content
QUICK REVIEW

[논문 리뷰] Faster Causal Attention Over Large Sequences Through Sparse Flash Attention

Matteo Pagliardini, Daniele Paliotta|arXiv (Cornell University)|2023. 06. 01.
Topic Modeling인용 수 5
한 줄 요약

이 논문은 인과 자기주의에서 동적이고 비정규적인 희소성 패턴을 지원하는 FlashAttention의 GPU 커널 확장을 제공하는 Sparse Causal Flash Attention (SCFA)를 소개한다. 이는 장문의 시퀀스에서 효율적인 계산을 가능하게 하며, 모델의 퍼플렉서티를 희생시키지 않은 채 FlashAttention 대비 최대 3.3배 빠른 훈련 속도를 달성한다. SCFA는 최적화된 희소 커널 실행을 통해 해시 기반 및 쿼리/키 제거 희소성 패턴을 효율적으로 처리한다.

ABSTRACT

Transformer-based language models have found many diverse applications requiring them to process sequences of increasing length. For these applications, the causal self-attention -- which is the only component scaling quadratically w.r.t. the sequence length -- becomes a central concern. While many works have proposed schemes to sparsify the attention patterns and reduce the computational overhead of self-attention, those are often limited by implementations concerns and end up imposing a simple and static structure over the attention matrix. Conversely, implementing more dynamic sparse attentions often results in runtimes significantly slower than computing the full attention using the Flash implementation from Dao et al. (2022). We extend FlashAttention to accommodate a large class of attention sparsity patterns that, in particular, encompass key/query dropping and hashing-based attention. This leads to implementations with no computational complexity overhead and a multi-fold runtime speedup on top of FlashAttention. Even with relatively low degrees of sparsity, our method improves visibly upon FlashAttention as the sequence length increases. Without sacrificing perplexity, we increase the training speed of a transformer language model by $2.0 imes$ and $3.3 imes$ for sequences of respectively $8k$ and $16k$ tokens.

연구 동기 및 목표

  • 장시퀀스 트랜스포터 모델에서 인과 자기주의의 계산 병목 현상을 해결하기 위해.
  • 해싱 또는 쿼리/키 제거와 같은 동적이고 비정규적인 주의 희소성 패턴에 대해 효율적이고 고성능의 추론과 훈련을 가능하게 하기 위해.
  • 현대의 희소 주의 메커니즘에서 흔한 비삼각형 인과 마스크에 대해 FlashAttention의 효율성을 확장하기 위해.
  • 계산 복잡도를 증가시키지 않으면서도 모델 품질을 훼손하지 않고 FlashAttention 대비 실용적인 속도 향상을 달성하기 위해.
  • 넓은 범위의 희소성 패턴을 최소한의 구현 오버헤드로 지원하는 유연하고 오픈소스 GPU 커널을 제공하기 위해.

제안 방법

  • 쿼리당 허용 가능한 키의 범위로 표현 가능한 임의의 희소성 패턴을 지원하도록 FlashAttention을 확장하여, 비정규적인 인과 마스크를 가능하게 한다.
  • Triton을 활용해 저수준 최적화를 수행하는 GPU 커널을 도입하여, 희소 인과 주의를 계산할 때 계산 복잡도 오버헤드 없이 처리한다.
  • 기하학적 해싱(예: Reformer 스타일의 LSH)을 통해 동적 희소성에 적용하지만, 근사치 없이 정확한 계산을 보장한다.
  • 헤드별 세밀한 쿼리 및 키 제거를 구현하여, 전체 헤드 프루닝 없이도 비례적인 계산 감소를 가능하게 한다.
  • 희소 액세스 패턴에도 불구하고 높은 메모리 대역폭 활용도를 유지하기 위해 블록 기반 메모리 액세스 패턴을 사용한다.
  • 커널 런처 오버헤드를 최소화하고 현대 GPU의 최적의 사용률을 극대화하기 위해 맞춤형 커널 융합 전략을 적용한다.
Figure 3: Comparing several hash-based sparse attention implementations with FlashAttention. Similarly to QK-dropping-based sparsity in Fig. 7 , due to the non-triangular causal mask resulting from re-ordering the tensors based on the hash buckets (see Fig. 1 ), a naive implementation would force th
Figure 3: Comparing several hash-based sparse attention implementations with FlashAttention. Similarly to QK-dropping-based sparsity in Fig. 7 , due to the non-triangular causal mask resulting from re-ordering the tensors based on the hash buckets (see Fig. 1 ), a naive implementation would force th

실험 결과

연구 질문

  • RQ1FlashAttention을 비삼각형 인과 마스크를 지원하도록 확장할 수 있을까? 이는 비정규적인 희소성 패턴에 대한 효율적 계산을 가능하게 할 수 있는가?
  • RQ2추가적인 계산 복잡도를 유발하지 않으면서도, 희소 주의 메커니즘에 대해 FlashAttention 대비 실용적인 속도 향상을 달성할 수 있을까?
  • RQ3동적이고 세밀한 희소성(예: 쿼리/키 제거)은 모델 성능을 유지하면서도 훈련 속도 향상에 눈에 띄는 영향을 미칠 수 있을까?
  • RQ4SCFA를 통한 해시 기반 주의의 성능은 원본 Reformer LSH와 비교해 속도와 커버리지 측면에서 어떻게 다를까?
  • RQ5시퀀스 길이가 증가함에 따라, 특히 8k 토큰을 초과할 경우 SCFA는 훈련 효율을 유지하거나 향상시킬 수 있을까?

주요 결과

  • SCFA는 8k 및 16k 토큰 시퀀스에서 각각 최대 2.0배와 3.3배 빠른 훈련 속도를 달성했으며, 퍼플렉서티에 영향을 주지 않았다.
  • 8192토큰 시퀀스에서 8개 및 16개 버킷을 사용할 경우, SCFA를 통한 해시 기반 희소 주의는 FlashAttention 기준 대비 훈련 시간을 각각 1.4배와 1.8배 단축시켰다.
  • 훈련 중 일관된 속도 향상이 이루어졌으며, H-LM 모델은 초기 단계부터 가속화되며 높은 처리량을 유지했다.
  • SCFA의 해시 기반 주의는 Reformer의 LSH보다 더 빠른 런타임을 보였고, 정확한 계산을 수행함으로써 저커버리지 및 근사 오류 문제를 피했다.
  • SCFA를 통한 쿼리 및 키 제거는 제거된 쌍의 비율에 비례해 계산을 비례적으로 감소시켜, 세밀한 효율성 제어를 가능하게 했다.
  • 런타임 향상은 시퀀스 길이와 희소성에 따라 스케일링되며, 버킷 수가 일정 수준 이상에 도달하면 점진적으로 감소하는 경향을 보였다.
(a) Forward attn. runtimes
(a) Forward attn. runtimes

더 나은 연구,지금 바로 시작하세요

논문 읽기부터 검토까지, 연구 시간을 획기적으로 줄여보세요.

카드 등록 없음 · 무료 플랜 제공

이 리뷰는 AI가 만들고, 인간 에디터가 검토했습니다.