Skip to main content
QUICK REVIEW

[논문 리뷰] Gated Linear Attention Transformers with Hardware-Efficient Training

Yang Song-lin, Bailin Wang|arXiv (Cornell University)|2023. 12. 11.
Neural Networks and Applications인용 수 5
한 줄 요약

이 논문은 표준 소프트맥스 어텐션을 데이터에 의존하는 게이팅된 선형 어텐션 메커니즘으로 대체하는 하드웨어 효율적인 트랜스포머 변종인 게이팅된 선형 어텐션(Gated Linear Attention, GLA)을 제안한다. 메모리 액세스 패턴을 최적화하고, I/O를 고려한 새로운 알고리즘인 FLASHLINEARATTENTION을 도입함으로써, 짧은 시퀀스(예: 1K)에서도 FlashAttention-2를 능가하는 더 빠른 훈련을 달성하면서도 선형 시간 추론을 유지한다. GLA 트랜스포머는 언어 모델링에서 LLaMA, RetNet, Mamba와 같은 강력한 베이스라인을 따라하거나 능가하며, 특히 길이 일반화와 기억 집약적인 작업에서 뛰어난 성능을 발휘한다.

ABSTRACT

Transformers with linear attention allow for efficient parallel training but can simultaneously be formulated as an RNN with 2D (matrix-valued) hidden states, thus enjoying linear-time inference complexity. However, linear attention generally underperforms ordinary softmax attention. Moreover, current implementations of linear attention lack I/O-awareness and are thus slower than highly optimized implementations of softmax attention. This work describes a hardware-efficient algorithm for linear attention that trades off memory movement against parallelizability. The resulting implementation, dubbed FLASHLINEARATTENTION, is faster than FLASHATTENTION-2 (Dao, 2023) as a standalone layer even on short sequence lengths (e.g., 1K). We then generalize this algorithm to a more expressive variant of linear attention with data-dependent gates. When used as a replacement for the standard attention layer in Transformers, the resulting gated linear attention (GLA) Transformer is found to perform competitively against the LLaMA-architecture Transformer (Touvron et al., 2023) as well recent linear-time-inference baselines such as RetNet (Sun et al., 2023a) and Mamba (Gu & Dao, 2023) on moderate-scale language modeling experiments. GLA Transformer is especially effective at length generalization, enabling a model trained on 2K to generalize to sequences longer than 20K without significant perplexity degradations. For training speed, the GLA Transformer has higher throughput than a similarly-sized Mamba model.

연구 동기 및 목표

  • 트랜스포머에서 선형 어텐션과 표준 소프트맥스 어텐션 간의 성능 격차를 해소하기 위해.
  • 메모리 이동을 줄이는 I/O 인식형 하드웨어 최적화 알고리즘을 설계하여 선형 어텐션의 훈련 효율성을 향상시키기 위해.
  • 데이터에 의존하는 게이팅을 통해 선형 어텐션의 표현력을 향상시켜 장수열 및 기억 집약적인 작업에서의 성능을 개선하기 위해.
  • 선형 시간 추론 복잡도를 유지하면서도 표준 트랜스포머와 경쟁 가능한 성능을 달성하기 위해.
  • 훈련 시퀀스 길이를 초월한 더 긴 시퀀스 길이로의 일반화 성능가지 향상시키기 위해.

제안 방법

  • 현대 GPU에 최적화된 메모리 액세스와 병렬성을 고려한 하드웨어 효율적인 선형 어텐션 알고리즘인 FLASHLINEARATTENTION을 제안한다.
  • I/O 효율성과 훈련 속도의 균형을 맞추기 위해 쿨럼브 기반 병렬 훈련 방식을 도입하며, 이는 쿨럼브 간 재귀와 쿨럼브 내 병렬 계산을 결합한다.
  • 숨겨진 상태 갱신이 데이터에 의존하는 게이트에 의해 조절되는 게이팅된 선형 어텐션 메커니즘을 설계하여 모델링 능력을 향상시킨다.
  • 누적 합 기법을 사용해 학습 가능한 파rameter α와 β에 대한 폐쇄형 기울기를 유도함으로써 고대역폭 메모리에 중간 상태를 저장할 필요 없이 구현한다.
  • FLASHLINEARATTENTION 알고리즘을 게이팅된 변형을 지원하도록 일반화하여 GLA-Transformer의 효율적 훈련을 가능하게 한다.
  • RNN 유사 재귀를 유지하면서도 쿨럼브 기반 분할을 통해 병렬 훈련을 가능하게 하는 수정된 어텐션 계산 방식을 적용한다.

실험 결과

연구 질문

  • RQ1하드웨어 최적화된 선형 어텐션의 구현이, FlashAttention-2와 같이 정교하게 최적화된 소프트맥스 어텐션보다 짧은 시퀀스에서도 성능을 뛰어넘을 수 있는가?
  • RQ2선형 어텐션에 데이터에 의존하는 게이팅을 도입하면 고정된 감쇠 인자나 일반적인 선형 어텐션보다 성능 향상이著명한가?
  • RQ3GLA-Transformer가 선형 시간 추론 복잡도를 유지하면서도 표준 트랜스포머(예: LLaMA)와 경쟁 가능한 성능을 달성할 수 있는가?
  • RQ4GLA-Transformer는 훈련 길이를 초월한 더 긴 시퀀스 길이로 일반화하는 데 얼마나 효과적인가, 특히 기억 집약적인 작업에서의 성능은?
  • RQ5GLA-Transformer의 훈련 스루풋은 Mamba나 RetNet과 같은 최신 선형 시간 모델과 비교해 어떻게 되는가?

주요 결과

  • FLASHLINEARATTENTION은 짧은 시퀀스(예: 1K 토큰)에서도 FlashAttention-2를 뛰어넘는 속도를 보이며, 뛰어난 I/O 효율성을 입증한다.
  • GLA-Transformer는 언어 모델링 벤치마크에서 경쟁적인 성능을 달성하며, LLaMA 아키텍처 모델과 최근의 선형 시간 모델인 RetNet, Mamba를 따라하거나 능가한다.
  • 340M 및 1.3B 파라미터 모델에서, GLA-Transformer는 11개의 제로샷 태스크 평균 정확도 48.0%와 55.5%를 기록하며, 여러 벤치마크에서 Mamba와 RetNet를 능가한다.
  • GLA-Transformer는 20K 토큰을 초월하는 긴 시퀀스로도 효과적으로 일반화되며, 2K 시퀀스로 훈련된 경우 퍼플렉서티 저하가 최소한이다.
  • 유사한 크기의 Mamba 모델보다 GLA-Transformer가 더 높은 훈련 스루풋을 기록하여, 대규모 프리트레인에서의 더 나은 확장성을 보여준다.
  • BoolQ나 ARC와 같은 기억 집약적인 작업에서 강력한 성능을 보이며, 개선된 기억 유지력과 장거리 의존성 모델링 능력을 시사한다.

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

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

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

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