[논문리뷰] FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness

2026. 1. 7. 19:34논문리뷰

FlashAttention: 빠르고 메모리 효율적인 정확한 Attention 연산

제목에서 알 수 있듯이, FlashAttention은 GPU 계층에서 IO를 줄여 빠르고 메모리 효율적인 정확한 Attention 연산을 가능하게 하는 방법을 제시한 논문입니다.

Transformer 모델은 자연어 처리, 이미지 분류 등 다양한 응용 분야에서 가장 널리 사용되는 아키텍처로 자리 잡았고, 점점 더 크고 깊어졌습니다. 그러나 여전히 긴 컨텍스트를 처리하는 데 어려움이 있습니다. 이는 Transformer의 self-attention 모듈이 시퀀스 길이에 따라 시간 및 메모리 복잡도가 제곱으로 증가하기 때문에 긴 시퀀스에서 처리 속도가 느리고 메모리 소모가 많기 때문입니다.

기존의 Attention 방법들은 계산 복잡도를 줄이기 위해 모델 품질을 희생하며 이 문제를 해결하려 했지만, 실질적인 시간 단축은 이루지 못하는 경우가 많았습니다. 이는 FLOP(부동소수점 연산) 감소에만 집중하고, 메모리 접근(IO)으로 인한 오버헤드를 간과한 결과입니다. 저자는 Attention 알고리즘이 GPU 메모리 레벨 간의 읽기 및 쓰기 과정을 고려하지 않았다고 판단하고, FlashAttention을 제안합니다.

 

FlashAttention은 GPU 고대역폭 메모리(HBM)와 GPU 온칩 SRAM 사이의 메모리 읽기/쓰기 횟수를 줄여, 빠르고 정확한 Attention 연산을 가능하게 하여 학습 및 추론 속도를 크게 향상시킬 수 있는 방법을 제시합니다.

아래 이미지는 HuggingFace에서 제시한 논문의 핵심 내용을 정리한 그림입니다.

기존의 Standard Attention 메커니즘은 고대역폭 메모리(HBM)을 통해 키, 쿼리 및 값을 읽고 씁니다. 이 방식에서 HBM에서 키, 쿼리 및 값을 로드하고 쓰는 비용이 높습니다. 또한, HBM에서 GPU 온칩 SRAM으로 데이터를 로드하고 다시 HBM에 쓰는 과정을 모든 Attention 단계에서 반복해야 합니다. 반면, FlashAttention은 키, 쿼리, 값을 한 번만 로드하고, Attention 연산을 누적하여 융합한 후 다시 사용할 수 있게 만들어, HBM과 GPU 온칩 SRAM 사이의 메모리 I/O 횟수를 줄여 빠르고 정확한 Attention 연산을 가능하게 합니다.

예를 들어, A100 GPU는 40-80GB의 고대역폭 메모리(HBM)를 갖추고 있으며, 이를 통해 1.5-2.0TB/s의 대역폭을 제공합니다. 텐서가 to device를 통해 HBM에 저장되고, 연산에 사용될 때 HBM에서 on-chip SRAM으로 데이터를 읽어와 실질적인 연산을 수행합니다. 연산이 끝난 후, 결과는 다시 HBM에 쓰여집니다.

on-chip SRAM은 SM(Streaming Multiprocessor)당 192KB의 용량을 가지며, 약 19TB/s의 대역폭을 제공합니다. 이는 HBM보다 훨씬 빠른 속도로 데이터를 처리할 수 있습니다.

계산 성능이 메모리 속도를 초과하게 되면서, 작업들은 점점 메모리(HBM) 접근에 의해 병목 현상이 발생하고 있습니다. 이에 따라, HBM 접근의 병목을 피하고 SRAM을 효율적으로 활용하는 것이 점점 더 중요해지고 있습니다.

 

Q, K, V는 크기가 N×d인 행렬로, 이 행렬들은 HBM에 저장됩니다.

  1. Step 1: Q와 K를 HBM에서 불러오고, Q와 K의 행렬 곱(QKᵀ)을 계산하여 N×N 크기의 Attention score S를 구합니다. 이 결과는 다시 HBM에 저장됩니다.
  2. Step 2: HBM에서 Attention score인 S를 다시 읽어와 softmax(P)를 계산하고, 그 결과를 HBM에 기록합니다.
  3. Step 3: P와 V 행렬을 HBM에서 읽어온 후, P와 V의 행렬 곱을 계산하여 최종 Attention 출력 O를 구합니다. 계산된 O는 다시 HBM에 저장됩니다.

