AI 모델 및 플랫폼

플래시 어텐션: 트랜스포머 효율성 혁신

mm
Unite.AI를 Google의 선호 소스에 추가

트랜스포머 모델이 크기와 복잡성이 증가함에 따라, 특히 긴 시퀀스를 다룰 때 계산 효율성과 메모리 사용량에 대한重大한 도전을 직면합니다. 플래시 어텐션은 트랜스포머 모델에서 어텐션 메커니즘을 구현하고 확장하는 방식을 혁신적으로 바꿀 수 있는 최적화 기법입니다.

이 всесторон한 가이드에서, 우리는 플래시 어텐션의 핵심 개념, 구현 세부 사항, 및 기계 학습 분야에 미치는深刻한 영향을 탐구할 것입니다.

문제: 어텐션이 비싼 것

플래시 어텐션을 이해하기 전에, 먼저 플래시 어텐션이 해결하려는 문제를 이해해야 합니다. 어텐션 메커니즘은 강력하지만, 특히 긴 시퀀스에서重大한 계산 비용이 있습니다.

표준 어텐션: 간단한 요약

트랜스포머 모델의 표준 어텐션 메커니즘은 다음 방정식으로 요약할 수 있습니다:

Attention(Q, K, V) = softmax(QK^T / √d) V

여기서 Q, K, 및 V는 쿼리, 키, 및 값 행렬입니다. 그리고 d는 키 벡터의 차원입니다.

이 공식은 우아하지만, 구현에서는 몇 가지 비효율성을 초래합니다:

  1. 메모리 병목: 중간 어텐션 행렬(QK^T)의 크기는 N x N입니다. 여기서 N은 시퀀스 길이입니다. 긴 시퀀스에서는 이는 빠르게 사용 가능한 GPU 메모리를 소진할 수 있습니다.
  2. 중복 메모리 액세스: 표준 구현에서, 어텐션 행렬은 계산되어 고대역폭 메모리(HBM)에 저장되고 softmax 연산을 위해 다시 읽어옵니다. 이 중복 메모리 액세스는 주요 병목 현상입니다.
  3. GPU 컴퓨팅의 활용도: 현대 GPU는 메모리 대역폭보다 훨씬 더 많은 컴퓨팅 능력을 가지고 있습니다. 표준 어텐션 구현은 메모리 제한이므로, GPU의 많은 컴퓨팅 가능성이 활용되지 않습니다.

다음은 표준 어텐션 구현을 보여주는 단순한 파이썬 코드 스니펫입니다:

</pre>
import torch

<p>def standard_attention(Q, K, V):
# Q, K, V shape: (batch_size, seq_len, d_model)
d_k = K.size(-1)
scores = torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(d_k))
attention_weights = torch.softmax(scores, dim=-1)
return torch.matmul(attention_weights, V)</p>

이 구현은 간단하지만, 위에서 언급한 비효율성을 가지고 있습니다. scores 텐서는 (batch_size, seq_len, seq_len) 모양을 가지며, 긴 시퀀스에서는 금지적으로 커질 수 있습니다.

플래시 어텐션 등장

플래시 어텐션은 Tri Dao와 동료们에 의해 2022년에 도입된 접근 방식으로, 어텐션을 계산하는 방식입니다. 이는 메모리 사용량을 크게 줄이고 계산 효율성을 개선합니다. 플래시 어텐션의 핵심 아이디어는:

  1. 타일링: 큰 어텐션 행렬을 빠른 온칩 SRAM에 맞는 작은 타일로 나눕니다.
  2. 재계산: 전체 어텐션 행렬을 저장하는 대신, 역방향 전달 동안 필요한 부분을 재계산합니다.
  3. IO-Aware 구현: 알고리즘을 최적화하여 데이터를 다른 레벨의 GPU 메모리 계층 사이에서 이동하는 것을 최소화합니다.

플래시 어텐션 알고리즘

플래시 어텐션의 핵심은 어텐션 메커니즘을 계산하는 방식입니다. 전체 어텐션 행렬을 한 번에 계산하는 대신, 블록으로 처리하여 현대 GPU의 메모리 계층을 활용합니다.

