AI VIDEO BRIEFING

플래시 어텐션 원리 쉽게 이해하기: GPU 메모리 계층과 온라인 소프트맥스로 어텐션 가속

트랜스포머의 어텐션을 수 배 빠르게 만드는 플래시 어텐션의 원리를 정리했다. 어텐션의 O(N제곱) 병목, GPU 메모리 계층, 온라인 소프트맥스, 그리고 대역폭 한계를 이용한 커널 융합을 쉽게 설명한다.

플래시 어텐션은 왜 빠른가: 수학이 GPU 메모리에 순응할 때 영상 대표 이미지

핵심 메시지

  • 플래시 어텐션은 근사나 정확도 희생 없이 어텐션을 계산하는 더 빠른 방법으로, RTX 3090에서도 4배 이상의 속도 향상을 준다.
  • 원래 어텐션은 시퀀스 길이에 대해 O(N제곱)로 확장돼, 거대한 가중치 행렬을 GPU의 느린 HBM 메모리에 통째로 올려야 하는 것이 병목이다.
  • 핵심 통찰은 '입출력(IO) 인식'이다. 현대 GPU는 연산보다 데이터 이동에 발목이 잡히는 대역폭 병목 상태이므로, 계산을 조금 늘리더라도 메모리 왕복을 줄이는 편이 이득이다.
  • 온라인 소프트맥스로 실행 중 최댓값·합계를 갱신하고 연산을 하나로 융합하면, 거대한 가중치 행렬을 만들지 않고 어텐션을 선형 메모리로 처리할 수 있다.

쉽게 이해하기

트랜스포머는 게임 프레임 생성부터 추천 알고리즘까지 우리 삶 곳곳에 들어와 있고, 그 핵심인 어텐션은 입력 토큰들이 서로 정보를 주고받게 해준다. 그런데 이 연산의 원래 구현은 매우 느렸다. 영상은 2017년 어텐션 논문에 필적하는 진전으로 평가받는 플래시 어텐션이 왜, 어떻게 빠른지를 단계적으로 풀어낸다.

느림의 원인은 시퀀스 길이에 있다. 은닉 차원은 대체로 1만 안팎에 머물지만 시퀀스 길이는 수백만까지 커졌는데, 가중치 행렬은 시퀀스 길이의 제곱으로 커진다. 예컨대 100만 x 100만 행렬을 bfloat16으로 담으면 그 하나만으로 2테라바이트가 넘는 VRAM이 필요하다. 시퀀스가 길어질수록 이 행렬을 다루는 일이 학습·추론의 병목이 된다.

GPU 메모리는 계층 구조다. 모든 데이터가 담기는 HBM(VRAM)은 용량은 크지만 대역폭이 가장 낮고, 모든 SM이 공유하는 L2 캐시는 더 빠르며, 각 SM 전용인 L1 캐시는 가장 빠르지만 용량이 수백 킬로바이트로 작다. H100 기준 HBM은 초당 3테라바이트, L1은 33테라바이트로 10배 이상 차이가 난다. 그래서 가능하면 L1을, 안 되면 L2를, 정말 어쩔 수 없을 때만 HBM을 쓰는 것이 원칙이다.

파이토치는 기본적으로 즉시 실행(eager) 모드라, 연산마다 피연산자를 HBM에서 꺼내 계산하고 중간 결과를 다시 HBM에 쓴다. 거대한 가중치 행렬을 반드시 실체화해야 한다는 점이 문제의 핵심이다. 해법은 원본 Q·K·V의 값을 캐시에 최대한 올려 그 안에서 모든 연산을 마치고 최종 결과만 HBM에 한 번 쓰는 것, 즉 여러 연산을 하나로 '융합'해 거대한 중간 결과를 만들지 않는 것이다.

이를 가능하게 하는 것이 온라인 소프트맥스다. 실행 중 최댓값과 합계를 버퍼로 추적하며, 새 최댓값이 나오면 이전 합계를 지수 규칙으로 한 번에 보정한다. 여기에 값 행렬과의 곱을 함께 갱신하면 어텐션 전체를 한 번의 순회로 처리할 수 있다. 연산량은 오히려 조금 늘지만, 대역폭 병목 상태에서는 데이터 이동을 크게 줄이는 쪽이 실제 벽시계 시간을 단축한다. 실제 사용은 파이토치의 scaled_dot_product_attention 함수 한 줄로 끝난다.

주요 인사이트

  • 어떤 규모로든 돌아갈 알고리즘은 그것을 실행할 하드웨어의 제약과 조화롭게 설계돼야 한다는 것이 이 이야기의 가장 큰 교훈이다.
  • 플래시 어텐션은 연산을 줄이지 않는다. 오히려 늘린다. 대신 데이터 이동을 크게 줄여, 연산보다 대역폭에 발목 잡힌 현대 GPU에서 '산술 강도'를 높이는 방식으로 빨라진다.
  • 메모리 복잡도가 O(N제곱)에서 선형으로 바뀌는 것이 실제 속도 향상의 근원이며, 이는 근사 없이 정확한 결과를 준다.
  • Torch Compile은 선형-ReLU 같은 단순한 융합은 잘 찾지만 플래시 어텐션 같은 고급 융합은 인식하지 못해, 파이토치에 명시적으로 알려줘야 한다.

자주 묻는 질문

플래시 어텐션은 정확도를 희생하나요?

아니요. 근사나 편법 없이 정확히 같은 어텐션을 계산합니다. 다만 계산 순서를 재구성해 GPU 메모리 왕복을 줄임으로써 속도를 얻습니다. RTX 3090에서도 4배 이상 빨라집니다.

왜 연산을 더 하는데 더 빨라지나요?

현대 GPU는 연산 능력은 뛰어나지만 메모리 대역폭이 이를 따라가지 못해, SM이 데이터를 기다리는 시간이 깁니다. 플래시 어텐션은 연산을 조금 늘리는 대신 느린 HBM으로의 데이터 이동을 크게 줄여, 대역폭 병목을 완화하기 때문에 실제 시간이 단축됩니다.

실제 코드에서는 어떻게 쓰나요?

파이토치의 scaled_dot_product_attention 함수 한 줄로 대체하면 됩니다. GPU가 지원하면 기본 백엔드로 플래시 어텐션(정확히는 파이토치가 쓰는 FlashAttention 2)이 선택되고, 아니면 호환 백엔드가 사용됩니다.

원문과 출처

이 글은 원본 영상의 자막을 바탕으로 한국어 독자를 위해 요약했습니다. 전체 맥락과 최신 정보는 원문에서 확인하세요.

YouTube 원본 영상 보기 ↗

관련 AI 소식