여기서 모든 주요 데이터(입력 Q, K, V 및 중간 결과 S, P)는 HBM에 저장됩니다. 그러나 이 알고리즘의 큰 단점 중 하나는 HBM에서 데이터를 읽고 쓰는 횟수가 많다는 점입니다. 특히, 중간 결과인 S와 P는 N×N 크기의 행렬로, O(N²) 크기의 행렬이 읽기/쓰기되어야 합니다.

따라서 저자들은 S와 P와 같은 중간값들을 HBM에서 읽고 쓰지 않고, 효율적으로 Attention 연산을 수행할 수 있는 방법에 대해 고민하고 해결책을 제시합니다.

 

저자들은 Attention 연산에서 memory-bound 연산 비율이 높은지, compute-bound 연산 비율이 높은지를 측정해보았습니다. GPT-2 모델에서 실험한 결과는 아래의 그래프와 같은 결과를 보였습니다.

실험 결과, Q와 K의 matrix multiplication과 마지막에 Attention logit과 V의 matrix multiplication보다 masking, softmax, dropout 등 memory-bound에 소요되는 시간이 더 크다는 것을 확인했습니다.

따라서 저자들은 HBM과 GPU on-chip SRAM 사이의 메모리 I/O 횟수를 줄여 빠르고 정확한 Attention 연산을 가능하게 하는 방법을 제안합니다. 그들은 Tiling Softmax라는 기법을 제안하며, 이를 통해 입력을 Block 단위로 나누고 여러 번에 걸쳐 Softmax 연산을 순차적으로 수행하여, Recomputation을 통해 메모리 사용량과 접근 횟수를 줄이는 방안을 제시합니다.

 

먼저, 저자들은 대규모 시퀀스 데이터를 효율적으로 처리하는 방법으로 Tiling Softmax를 제안합니다.

SRAM은 연산만을 위한 메모리 계층이기 때문에 용량이 크지 않습니다. 따라서 매우 큰 행렬을 연산하려면 Q와 K 행렬을 한 번에 모두 SRAM으로 가져올 수 없습니다. 이를 해결하기 위해, GPU는 Q와 K 행렬을 Block 단위로 나누어 SRAM으로 가져오는 방식으로 연산을 진행합니다.

Softmax 연산은 전체 입력에 대한 값을 필요로 하는 연산입니다. 즉, Softmax 수식을 보면 S의 일부분만 있어도 분자는 구할 수 있지만, 전체 S의 합인 분모를 구하려면 전체 S가 필요합니다. 기존의 Softmax에서는 전체 S를 HBM에 저장하고 다시 SRAM으로 읽어오기 때문에 문제가 되지 않지만, Tiling Softmax에서는 전체 행렬을 Block 단위로 나누어 연산하므로 전체 S에 대한 합을 계산해야 하는 문제가 발생합니다.

저자들은 이 문제를 해결할 수 있는 방법을 제안합니다.

 

먼저, Tensor를 여러 개의 Block으로 나누어 연산할 때, 하나의 Block에서의 Attention 연산 과정을 수식으로 나타내면 다음과 같습니다.

하나의 Block에 대한 벡터 𝑥가 주어지면, 벡터 𝑥의 최대값을 찾는 함수 𝒎(𝒙)를 통해 최대값을 구합니다. 이때 𝑓(𝑥)는 Softmax의 분자값으로, 벡터 𝑥의 각 요소에서 𝒎(𝑥) 값을 빼고(Normalization) Exponential 변환을 진행합니다. 이렇게 하면 𝑓(𝑥) 값이 계산될 때, 너무 큰 값으로 인해 발생할 수 있는 Overflow를 방지할 수 있습니다. (이는 Safe Softmax에서 활용되는 연산입니다.)

그 후, 𝑙(𝑥)는 𝒎(𝑥)를 사용하여 Normalization된 벡터 𝑥의 모든 요소의 합을 구합니다. 따라서 벡터의 모든 요소의 합인 𝑓(𝑥)로, 하나의 Block 내에서 이루어지는 Softmax 연산의 결과를 얻을 수 있습니다.

 

각 Block이 Softmax를 개별적으로 수행했을 때, 이를 합쳐 전체 행렬에 대한 Softmax 값을 구할 수 있는지에 대한 연산 과정을 설명합니다.

주어진 두 개의 Block 𝑥^((1))과 𝑥^((2))는 이미 각 Block에서 Softmax를 수행했다고 가정합니다. 전체 벡터 𝑥의 최댓값인 𝒎(𝑥)는 𝑚(𝑥^((1)))과 𝑚(𝑥^((2)))의 최댓값을 이용해 간단하게 구할 수 있습니다.