다음은 알고리즘의 고수준 개요입니다:

  1. 입력: 행렬 Q, K, V가 HBM에 있고 온칩 SRAM의 크기는 M입니다.
  2. 블록 크기는 사용 가능한 SRAM에 따라 계산됩니다.
  3. 출력 행렬 O와 보조 벡터 l 및 m을 초기화합니다.
  4. 입력 행렬을 SRAM에 맞는 블록으로 나눕니다.
  5. 두 개의 중첩된 루프가 이러한 블록을 처리합니다:
    • 외부 루프는 K와 V 블록을 로드합니다
    • 내부 루프는 Q 블록을 로드하고 계산을 수행합니다
  6. 온칩 계산에는 행렬 곱셈, softmax, 및 출력 계산이 포함됩니다.
  7. 결과는 각 블록을 처리한 후 HBM에 다시 작성됩니다.

이 블록별 계산은 플래시 어텐션이 작은 메모리 풋프린트를 유지하면서 정확한 어텐션을 계산할 수 있도록 합니다.

플래시 어텐션의 수학

플래시 어텐션이 작동하는 핵심은 블록별로 softmax를 계산할 수 있는 수학적 트릭입니다. 이 논문은 두 가지 주요 공식을 소개합니다:

  1. 소프트맥스 분해:
    softmax(x) = exp(x - m) / Σexp(x - m)

    여기서 m은 x의 최대 값입니다.

  2. 소프트맥스 머지:
    softmax(x ∪ y) = softmax(softmax(x) * e^(m_x - m), softmax(y) * e^(m_y - m))

    여기서 m = max(m_x, m_y)입니다.

이 공식들은 플래시 어텐션이 각 블록에 대한 부분적인 softmax 결과를 계산하고 이를 올바르게 결합하여 최종 결과를 얻을 수 있도록 합니다.

구현 세부 사항

플래시 어텐션의 핵심 개념을 보여주는 단순화된 구현입니다:

import torch

<p>def flash_attention(Q, K, V, block_size=256):
batch_size, seq_len, d_model = Q.shape</p>

