어텐션 메모리 계산기
입력
| 시퀀스 길이 | 4,096 |
|---|---|
| 어텐션 헤드 수 | 32 |
| 배치 크기 | 1 |
| 정밀도 | FP16 / BF16 (2바이트) |
어텐션 메모리 계산기
표준 트랜스포머 어텐션에서 실체화된 어텐션 점수 행렬의 메모리를 시퀀스 길이, 헤드 수, 배치 크기, 요소당 바이트로부터 추정합니다. FlashAttention이 제거하는 바로 그 제곱 항입니다.
입력
워크로드
결과
값을 입력하면 계산 결과가 표시됩니다.
어텐션 메모리
표준 트랜스포머 어텐션은 모든 토큰을 다른 모든 토큰과 비교해, 소프트맥스 이전에 점수 행렬을 만듭니다. 이 행렬은 시퀀스 길이에 대해 정사각이므로 메모리가 컨텍스트의 제곱으로 커집니다. 이것이 그 유명한 어텐션의 제곱 비용입니다. 이 계산기는 그 실체화된 행렬의 크기를 시퀀스 길이, 어텐션 헤드 수, 배치 크기, 점수 하나당 바이트로부터 추정합니다. 이는 정확히 FlashAttention이 저장을 피하는 메모리입니다.
제곱 항
개의 토큰으로 이루어진 시퀀스에 대해 어텐션은 점수 행렬을 만듭니다. 행 열에는 토큰 가 토큰 를 얼마나 어텐션하는지가 담깁니다. 이를 저장하는 데에는 에 비례하는 메모리가 듭니다. 컨텍스트에 선형으로 커지는 모델 가중치나 키-값 캐시와 달리, 점수 행렬은 그 제곱으로 커지므로 긴 컨텍스트에서는 가장 큰 중간 버퍼가 되며, 이것이 단순한 어텐션이 메모리를 소진하는 이유입니다.
공식
각 헤드는 자기 점수 행렬을 만들고 배치의 각 시퀀스는 자기 사본을 들고 있으므로, 바이트 단위 메모리는 다음과 같습니다.
Abytes=B⋅H⋅s2⋅e여기서 는 배치 크기, 는 어텐션 헤드 수, 는 시퀀스 길이, 는 저장 점수당 바이트입니다. 로 나누면 기가바이트가 됩니다. 긴 컨텍스트를 비싸게 만드는 것은 에 붙은 제곱입니다. 헤드 수와 배치는 선형으로만 곱해집니다.
이 메모리가 무엇인가
이 값은 표준 어텐션이 쿼리-키 곱과 값에 대한 소프트맥스 가중합 사이에 메모리에 써넣는, 실체화된 점수 행렬을 측정합니다. 이는 단순한 구현이 반드시 들고 있어야 하는 활성화 버퍼입니다. 이것이 지속되는지는 상황에 따라 다릅니다. 평범한 순전파에서는 일시적이어서 레이어 사이에서 재사용할 수 있지만, 학습 프레임워크는 체크포인팅이나 융합 커널이 개입하지 않는 한 역전파를 위해 레이어별 사본을 유지할 수 있습니다.
FlashAttention
FlashAttention은 전체 행렬을 한 번도 써넣지 않고 동일한 출력을 계산합니다. 키와 값을 작은 블록 단위로 훑으며 실행 중인 소프트맥스 통계를 유지하므로, 한 번에 한 블록만 메모리에 두면 됩니다. 이로써 저장 공간이 시퀀스 길이에 대해 제곱에서 선형으로 바뀝니다. 이 계산기가 크기를 재는 행렬이 바로 FlashAttention이 저장을 거부하는 대상이므로, 여기서 나온 결과는 융합 어텐션 커널이 주어진 컨텍스트 길이에서 절약하는 메모리의 좋은 근사치입니다.
계산 예시
4,096 토큰짜리 시퀀스 하나를 어텐션 헤드 32개, 16비트 정밀도로 처리한다고 합시다.
Abytes=1×32×40962×2=1,073,741,824한 레이어의 점수 행렬이 약 1.07 GB입니다. 컨텍스트를 8,192 토큰으로 2배 늘리면 제곱 항이 위력을 발휘해, 같은 식이 대략 4.29 GB가 됩니다. 길이가 2배인데 메모리는 4배입니다. 같은 어텐션 예산의 키-값 측면은 KV 캐시 크기 계산기에서, 모델 가중치 부분은 LLM 추론 VRAM 계산기에서 다룹니다.
자주 묻는 질문 (FAQ)
어텐션 메모리는 왜 시퀀스 길이의 제곱에 비례하나요?
어텐션은 모든 토큰을 다른 모든 토큰과 비교해, 양쪽 차원이 모두 시퀀스 길이인 점수 행렬을 만듭니다. 따라서 그 전체 행렬을 저장하는 데에는 시퀀스 길이의 제곱에 비례하는 메모리가 듭니다. 컨텍스트가 2배가 되면 점수 행렬은 4배가 되며, 이 때문에 긴 컨텍스트에서 표준 어텐션이 메모리에 묶입니다. 제곱 항이 가중치와 활성화의 선형 비용을 앞지르기 때문입니다.
FlashAttention은 이 메모리를 어떻게 줄이나요?
FlashAttention은 전체 점수 행렬을 한 번도 실체화하지 않고 같은 결과를 계산합니다. 키와 값을 작은 블록 단위로 흘려보내며 실행 중인 소프트맥스 통계를 유지하므로, 한 번에 한 블록만 들고 있으면 됩니다. 이 계산기가 측정하는 제곱 행렬이 바로 FlashAttention이 저장을 피하는 메모리이며, 그 덕분에 훨씬 긴 시퀀스로 확장할 수 있습니다. 연산 자체는 그대로이고, 중간 저장 공간만 제곱에서 선형으로 줄어듭니다.
이 메모리는 레이어당인가요, 아니면 모델 전체인가요?
구현에 따라 다릅니다. 여기서 나온 값은 한 레이어의 점수 행렬 크기입니다. 단순한 순전파에서는 그 버퍼를 레이어 사이에서 해제하고 재사용할 수 있으므로 레이어 수만큼 곱해지지 않고 일시적입니다. 다만 학습에서는 그래디언트 체크포인팅이나 FlashAttention을 쓰지 않는 한, 프레임워크가 역전파를 위해 레이어별 어텐션 텐서를 유지할 수 있으며, 그 경우 합계는 레이어 수에 비례해 커질 수 있습니다.
면책조항
이 추정치는 명시된 정밀도에서 실체화된 점수 행렬만 다루며, 쿼리·키·값 텐서, 출력 투영, 그 밖의 활성화는 제외합니다. 이는 표준 어텐션을 설명한 것으로, FlashAttention 같은 커널은 이 행렬을 아예 저장하지 않습니다.