이후 𝑓(𝑥) 연산은 이전 페이지에서 설명한 방식과 동일하게 𝑚(𝑥^((1)))과 𝑚(𝑥^((2)))을 사용해 𝑓(𝑥^((1)))과 𝑓(𝑥^((2)))를 Normalization한 후, 이를 다시 𝒎(𝑥)으로 Denormalization하여 새로운 𝑓(𝑥) 값을 구합니다.

마찬가지로 𝒍(𝑥)도 각 Block에서 계산된 𝑙(𝑥^((1)))과 𝑙(𝑥^((2))) 값을 사용하여 Normalization된 값을 구한 뒤, Denormalization하고 𝒎(𝑥)으로 다시 Normalization하여 두 값을 더한 후, 전체 Softmax의 분모 값인 𝒍(𝑥)을 구합니다.

이 과정을 통해 각 Block이 개별적으로 Softmax 연산을 하더라도, 두 Block을 합친 전체 Block에 대한 Softmax 값을 구할 수 있음을 알 수 있습니다. 또한, 이 과정을 통해 𝑚(𝑥), 𝑓(𝑥), 𝑙(𝑥) 값만 있으면 전체 벡터 𝑥의 Softmax 값을 쉽게 계산할 수 있습니다.

이렇게 Block 단위로 계산을 진행하면, 큰 벡터를 한 번에 처리하지 않아도 되어 메모리 효율성을 높이면서도 계산 복잡도를 줄일 수 있습니다.

 

1~4번 과정은 해당 블록의 크기와 개수를 설정하고 초기화하는 과정입니다.

먼저, HBM(High Bandwidth Memory)에 위치하는 𝑁 x 𝑑 크기의 𝑄, 𝐾, 𝑉 행렬과 𝑀 크기의 SRAM이 존재합니다.

그리고 Block 크기는 SRAM 메모리 크기를 고려하여, 열 방향으로 Block 크기를 설정합니다. 최종적으로는 𝐾와 𝑉 행렬에 대해 𝐵_𝑐 x 𝑑 크기의 Block이 총 𝑇_𝑐 개 생성됩니다. 또한, O와 Q에 대해서는 𝐵_𝑟 x 𝑑 크기의 Block이 총 𝑇_𝑟 개 생성됩니다.

 

이후 각 Block에 대해 이중 루프가 발생합니다. 하나의 𝐾, 𝑉 Block에 대해 𝑂, 𝑄, 𝑚, 𝑙 값이 𝑇_𝑟번 반복됩니다.

즉, 𝐾, 𝑉가 한 번 반복될 때마다, 𝑂, 𝑄, 𝑚, 𝑙 값은 𝑇_𝑟번 반복되며, 해당하는 인덱스의 Block을 SRAM으로 로드하면서 연산이 수행됩니다.

하나의 𝐾, 𝑉 Block에 대해 이 과정은 다음과 같이 진행됩니다.

 

SRAM으로 읽어온 값을 이용해, 𝑄_𝑖와 전치된 𝐾_𝑗의 matrix multiplication을 계산하여, 해당 𝑗번째 𝐾의 Block에 대한 𝑖번째 𝑄_𝑖 Block의 결과인 𝐵_𝑟 x 𝐵_𝑐 크기의 Attention Score 𝑆_𝑖𝑗를 구합니다.

 

이전에 설명한 한 Block에 대한 Softmax 연산을 수행하는 과정입니다.

먼저, 𝑗번째 𝐾_𝑗의 Block에 대해 𝑖번째 𝑄_𝑖와의 matrix multiplication을 통해 계산된 Attention Score인 𝑆_𝑖𝑗에서 각 행의 최대값인 𝑚̃_𝑖𝑗를 구합니다.

그 후, 𝑆_𝑖𝑗에서 최대값 𝑚̃_𝑖𝑗를 뺀 후, Exponential을 적용하여 Normalization을 진행합니다. 이때 계산되는 값은 Softmax의 분자값인 𝑃̃_𝑖𝑗입니다.

마지막으로, rowsum을 통해 𝑃̃_𝑖𝑗의 모든 값을 더하여 Softmax의 분모값인 𝑙̃_𝑖𝑗를 구합니다.

이로써 𝑗번째 𝐾_𝑗의 Block에 대해 𝑖번째 𝑄_𝑖를 가지고 한 Block 내에서의 Softmax 연산을 수행하는 과정을 설명하였습니다.

 

각 𝑄_𝑖에 대해, 이전 Block까지 계산된 값인 𝑚_𝑖와 𝑙_𝑖, 그리고 현재 Block의 𝑚̃_𝑖𝑗와 𝑙̃_𝑖𝑗를 합쳐 새로운 𝑚_𝑖^𝑛𝑒𝑤와 𝑙_𝑖^𝑛𝑒𝑤를 업데이트하는 과정입니다.