<p># 출력과 실행 중인 통계를 초기화
O = torch.zeros_like(Q)
L = torch.zeros((batch_size, seq_len, 1))
M = torch.full((batch_size, seq_len, 1), float(&#039;-inf&#039;))</p>

<p>for i in range(0, seq_len, block_size):
Q_block = Q[:, i:i+block_size, :]</p>

<p>for j in range(0, seq_len, block_size):
K_block = K[:, j:j+block_size, :]
V_block = V[:, j:j+block_size, :]</p>

<p># 이 블록에 대한 어텐션 점수를 계산
S_block = torch.matmul(Q_block, K_block.transpose(-2, -1)) / (d_model ** 0.5)</p>

<p># 실행 중인 최대값을 업데이트합니다
M_new = torch.maximum(M[:, i:i+block_size], S_block.max(dim=-1, keepdim=True).values)</p>

<p># 지수함수를 계산
exp_S = torch.exp(S_block - M_new)
exp_M_diff = torch.exp(M[:, i:i+block_size] - M_new)</p>

<p># 실행 중인 합계를 업데이트합니다
L_new = exp_M_diff * L[:, i:i+block_size] + exp_S.sum(dim=-1, keepdim=True)</p>

<p># 이 블록에 대한 출력을 계산
O[:, i:i+block_size] = (
exp_M_diff * O[:, i:i+block_size] +
torch.matmul(exp_S, V_block)
) / L_new</p>

<p># 실행 중인 통계를 업데이트합니다
L[:, i:i+block_size] = L_new
M[:, i:i+block_size] = M_new</p>

return O

이 구현은 플래시 어텐션의 본질을 담고 있습니다. 입력을 블록으로 처리하며, 실행 중인 통계(M 및 L)를 유지하여 모든 블록에 걸쳐 softmax를 올바르게 계산합니다.

플래시 어텐션의 영향

플래시 어텐션의 도입은 기계 학습 분야, 특히 대규모 언어 모델 및 긴 컨텍스트 응용 분야에重大한 영향을 미쳤습니다. 주요 이점은:

  1. 메모리 사용량 감소: 플래시 어텐션은 메모리 복잡도를 O(N^2)에서 O(N)으로 줄입니다. 여기서 N은 시퀀스 길이입니다. 이는 동일한 하드웨어에서 훨씬 더 긴 시퀀스를 처리할 수 있음을 의미합니다.
  2. 속도 개선: 데이터 이동을 최소화하고 GPU 컴퓨팅 능력을 더 잘 활용함으로써, 플래시 어텐션은重大한 속도 향상을 달성합니다. 저자는 표준 구현에 비해 GPT-2에서 최대 3배 빠른 훈련을 보고했습니다.
  3. 정확한 계산: 다른 어텐션 최적화 기법과 달리, 플래시 어텐션은 근사값이 아닌 정확한 어텐션을 계산합니다.
  4. 확장성: 감소된 메모리 풋프린트는 훨씬 더 긴 시퀀스, 잠재적으로 수백만 토큰까지 확장할 수 있습니다.

실제 영향

플래시 어텐션의 영향은 학술 연구를 넘어선다. 많은 인기 있는 기계 학습 라이브러리와 모델에서 빠르게 채택되었습니다:

  • 허깅페이스 트랜스포머: 인기 있는 트랜스포머 라이브러리는 플래시 어텐션을 통합하여 사용자가 이를 쉽게 활용할 수 있도록 했습니다.
  • GPT-4 및 그 이상: 고급 언어 모델에서 플래시 어텐션과 유사한 기술을 사용할 수 있다는 추측이 있습니다.
  • 긴 컨텍스트 모델: 플래시 어텐션은 전체 책이나 긴 비디오와 같은 매우 긴 컨텍스트를 처리할 수 있는 새로운 모델 세대를 가능하게 했습니다.

플래시 어텐션: 최근 개발

표준 어텐션 대 플래시 어텐션

표준 어텐션 대 플래시 어텐션

플래시 어텐션-2

원래 플래시 어텐션의 성공을 기반으로, 동일한 팀은 2023년에 플래시 어텐션-2를 도입했습니다. 이 업데이트된 버전은 몇 가지 개선점을 제공합니다:

  1. 추가 최적화: 플래시 어텐션-2는 최대 70%의 이론적 피크 FLOPS를 달성하여 A100 GPU에서 더욱 나은 GPU 활용도를 달성합니다.
  2. 개선된 역방향 전달: 역방향 전달이 거의 전방향 전달과 같은 속도로 최적화되어 훈련에重大한 속도 향상을 제공합니다.
  3. 다양한 어텐션 변형 지원: 플래시 어텐션-2는 그룹 쿼리 어텐션 및 다중 쿼리 어텐션과 같은 다양한 어텐션 변형을 확장하여 지원합니다.

플래시 어텐션-3

2024년에 출시된 플래시 어텐션-3은 이 연구 라인의 최신 발전을 나타냅니다. 이는 성능을 더욱 개선하기 위한 몇 가지 새로운 기술을 도입합니다:

  1. 비동기 계산: 새로운 GPU 명령어의 비동기적인 특성을 활용하여 서로 다른 계산을 중첩합니다.
  2. FP8 지원: 저정밀 FP8 계산을 사용하여 계산을 더욱 빠르게 처리합니다.
  3. 비일관성 처리: 저정밀 형식에서 양자화 오류를 줄이는 기술입니다.

다음은 플래시 어텐션-3이 비동기 계산을 어떻게 활용하는지 보여주는 간단한 예입니다:

import torch
from torch.cuda.amp import autocast

<p>def flash_attention_3(Q, K, V, block_size=256):
with autocast(dtype=torch.float8): # FP8을 사용한 계산
# ... (이전 구현과 유사)</p>

<p># 비동기 계산 예
with torch.cuda.stream(torch.cuda.Stream()):
# GEMM을 비동기적으로 계산
S_block = torch.matmul(Q_block, K_block.transpose(-2, -1)) / (d_model ** 0.5)</p>

<p># 기본 스트림에서:
# 소프트맥스 계산을 준비합니다</p>

<p># 스트림을 동기화합니다
torch.cuda.synchronize()</p>

<p># 소프트맥스 및 출력 계산을 계속합니다
# ...</p>

return O

이 코드 스니펫은 플래시 어텐션-3이 비동기 계산과 FP8 정밀도를 어떻게 사용하는지 보여줍니다. 이는 실제 구현보다 훨씬 더 단순화된 예입니다.

프로젝트에서 플래시 어텐션 구현

플래시 어텐션을 자신의 프로젝트에서 활용하려면 몇 가지 옵션이 있습니다:

  1. 기존 라이브러리 사용: 많은 인기 있는 라이브러리에서 이미 플래시 어텐션을 구현하고 있습니다. 최신 버전으로 업데이트하고 적절한 플래그를 설정하면 충분할 수 있습니다.
  2. 사용자 정의 구현: 더 많은 제어 또는 특수한 사용 사례를 위해, 플래시 어텐션을 직접 구현할 수 있습니다. xformers 라이브러리는 좋은 참조 구현을 제공합니다.
  3. 하드웨어 특정 최적화: 특정 하드웨어(예: NVIDIA H100 GPU)를 사용하는 경우, 하드웨어 특정 기능을 최대한 활용하여 성능을 최적화할 수 있습니다.

다음은 허깅페이스 트랜스포머 라이브러리에서 플래시 어텐션을 사용하는 예입니다:

from transformers import AutoModel, AutoConfig

<p># 플래시 어텐션 활성화
config = AutoConfig.from_pretrained(&quot;bert-base-uncased&quot;)
config.use_flash_attention = True</p>

<p># 플래시 어텐션을 사용하는 모델 로드
model = AutoModel.from_pretrained(&quot;bert-base-uncased&quot;, config=config)</p>

<p># 모델을 일반적으로 사용합니다
# ...

도전과 미래 방향

플래시 어텐션이 트랜스포머 효율성을 개선하는 데重大한 발전을 이루었지만, 여전히 도전과 미래 연구 방향이 있습니다:

  1. 하드웨어 특이성: 현재 구현은 특정 GPU 아키텍처에 최적화되어 있습니다. 이러한 최적화를 다른 하드웨어에 일반화하는 것은 여전히 도전입니다.
  2. 다른 기술과의 통합: 플래시 어텐션을 다른 최적화 기술(예: 가지치기, 양자화, 모델 압축)과 결합하는 것은 활발한 연구 영역입니다.
  3. 다른 도메인으로의 확장: 플래시 어텐션이 컴퓨터 비전 및 멀티모달 모델과 같은 다른 도메인으로 확장되는 것은 진행 중인 노력입니다.
  4. 이론적 이해: 플래시 어텐션이 왜如此 잘 작동하는지에 대한 더 깊은 이해는 더욱 강력한 최적화를 이끌어낼 수 있습니다.

결론

플래시 어텐션은 현대 GPU의 메모리 계층을巧妙하게 활용하고 수학적 트릭을 사용하여, 정확한 어텐션을 계산하면서도 메모리 사용량과 계산 효율성을 크게 개선합니다.

이 기사에서 탐구한 바와 같이, 플래시 어텐션의 영향은 단순한 최적화 기술을 넘어섭니다. 이는 더욱 강력하고 효율적인 모델을 개발하는 데 기여했습니다.

지난 5년 동안私は Machine Learning과 Deep Learning의 매혹적인 세계에 몰두해 왔습니다.私の熱情と専門知識は私を50以上의多様한 소프트웨어 엔지니어링 프로젝트에 기여하게 했으며, 특히 AI/ML에 중점을 두었습니다.私の継続的な 호기심은 또한私를自然어 처리로 끌어들였습니다.私は이 분야를さらに 탐구하기를熱望합니다.