10번째 줄까지는 각 Block에 대한 연산이 이루어졌다면, 이제는 이전 Block들과 통합하는 과정을 진행합니다.

지금까지 계산한 𝑚_𝑖와 𝑚̃_𝑖𝑗 중 더 큰 값을 선택하여 𝑚_𝑖^𝑛𝑒𝑤를 업데이트합니다. 그 후, 𝑙_𝑖^𝑛𝑒𝑤를 계산합니다. 𝑙_𝑖는 이미 이전에 𝑚_𝑖로 Normalization이 된 값이므로, 𝑚_𝑖를 Denormalization한 후 𝑚_𝑖^𝑛𝑒𝑤로 다시 Normalization하여 새로운 𝑙_𝑖^𝑛𝑒𝑤를 얻습니다.

𝑙̃_𝑖𝑗도 10번째 줄에서 𝑚̃_𝑖𝑗로 Normalization이 된 값이므로, 𝑚̃_𝑖𝑗를 Denormalization한 후, 𝑚_𝑖^𝑛𝑒𝑤로 다시 Normalization하고 두 값을 더하여, 이전 모든 Block들의 합과 현재 Block의 원소들을 모두 합산하여 현재 Block까지의 모든 원소의 합을 구합니다.

 

이제 𝑖번째 Block까지의 Attention 연산의 최종 출력인 𝑂_𝑖에 대한 계산을 진행합니다.

  • 빨간 박스: 𝑖번째 Block까지의 전체 Softmax의 분모값
  • 파란 박스: 𝑖번째 Block까지의 전체 Softmax의 분자값 (새로 업데이트된 최대값 𝑚_𝑖^𝑛𝑒𝑤로 Normalization된 값)

먼저, 분모값은 𝑙_𝑖^𝑛𝑒𝑤로, 11번째 줄에서 이미 업데이트된 𝑙_𝑖^𝑛𝑒𝑤를 그대로 사용합니다.

분자값은 𝑑𝑖𝑎𝑔(𝑙_𝑖) 𝑒^(𝑚_𝑖 − 𝑚_𝑖^𝑛𝑒𝑤)}로, 𝑂_𝑖는 이전 Block까지의 Softmax 값이기 때문에, 𝑒^(𝑚_𝑖 − 𝑚_𝑖^𝑛𝑒𝑤) 부분에서 이전에 사용한 Normalization 상수 𝑚_𝑖를 Denormalization한 후, 11번째 줄에서 업데이트된 𝑚_𝑖^𝑛𝑒𝑤로 Normalization을 다시 진행합니다.

또한, Softmax의 식을 보면 𝑓(𝑥) / 𝑙(𝑥) 형태입니다. 이때 𝑂_𝑖는 𝑓(𝑥) / 𝑙(𝑥) 형태의 Softmax에 𝑉와 matrix multiplication한 결과입니다. 따라서 이전 Block까지 사용된 분모값이 곱해져 있기 때문에, 이전 분모값인 𝑙(𝑥)를 다시 곱해줘야 이전 Block까지의 분자값인 𝑓(𝑥)만 남게 됩니다.

현재 Block에 대해 계산된 Softmax의 분자값인 𝑃̃_𝑖𝑗 또한 현재 Block의 Normalization 상수 𝑚_𝑖^𝑛𝑒𝑤로 Denormalization 후, 업데이트된 𝑚_𝑖^𝑛𝑒𝑤로 다시 Normalization을 진행한 후, 𝑉와 matrix multiplication을 통해 현재 Block까지의 Softmax 값을 구합니다.

이렇게 계산된 현재 Block의 최종값인 𝑂_𝑖를 업데이트한 후, SRAM에서 HBM으로 Write합니다.

이 연산은 𝐾_1, 𝑉_1에 대해 𝑇_𝑟번 반복하고, 이를 𝑇_𝑐번 반복하여 총 𝑇_𝑟 x 𝑇_𝑐번 반복합니다. 마지막으로, 모든 Block에 대한 값인 𝑂를 return합니다.

 

Standard Attention의 메모리 접근 횟수에 대한 설명입니다.

Attention Matrix인 𝑁 x 𝑑 크기의 Q, K, 𝑉에서 𝑂(𝑁𝑑) 만큼 I/O가 발생합니다.

또한, 중간 결과로 생성되는 𝑆, 𝑃 행렬에 의해 추가적으로 𝑂(𝑁²) 크기의 메모리가 생성됩니다.

따라서, 총 Standard Attention에서 발생하는 I/O는 𝑂(𝑁𝑑 + 𝑁²) 만큼 발생합니다.

 

FlashAttention의 메모리 접근 횟수에 대한 설명입니다.

  • 빨간 박스: 𝐾와 𝑉는 𝑗번째 Block이 𝑇_𝑐번 반복되므로 𝑁 x 𝑑 크기만큼 I/O가 발생합니다.
  • 초록 박스: 𝑙과 𝑚 벡터는 각각 𝑁 크기의 Scalar 값을 담고 있으며, 𝑇_𝑐번 반복되므로 𝑇_𝑐 x 𝑁만큼 I/O가 발생합니다.
  • 파란 박스: 𝑁 x 𝑑 크기의 벡터가 𝑇_𝑐번 반복되므로 𝑇_𝑐 x 𝑁𝑑만큼 I/O가 발생합니다.

이들을 종합하면, 𝐵_𝑐 = [𝑀 / 4𝑑]이고, 𝑇_𝑐 = [𝐵_𝑐 / 𝑀]이므로 𝑇_𝑐 = [𝑁𝑑 / 𝑀]로 정리됩니다. 따라서 최종적으로 FlashAttention의 메모리 접근 횟수는 𝑂((𝑁²𝑑²) / 𝑀)로, 이전의 Standard Attention과 동일하게 𝑁²만큼의 접근이 발생합니다.

하지만 저자들은 𝑑²가 𝑀보다 훨씬 작으며, 실험적으로 훨씬 빠른 성능을 확인했다고 보고하였습니다.

 

저자들은 HBM 접근 횟수가 Attention 실행 시간의 주요 결정 요인임을 확인했습니다. 그림을 보면 FlashAttention이 매 Block마다 Standard Attention보다 더 많은 FLOP 수를 가짐에도 불구하고, 훨씬 적은 HBM 접근 횟수를 가지므로 실행 시간이 훨씬 빠르다는 사실을 알 수 있습니다.

또한, 오른쪽 그래프는 FlashAttention의 Block 크기 𝐵_𝑐를 변화시키면서 HBM 접근 횟수와 실행 시간을 측정한 결과를 보여줍니다. Block 크기가 커질수록 HBM 접근 횟수가 줄어들어 실행 시간이 감소하는 경향을 보였습니다. 하지만, Block 크기가 충분히 커지면(256 이상) 실행 시간은 다른 요인(산술 연산)에 의해 병목 현상을 겪게 됩니다. 또한, Block 크기가 커질수록 작은 SRAM 크기에 맞지 않게 되어 효율성이 떨어지는 문제가 발생합니다.

 

GPT-2의 훈련 시간에 대한 실험 그래프를 보면, FlashAttention을 사용한 GPT-2는 HuggingFace 및 Megatron-LM 구현보다 훨씬 더 빠른 학습 시간을 기록한 것을 확인할 수 있습니다. FlashAttention은 HuggingFace와 Megatron-LM 대비 약 3배 이상의 학습 속도 향상을 보였습니다.

또한, FlashAttention은 모델의 정의된 구성을 변경하지 않으므로, 다른 구현들과 동일한 언어 모델의 성능을 평가하는 지표인 perplexity도 동일하게 유지되었습니다.

Perplexity는 언어 모델의 성능을 평가하는 지표로, 모델이 얼마나 정확하게 예측을 하는지를 나타냅니다. 수학적으로는 확률의 역수로 해석되며, 낮을수록 모델이 더 정확하게 예측한다는 의미입니다. 언어 모델이 텍스트의 다음 단어를 예측할 때, 확률이 높은 단어들에 집중할수록 perplexity는 낮아집니다.

Perplexity는 또한 모델이 얼마나 불확실한 상태에 있는지를 나타내며, 특정 문맥에서 가능한 모든 단어의 분포를 예측하는 데 사용됩니다. 예를 들어, 언어 모델이 완벽하게 예측한다면, perplexity는 1에 가깝게 나옵니다.

 

아래는 제가 참고한 자료에 대한 url입니다.
논문 paper: https://arxiv.org/abs/2205.14135

 

FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness

Transformers are slow and memory-hungry on long sequences, since the time and memory complexity of self-attention are quadratic in sequence length. Approximate attention methods have attempted to address this problem by trading off model quality to reduce

arxiv.org

GitHub: https://github.com/Dao-AILab/flash-attention

 

GitHub - Dao-AILab/flash-attention: Fast and memory-efficient exact attention

Fast and memory-efficient exact attention. Contribute to Dao-AILab/flash-attention development by creating an account on GitHub.

github.com